275 lines
11 KiB
Python
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()
|