"""Provider 凭据 CRUD、乐观替换、状态持久化与定时检查。""" from __future__ import annotations from collections.abc import Callable from datetime import datetime, timedelta, timezone from uuid import UUID from sqlalchemy import delete, select, update from sqlalchemy.exc import IntegrityError from core.external_systems.crypto import configured as crypto_configured from core.external_systems.crypto import encrypt_secret from core.storage import session_scope from core.storage.models import ProviderCredential from .registry import PROVIDERS, get_provider from .runtime import _env_values, resolve_credentials from .testing import TestResult, classify_response, test_provider _BAD = {"low_balance", "exhausted", "auth_error", "unreachable"} _NOTIFY_COOLDOWN = timedelta(hours=24) class ProviderCredentialError(RuntimeError): pass class RevisionConflict(ProviderCredentialError): pass def _hint(value: str) -> str: value = str(value or "") return f"***{value[-4:]}" if len(value) >= 4 else "***" def _row_payload(provider, row: ProviderCredential | None) -> dict: database = bool(row and all(field.name in row.credentials for field in provider.fields)) env_values = _env_values(provider.provider_id) env_source = "env" if len(env_values) == len(provider.fields) else "missing" source = "database" if database else env_source hints = row.credential_hint if database else { field.name: _hint(env_values.get(field.name, "")) for field in provider.fields if env_values.get(field.name) } return { "provider_id": provider.provider_id, "display_name": provider.display_name, "category": provider.category, "fields": [ {"name": field.name, "label": field.label, "required": True, "hint": hints.get(field.name, "")} for field in provider.fields ], "configured": source != "missing", "source": source, "revision": row.revision if row else 0, "test_status": row.test_status if row else "untested", "test_detail": row.test_detail if row else None, "last_tested_at": row.last_tested_at.isoformat() if row and row.last_tested_at else None, "balance": ( {"amount": str(row.balance_amount), "currency": row.balance_currency} if row and row.balance_amount is not None else None ), "balance_supported": provider.balance_supported, "billable_test": provider.billable, "alerting": bool(row and row.test_status in _BAD), } def list_providers() -> list[dict]: with session_scope() as session: rows = { row.provider_id: row for row in session.execute(select(ProviderCredential)).scalars() } return [_row_payload(provider, rows.get(provider.provider_id)) for provider in PROVIDERS] def _validate_values(provider_id: str, values: dict[str, str]) -> dict[str, str]: provider = get_provider(provider_id) expected = {field.name for field in provider.fields} cleaned = {str(k): str(v).strip() for k, v in values.items()} if set(cleaned) != expected or any(not value for value in cleaned.values()): raise ProviderCredentialError("凭据字段不完整或包含未知字段") return cleaned def _notification_transition( old_status: str, old_notified_at: datetime | None, result: TestResult, now: datetime ) -> str | None: if result.status == "normal" and old_status in _BAD: return "recovered" if result.status not in _BAD: return None if old_status not in _BAD or old_status != result.status: return result.status if old_notified_at is None or now - old_notified_at >= _NOTIFY_COOLDOWN: return result.status return None def _apply_result( provider_id: str, result: TestResult, *, notify: Callable[[str, str, TestResult], None] | None = None, ) -> None: now = datetime.now(timezone.utc) event: str | None = None with session_scope() as session: row = session.execute( select(ProviderCredential).where(ProviderCredential.provider_id == provider_id) ).scalar_one_or_none() if row is None: row = ProviderCredential( provider_id=provider_id, credentials={}, credential_hint={}, revision=1 ) session.add(row) session.flush() old_status = row.test_status event = _notification_transition(old_status, row.last_notified_at, result, now) row.test_status = result.status row.test_detail = result.detail[:500] row.balance_amount = result.balance_amount row.balance_currency = result.balance_currency row.last_tested_at = now if event == "recovered": row.last_notified_at = None row.last_notified_status = None elif event: row.last_notified_at = now row.last_notified_status = result.status if event and notify: notify(provider_id, event, result) def replace_credentials( provider_id: str, values: dict[str, str], *, expected_revision: int, updated_by: UUID, request=None, notify: Callable[[str, str, TestResult], None] | None = None, ) -> dict: if not crypto_configured(): raise ProviderCredentialError( "未配置 ZCBOT_CREDENTIAL_MASTER_KEY,禁止保存数据库凭据" ) provider = get_provider(provider_id) cleaned = _validate_values(provider_id, values) kwargs = {"request": request} if request is not None else {} result = test_provider(provider_id, cleaned, **kwargs) if not result.accepted: raise ProviderCredentialError(f"候选凭据测试失败:{result.detail}") encrypted = { field.name: encrypt_secret( cleaned[field.name], aad=f"provider:{provider_id}:{field.name}" ) for field in provider.fields } hints = {field.name: _hint(cleaned[field.name]) for field in provider.fields} now = datetime.now(timezone.utc) try: with session_scope() as session: if expected_revision == 0: session.add(ProviderCredential( provider_id=provider_id, credentials=encrypted, credential_hint=hints, revision=1, test_status=result.status, test_detail=result.detail, balance_amount=result.balance_amount, balance_currency=result.balance_currency, last_tested_at=now, updated_by=updated_by, )) revision = 1 else: changed = session.execute( update(ProviderCredential) .where(ProviderCredential.provider_id == provider_id) .where(ProviderCredential.revision == expected_revision) .values( credentials=encrypted, credential_hint=hints, revision=expected_revision + 1, test_status=result.status, test_detail=result.detail, balance_amount=result.balance_amount, balance_currency=result.balance_currency, last_tested_at=now, updated_by=updated_by, updated_at=now, ) ) if int(changed.rowcount or 0) != 1: raise RevisionConflict("凭据已被其他管理员更新,请刷新后重试") revision = expected_revision + 1 except IntegrityError as exc: raise RevisionConflict("凭据已被其他管理员更新,请刷新后重试") from exc if result.status == "low_balance" and notify: notify(provider_id, result.status, result) with session_scope() as session: session.execute( update(ProviderCredential) .where(ProviderCredential.provider_id == provider_id) .where(ProviderCredential.revision == revision) .values(last_notified_at=now, last_notified_status=result.status) ) return {"provider_id": provider_id, "revision": revision, "test_status": result.status, "test_detail": result.detail} def test_current_credentials( provider_id: str, *, request=None, notify: Callable[[str, str, TestResult], None] | None = None, ) -> TestResult: resolved = resolve_credentials(provider_id) provider = get_provider(provider_id) if len(resolved.values) != len(provider.fields): raise ProviderCredentialError("Provider 凭据未完整配置") kwargs = {"request": request} if request is not None else {} result = test_provider(provider_id, resolved.values, **kwargs) _apply_result(provider_id, result, notify=notify) return result def delete_override(provider_id: str, *, expected_revision: int) -> None: get_provider(provider_id) with session_scope() as session: changed = session.execute( delete(ProviderCredential) .where(ProviderCredential.provider_id == provider_id) .where(ProviderCredential.revision == expected_revision) ) if int(changed.rowcount or 0) != 1: raise RevisionConflict("凭据已被更新或不存在,请刷新后重试") def record_business_failure( provider_id: str, *, status_code: int = 0, detail: str = "", notify: Callable[[str, str, TestResult], None] | None = None, ) -> str | None: result = classify_response(status_code, detail) if result.status not in {"auth_error", "exhausted"}: return None if notify is None: from .monitor import send_notification notify = send_notification _apply_result(provider_id, result, notify=notify) return result.status def due_provider_ids(now: datetime | None = None) -> list[str]: now = now or datetime.now(timezone.utc) with session_scope() as session: rows = { row.provider_id: row for row in session.execute(select(ProviderCredential)).scalars() } due = [] for provider in PROVIDERS: resolved = resolve_credentials(provider.provider_id) if len(resolved.values) != len(provider.fields): continue row = rows.get(provider.provider_id) if row is None or row.last_tested_at is None or ( now - row.last_tested_at >= timedelta(seconds=provider.check_interval_seconds) ): due.append(provider.provider_id) return due