ForcePilot/backend/package/yuxi/channel/security/secrets_manager.py

116 lines
3.6 KiB
Python
Raw Normal View History

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()