74 lines
2.4 KiB
Python
74 lines
2.4 KiB
Python
"""每个新外部请求调用一次的数据库优先/env fallback 凭据解析器。"""
|
||
from __future__ import annotations
|
||
|
||
import os
|
||
from dataclasses import dataclass
|
||
|
||
from sqlalchemy import select
|
||
from sqlalchemy.exc import SQLAlchemyError
|
||
|
||
from core.external_systems.crypto import decrypt_secret
|
||
|
||
from .registry import BY_ENV, get_provider
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class ResolvedCredentials:
|
||
provider_id: str
|
||
source: str
|
||
values: dict[str, str]
|
||
|
||
|
||
def _env_values(provider_id: str) -> dict[str, str]:
|
||
provider = get_provider(provider_id)
|
||
return {
|
||
field.name: value
|
||
for field in provider.fields
|
||
if (value := (os.getenv(field.env) or "").strip())
|
||
}
|
||
|
||
|
||
def resolve_credentials(provider_id: str) -> ResolvedCredentials:
|
||
provider = get_provider(provider_id)
|
||
row = None
|
||
try:
|
||
from core.storage import session_scope
|
||
from core.storage.models import ProviderCredential
|
||
with session_scope() as session:
|
||
row = session.execute(
|
||
select(ProviderCredential).where(
|
||
ProviderCredential.provider_id == provider_id
|
||
)
|
||
).scalar_one_or_none()
|
||
except (RuntimeError, SQLAlchemyError):
|
||
# CLI、migration 前部署或 DB 短暂不可用时保持历史 env 行为。
|
||
pass
|
||
if row is not None and all(field.name in row.credentials for field in provider.fields):
|
||
# 已存在完整 DB 覆盖时,master key/AAD 错误必须显式失败,不能静默绕回 env。
|
||
values = {
|
||
field.name: decrypt_secret(
|
||
row.credentials[field.name], aad=f"provider:{provider_id}:{field.name}"
|
||
)
|
||
for field in provider.fields
|
||
}
|
||
return ResolvedCredentials(provider_id, "database", values)
|
||
values = _env_values(provider_id)
|
||
source = "env" if len(values) == len(provider.fields) else "missing"
|
||
return ResolvedCredentials(provider_id, source, values)
|
||
|
||
|
||
def resolve_secret(provider_id: str, field: str = "api_key") -> str:
|
||
return resolve_credentials(provider_id).values.get(field, "")
|
||
|
||
|
||
def resolve_env_secret(env_name: str) -> str:
|
||
binding = BY_ENV.get(env_name)
|
||
if binding is None:
|
||
return (os.getenv(env_name) or "").strip()
|
||
return resolve_secret(*binding)
|
||
|
||
|
||
def provider_available(provider_id: str) -> bool:
|
||
provider = get_provider(provider_id)
|
||
return len(resolve_credentials(provider_id).values) == len(provider.fields)
|