116 lines
3.6 KiB
Python
116 lines
3.6 KiB
Python
|
|
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()
|