import logging import os import threading from collections.abc import Callable from enum import StrEnum from pathlib import Path logger = logging.getLogger(__name__) class SecretSource(StrEnum): ENV = "env" FILE = "file" class SecretsRuntime: def __init__(self): self._secrets: dict[str, str] = {} self._snapshots: list[dict[str, str]] = [] self._resolvers: list[Callable[[str], str | None]] = [] self._lock = threading.Lock() def register_resolver(self, resolver: Callable[[str], str | None]) -> None: with self._lock: self._resolvers.append(resolver) def load_from_env(self, var_name: str, secret_name: str | None = None) -> None: value = os.getenv(var_name) if value is not None: name = secret_name or var_name with self._lock: self._secrets[name] = value logger.debug("Loaded secret '%s' from env var %s", name, var_name) def load_from_file(self, file_path: str | Path, secret_name: str | None = None) -> None: path = Path(file_path) if not path.exists(): logger.warning("Secret file not found: %s", path) return try: value = path.read_text(encoding="utf-8").strip() except Exception: logger.exception("Failed to read secret file: %s", path) return name = secret_name or path.stem with self._lock: self._secrets[name] = value logger.debug("Loaded secret '%s' from file %s", name, path) def snapshot(self) -> None: with self._lock: self._snapshots.append(dict(self._secrets)) def rollback(self) -> None: with self._lock: if not self._snapshots: raise RuntimeError("No snapshot available for rollback") self._secrets = self._snapshots.pop() def plan(self) -> list[str]: with self._lock: return sorted(self._secrets.keys()) def validate(self, required: list[str]) -> list[str]: with self._lock: return [name for name in required if name not in self._secrets] def apply(self, updates: dict[str, str]) -> None: with self._lock: self._snapshots.append(dict(self._secrets)) try: self._secrets.update(updates) except Exception: self._secrets = self._snapshots.pop() raise def get(self, name: str) -> str | None: with self._lock: return self._secrets.get(name) def resolve(self, value: str) -> str: with self._lock: resolvers = list(self._resolvers) for resolver in resolvers: result = resolver(value) if result is not None: return result if value.startswith("$env:"): var_name = value[5:] env_val = os.getenv(var_name) if env_val is not None: return env_val if value.startswith("$secret:"): secret_name = value[8:] with self._lock: secret_val = self._secrets.get(secret_name) if secret_val is not None: return secret_val return value def trim_credential(self, value: str | None) -> str | None: if value is None: return None if value.startswith("${") and value.endswith("}"): return None return value @property def secret_names(self) -> list[str]: with self._lock: return list(self._secrets.keys()) secrets_runtime = SecretsRuntime()