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()