zcbot/core/provider_credentials/runtime.py

74 lines
2.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""每个新外部请求调用一次的数据库优先/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)