zcbot/tests/test_provider_credentials.py

275 lines
11 KiB
Python

import json
import os
import unittest
from contextlib import contextmanager
from datetime import datetime, timedelta, timezone
from importlib import util
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import MagicMock, Mock, patch
from uuid import uuid4
import httpx
from core.capabilities import ModelCapabilities
from core.llm import LLM
from core.provider_credentials.registry import BY_ID
from core.provider_credentials.runtime import resolve_credentials
from core.provider_credentials.service import (
ProviderCredentialError,
RevisionConflict,
_hint,
_notification_transition,
_row_payload,
replace_credentials,
)
from core.provider_credentials.testing import (
TestResult,
classify_response,
test_provider,
)
def response(status, body):
content = json.dumps(body).encode() if isinstance(body, dict) else str(body).encode()
return httpx.Response(status, content=content)
class ProviderRegistryTests(unittest.TestCase):
def test_credential_groups_are_explicit(self):
self.assertEqual(
[f.name for f in BY_ID["xfyun_iat"].fields],
["appid", "api_key", "api_secret"],
)
self.assertEqual(
[f.name for f in BY_ID["xfyun_lfasr"].fields],
["appid", "secret_key"],
)
self.assertNotIn("ZCBOT_DB_URL", [f.env for p in BY_ID.values() for f in p.fields])
def test_hint_only_reveals_last_four(self):
self.assertEqual(_hint("sk-abcdefgh"), "***efgh")
self.assertNotIn("abcdef", _hint("sk-abcdefgh"))
class ProviderClassificationTests(unittest.TestCase):
def test_http_402_and_429_are_not_conflated(self):
self.assertEqual(classify_response(402, "").status, "exhausted")
self.assertEqual(classify_response(429, "rate limit").status, "unreachable")
def test_deepseek_30_yuan_boundary(self):
def request_at(amount):
return lambda *a, **k: response(200, {
"is_available": True,
"balance_infos": [{"currency": "CNY", "total_balance": str(amount)}],
})
creds = {"api_key": "candidate-secret"}
self.assertEqual(
test_provider("deepseek", creds, request=request_at("29.99")).status,
"low_balance",
)
self.assertEqual(
test_provider("deepseek", creds, request=request_at("30.00")).status,
"normal",
)
def test_candidate_secret_is_not_returned_in_detail(self):
secret = "candidate-super-secret"
result = test_provider(
"deepseek", {"api_key": secret},
request=lambda *a, **k: response(401, {"message": secret}),
)
self.assertEqual(result.status, "auth_error")
self.assertNotIn(secret, result.detail)
class ProviderResolutionTests(unittest.TestCase):
def test_database_precedes_environment(self):
row = SimpleNamespace(credentials={"api_key": "cipher"})
class Result:
def scalar_one_or_none(self):
return row
class Session:
def execute(self, _statement):
return Result()
@contextmanager
def scope():
yield Session()
with (
patch.dict(os.environ, {"DEEPSEEK_API_KEY": "env-key"}),
patch("core.storage.session_scope", scope),
patch("core.provider_credentials.runtime.decrypt_secret", return_value="db-key") as decrypt,
):
resolved = resolve_credentials("deepseek")
self.assertEqual((resolved.source, resolved.values["api_key"]), ("database", "db-key"))
decrypt.assert_called_once_with("cipher", aad="provider:deepseek:api_key")
def test_environment_fallback_without_database(self):
with (
patch.dict(os.environ, {"DEEPSEEK_API_KEY": "env-key"}),
patch("core.storage.session_scope", side_effect=RuntimeError("no db")),
):
resolved = resolve_credentials("deepseek")
self.assertEqual((resolved.source, resolved.values), ("env", {"api_key": "env-key"}))
def test_each_xfyun_group_resolves_its_declared_fields(self):
env = {
"XFYUN_APPID": "app", "XFYUN_API_KEY": "iat-key",
"XFYUN_API_SECRET": "iat-secret", "XFYUN_LFASR_SECRET_KEY": "lf-key",
}
with patch.dict(os.environ, env), patch("core.storage.session_scope", side_effect=RuntimeError("no db")):
self.assertEqual(set(resolve_credentials("xfyun_iat").values), {"appid", "api_key", "api_secret"})
self.assertEqual(set(resolve_credentials("xfyun_lfasr").values), {"appid", "secret_key"})
class LLMCredentialLifecycleTests(unittest.TestCase):
def _caps(self):
return ModelCapabilities(
family="deepseek_v4", model_id="deepseek/model",
api_key_env="DEEPSEEK_API_KEY", thinking_transport="none",
)
def test_created_llm_uses_new_key_on_next_request(self):
state = {"key": "old-key"}
llm = LLM(self._caps(), credential_resolver=lambda _env: state["key"])
first = llm._build_kwargs([], None, None, None)
state["key"] = "new-key"
second = llm._build_kwargs([], None, None, None)
self.assertEqual(first["api_key"], "old-key")
self.assertEqual(second["api_key"], "new-key")
def test_built_stream_request_keeps_key_snapshot(self):
state = {"key": "old-key"}
calls = []
llm = LLM(self._caps(), credential_resolver=lambda _env: state["key"])
def completion(**kwargs):
calls.append(kwargs["api_key"])
state["key"] = "new-key"
return iter([{"chunk": 1}, {"chunk": 2}])
with patch("core.llm.litellm.completion", side_effect=completion):
self.assertEqual(len(list(llm.chat_stream([]))), 2)
self.assertEqual(calls, ["old-key"])
class ProviderMutationTests(unittest.TestCase):
def test_missing_master_key_refuses_database_save(self):
with (
patch("core.provider_credentials.service.crypto_configured", return_value=False),
self.assertRaisesRegex(ProviderCredentialError, "MASTER_KEY"),
):
replace_credentials(
"deepseek", {"api_key": "candidate"}, expected_revision=0,
updated_by=uuid4(),
)
def test_failed_candidate_never_opens_database_transaction(self):
with (
patch("core.provider_credentials.service.crypto_configured", return_value=True),
patch("core.provider_credentials.service.test_provider",
return_value=TestResult("auth_error", "认证失败")),
patch("core.provider_credentials.service.session_scope") as scope,
self.assertRaises(ProviderCredentialError),
):
replace_credentials(
"deepseek", {"api_key": "bad-key"}, expected_revision=0,
updated_by=uuid4(),
)
scope.assert_not_called()
def test_revision_conflict_does_not_overwrite(self):
class Session:
def execute(self, _statement):
return SimpleNamespace(rowcount=0)
@contextmanager
def scope():
yield Session()
with (
patch("core.provider_credentials.service.crypto_configured", return_value=True),
patch("core.provider_credentials.service.test_provider",
return_value=TestResult("normal", "正常")),
patch("core.provider_credentials.service.encrypt_secret", return_value="cipher"),
patch("core.provider_credentials.service.session_scope", scope),
self.assertRaises(RevisionConflict),
):
replace_credentials(
"deepseek", {"api_key": "new-key"}, expected_revision=7,
updated_by=uuid4(),
)
def test_alert_cooldown_and_recovery(self):
now = datetime.now(timezone.utc)
low = TestResult("low_balance", "余额低")
self.assertEqual(_notification_transition("normal", None, low, now), "low_balance")
self.assertIsNone(_notification_transition("low_balance", now, low, now))
self.assertEqual(
_notification_transition("low_balance", now - timedelta(days=1), low, now),
"low_balance",
)
self.assertEqual(
_notification_transition("auth_error", now, TestResult("normal", "正常"), now),
"recovered",
)
def test_admin_payload_contains_no_ciphertext(self):
provider = BY_ID["deepseek"]
row = SimpleNamespace(
credentials={"api_key": "ciphertext-secret"},
credential_hint={"api_key": "***1234"}, revision=2,
test_status="normal", test_detail="正常", last_tested_at=None,
balance_amount=None, balance_currency=None,
)
with patch("core.provider_credentials.service.resolve_credentials") as resolve:
resolve.return_value = SimpleNamespace(source="database", values={"api_key": "plain-secret"})
payload = _row_payload(provider, row)
encoded = json.dumps(payload, ensure_ascii=False)
self.assertNotIn("ciphertext-secret", encoded)
self.assertNotIn("plain-secret", encoded)
self.assertIn("***1234", encoded)
class ProviderLeaderTests(unittest.TestCase):
def test_unclaimed_round_does_not_test(self):
connection = Mock()
connection.exec_driver_sql.return_value.scalar.return_value = False
engine = MagicMock()
engine.connect.return_value.__enter__.return_value = connection
with (
patch("core.provider_credentials.monitor.get_engine", return_value=engine),
patch("core.provider_credentials.monitor.due_provider_ids") as due,
):
from core.provider_credentials.monitor import run_due_checks
self.assertEqual(run_due_checks(), 0)
due.assert_not_called()
class ProviderMigrationTests(unittest.TestCase):
def test_0039_creates_single_encrypted_credentials_table(self):
path = Path(__file__).resolve().parents[1] / "db" / "migrations" / "versions" / "20260902_1000_0039_provider_credentials.py"
spec = util.spec_from_file_location("migration_0039_test", path)
module = util.module_from_spec(spec)
assert spec.loader is not None
spec.loader.exec_module(module)
captured = {}
def create_table(name, *items):
captured["name"] = name
captured["columns"] = {item.name for item in items if hasattr(item, "name")}
with patch.object(module.op, "create_table", side_effect=create_table):
module.upgrade()
self.assertEqual(captured["name"], "provider_credentials")
self.assertTrue({"provider_id", "credentials", "credential_hint", "revision"} <= captured["columns"])
self.assertNotIn("api_key", captured["columns"])
if __name__ == "__main__":
unittest.main()