zcbot/core/provider_credentials/service.py

264 lines
10 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.

"""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