264 lines
10 KiB
Python
264 lines
10 KiB
Python
"""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
|