Coverage for agentos/core/secrets.py: 0%

225 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-09 09:19 +0800

1"""AgentOS Secrets Manager — production-grade secrets lifecycle. 

2 

3Backends: 

4- EnvSecretsBackend: os.environ (12-factor) 

5- VaultSecretsBackend: HashiCorp Vault (kv-v2) 

6- EncryptedFileBackend: Fernet-encrypted JSON file 

7- CompositeSecretsBackend: layered resolution (most specific first) 

8 

9Design: ~350 lines, thread-safe, lazy loading with caching. 

10""" 

11 

12from __future__ import annotations 

13 

14import json 

15import logging 

16import os 

17import time 

18from abc import ABC, abstractmethod 

19from dataclasses import dataclass 

20from pathlib import Path 

21from threading import RLock 

22from typing import Any 

23 

24logger = logging.getLogger(__name__) 

25 

26 

27# ============================================================================ 

28# Data types 

29# ============================================================================ 

30 

31 

32class SecretNotFoundError(Exception): 

33 """Requested secret key does not exist.""" 

34 

35 

36class BackendUnavailableError(Exception): 

37 """Secrets backend is unreachable or misconfigured.""" 

38 

39 

40@dataclass 

41class SecretsConfig: 

42 """Global secrets manager configuration.""" 

43 

44 cache_ttl: float = 300.0 # Seconds to cache fetched secrets 

45 max_cache_size: int = 1000 # Max cached entries 

46 fail_open: bool = False # If True, return None on backend error instead of raising 

47 allow_environment_fallback: bool = True # Try env vars before backend 

48 

49 

50# ============================================================================ 

51# Abstract backend 

52# ============================================================================ 

53 

54 

55class AbstractSecretsBackend(ABC): 

56 """Interface for secrets backends.""" 

57 

58 @abstractmethod 

59 async def get(self, key: str) -> str | None: 

60 """Retrieve a single secret value.""" 

61 

62 @abstractmethod 

63 async def get_all(self, prefix: str = "") -> dict[str, str]: 

64 """Retrieve all secrets matching prefix.""" 

65 

66 @abstractmethod 

67 async def health_check(self) -> bool: 

68 """Verify backend connectivity.""" 

69 

70 

71# ============================================================================ 

72# Env backend 

73# ============================================================================ 

74 

75 

76class EnvSecretsBackend(AbstractSecretsBackend): 

77 """OS environment variable backend — no deps, always available.""" 

78 

79 def __init__(self, prefix: str = ""): 

80 self._prefix = prefix 

81 

82 async def get(self, key: str) -> str | None: 

83 return os.environ.get(f"{self._prefix}{key}") 

84 

85 async def get_all(self, prefix: str = "") -> dict[str, str]: 

86 full_prefix = f"{self._prefix}{prefix}" 

87 return { 

88 k[len(self._prefix) :] if k.startswith(self._prefix) else k: v 

89 for k, v in os.environ.items() 

90 if k.startswith(full_prefix) 

91 } 

92 

93 async def health_check(self) -> bool: 

94 return True 

95 

96 

97# ============================================================================ 

98# Encrypted file backend 

99# ============================================================================ 

100 

101 

102class EncryptedFileBackend(AbstractSecretsBackend): 

103 """Fernet-encrypted JSON file backend. 

104 

105 Secrets file encrypted at rest. Use `cryptography` for encryption. 

106 """ 

107 

108 def __init__(self, file_path: str, encryption_key: str): 

109 self._path = Path(file_path) 

110 self._key = encryption_key 

111 self._cache: dict[str, str] | None = None 

112 self._cache_time: float = 0.0 

113 self._lock = RLock() 

114 

115 def _decrypt(self) -> dict[str, str]: 

116 """Decrypt and load secrets file.""" 

117 from cryptography.fernet import Fernet 

118 

119 if not self._path.exists(): 

120 return {} 

121 

122 fernet = Fernet(self._key.encode() if isinstance(self._key, str) else self._key) 

123 with open(self._path, "rb") as f: 

124 encrypted = f.read() 

125 decrypted = fernet.decrypt(encrypted) 

126 return json.loads(decrypted) 

127 

128 async def get(self, key: str) -> str | None: 

129 with self._lock: 

130 if self._cache is None or (time.monotonic() - self._cache_time) > 300: 

131 try: 

132 self._cache = self._decrypt() 

133 self._cache_time = time.monotonic() 

134 except Exception as exc: 

135 logger.error("Failed to decrypt secrets file: %s", exc) 

136 raise BackendUnavailableError(f"Decryption failed: {exc}") 

137 return self._cache.get(key) if self._cache else None 

138 

139 async def get_all(self, prefix: str = "") -> dict[str, str]: 

140 await self.get("") # Refresh cache 

141 if self._cache is None: 

142 return {} 

143 if not prefix: 

144 return dict(self._cache) 

145 return {k: v for k, v in self._cache.items() if k.startswith(prefix)} 

146 

147 async def health_check(self) -> bool: 

148 try: 

149 self._decrypt() 

150 return True 

151 except Exception: 

152 return False 

153 

154 

155# ============================================================================ 

156# Vault backend (HashiCorp Vault kv-v2) 

157# ============================================================================ 

158 

159 

160class VaultSecretsBackend(AbstractSecretsBackend): 

161 """HashiCorp Vault backend (kv-v2). 

162 

163 Requires hvac library or httpx for REST calls. 

164 """ 

165 

166 def __init__( 

167 self, 

168 url: str, 

169 token: str, 

170 mount_point: str = "secret", 

171 path_prefix: str = "", 

172 verify_ssl: bool = True, 

173 ): 

174 self._url = url.rstrip("/") 

175 self._token = token 

176 self._mount_point = mount_point 

177 self._path_prefix = path_prefix 

178 self._verify_ssl = verify_ssl 

179 self._cache: dict[str, str | None] = {} 

180 self._lock = RLock() 

181 

182 async def _call(self, method: str, path: str) -> Any: 

183 """Make authenticated Vault API call.""" 

184 import httpx 

185 

186 async with httpx.AsyncClient(verify=self._verify_ssl, timeout=10.0) as client: 

187 headers = {"X-Vault-Token": self._token} 

188 url = f"{self._url}/v1/{path}" 

189 resp = await client.request(method, url, headers=headers) 

190 if resp.status_code == 404: 

191 return None 

192 if resp.status_code >= 400: 

193 raise BackendUnavailableError( 

194 f"Vault API error {resp.status_code}: {resp.text[:200]}" 

195 ) 

196 return resp.json() 

197 

198 async def get(self, key: str) -> str | None: 

199 full_path = f"{self._path_prefix}/{key}" if self._path_prefix else key 

200 vault_path = f"{self._mount_point}/data/{full_path}" 

201 

202 with self._lock: 

203 if key in self._cache: 

204 return self._cache[key] 

205 

206 try: 

207 data = await self._call("GET", vault_path) 

208 if data is None: 

209 return None 

210 value = data.get("data", {}).get("data", {}).get("value") 

211 if value is None: 

212 value = data.get("data", {}).get("data", {}) 

213 with self._lock: 

214 self._cache[key] = str(value) if not isinstance(value, dict) else json.dumps(value) 

215 return self._cache[key] 

216 except BackendUnavailableError: 

217 raise 

218 

219 async def get_all(self, prefix: str = "") -> dict[str, str]: 

220 list_path = f"{self._mount_point}/metadata/{self._path_prefix}" 

221 result: dict[str, str] = {} 

222 

223 try: 

224 data = await self._call("LIST", list_path) 

225 if data and "data" in data and "keys" in data["data"]: 

226 for key in data["data"]["keys"]: 

227 if not prefix or key.startswith(prefix): 

228 value = await self.get(key) 

229 if value is not None: 

230 result[key] = value 

231 except BackendUnavailableError: 

232 raise 

233 

234 return result 

235 

236 async def health_check(self) -> bool: 

237 try: 

238 resp = await self._call("GET", "sys/health") 

239 return resp is not None and resp.get("initialized", False) 

240 except Exception: 

241 return False 

242 

243 

244# ============================================================================ 

245# Composite backend (layered) 

246# ============================================================================ 

247 

248 

249class CompositeSecretsBackend(AbstractSecretsBackend): 

250 """Resolve secrets from multiple backends in priority order. 

251 

252 First backend that returns a non-None value wins. 

253 """ 

254 

255 def __init__(self, backends: list[AbstractSecretsBackend]): 

256 self._backends = backends 

257 

258 async def get(self, key: str) -> str | None: 

259 for backend in self._backends: 

260 try: 

261 value = await backend.get(key) 

262 if value is not None: 

263 return value 

264 except Exception as exc: 

265 logger.debug("Backend %s failed for key=%s: %s", type(backend).__name__, key, exc) 

266 return None 

267 

268 async def get_all(self, prefix: str = "") -> dict[str, str]: 

269 result: dict[str, str] = {} 

270 for backend in reversed(self._backends): # Low-priority first 

271 try: 

272 batch = await backend.get_all(prefix) 

273 result.update(batch) 

274 except Exception as exc: 

275 logger.debug("Backend %s get_all failed: %s", type(backend).__name__, exc) 

276 return result 

277 

278 async def health_check(self) -> bool: 

279 for backend in self._backends: 

280 if await backend.health_check(): 

281 return True 

282 return False 

283 

284 

285# ============================================================================ 

286# Secrets Manager (high-level) 

287# ============================================================================ 

288 

289 

290class SecretsManager: 

291 """High-level secrets manager with caching and fail-open support. 

292 

293 Usage: 

294 sm = SecretsManager(EnvSecretsBackend("MYAPP_")) 

295 api_key = await sm.get("API_KEY") 

296 # or 

297 api_key = await sm.require("API_KEY") # raises if missing 

298 """ 

299 

300 def __init__(self, backend: AbstractSecretsBackend, config: SecretsConfig = SecretsConfig()): 

301 self._backend = backend 

302 self._config = config 

303 self._cache: dict[str, tuple[float, str | None]] = {} 

304 self._lock = RLock() 

305 

306 async def _cache_put(self, key: str, value: str | None): 

307 """Store value in cache with eviction.""" 

308 with self._lock: 

309 if len(self._cache) >= self._config.max_cache_size: 

310 sorted_keys = sorted(self._cache, key=lambda k: self._cache[k][0]) 

311 for old_key in sorted_keys[: len(self._cache) // 4]: 

312 del self._cache[old_key] 

313 self._cache[key] = (time.monotonic(), value) 

314 

315 async def get(self, key: str) -> str | None: 

316 """Get secret value. Returns None if not found.""" 

317 # Check cache 

318 with self._lock: 

319 if key in self._cache: 

320 ts, val = self._cache[key] 

321 if time.monotonic() - ts < self._config.cache_ttl: 

322 return val 

323 

324 # Fallback to environment 

325 if self._config.allow_environment_fallback: 

326 env_val = os.environ.get(key) 

327 if env_val is not None: 

328 await self._cache_put(key, env_val) 

329 return env_val 

330 

331 # Query backend 

332 try: 

333 value = await self._backend.get(key) 

334 except Exception as exc: 

335 if self._config.fail_open: 

336 logger.warning("Secrets backend error for key=%s (fail_open): %s", key, exc) 

337 return None 

338 raise BackendUnavailableError(f"Failed to fetch '{key}': {exc}") from exc 

339 

340 # Update cache 

341 await self._cache_put(key, value) 

342 

343 return value 

344 

345 async def require(self, key: str) -> str: 

346 """Get secret value. Raises SecretNotFoundError if missing.""" 

347 value = await self.get(key) 

348 if value is None: 

349 raise SecretNotFoundError(f"Required secret '{key}' not found") 

350 return value 

351 

352 async def get_all(self, prefix: str = "") -> dict[str, str]: 

353 """Get all secrets with given prefix.""" 

354 try: 

355 return await self._backend.get_all(prefix) 

356 except Exception as exc: 

357 if self._config.fail_open: 

358 logger.warning("Secrets get_all error (fail_open): %s", exc) 

359 return {} 

360 raise 

361 

362 async def health_check(self) -> bool: 

363 """Check backend health.""" 

364 return await self._backend.health_check() 

365 

366 def invalidate_cache(self, key: str | None = None): 

367 """Invalidate cache entries.""" 

368 with self._lock: 

369 if key is None: 

370 self._cache.clear() 

371 elif key in self._cache: 

372 del self._cache[key] 

373 

374 

375# ============================================================================ 

376# Convenience factory 

377# ============================================================================ 

378 

379 

380def create_secrets_manager( 

381 backend_type: str = "env", 

382 **kwargs: Any, 

383) -> SecretsManager: 

384 """Factory: create SecretsManager with common backends. 

385 

386 backend_type: 'env' | 'vault' | 'encrypted_file' | 'composite' 

387 """ 

388 if backend_type == "env": 

389 backend = EnvSecretsBackend(prefix=kwargs.get("prefix", "")) 

390 elif backend_type == "vault": 

391 backend = VaultSecretsBackend( 

392 url=kwargs["url"], 

393 token=kwargs["token"], 

394 mount_point=kwargs.get("mount_point", "secret"), 

395 path_prefix=kwargs.get("path_prefix", ""), 

396 verify_ssl=kwargs.get("verify_ssl", True), 

397 ) 

398 elif backend_type == "encrypted_file": 

399 backend = EncryptedFileBackend( 

400 file_path=kwargs["file_path"], 

401 encryption_key=kwargs["encryption_key"], 

402 ) 

403 elif backend_type == "composite": 

404 backends = kwargs["backends"] 

405 backend = CompositeSecretsBackend(backends) 

406 else: 

407 raise ValueError(f"Unknown backend type: {backend_type}") 

408 

409 return SecretsManager(backend, SecretsConfig(**kwargs.get("config", {})))