refactor(provider-keys): 拆分 keys 端点逻辑并补充配额刷新测试

将 admin keys 接口的创建、更新、查询、删除、导出与配额刷新逻辑迁移到 provider_keys 服务层

新增 auth_type 归一化、重复校验、响应构建与 quota_refresh(codex/kiro/antigravity)模块,并补充对应单元测试
This commit is contained in:
AAEE86
2026-02-28 14:09:36 +08:00
parent 4b02078b60
commit 08b89b7ef8
18 changed files with 3465 additions and 1368 deletions

View File

@@ -0,0 +1,207 @@
from __future__ import annotations
import sys
import types
from datetime import datetime, timezone
from types import SimpleNamespace
from typing import Any, cast
import pytest
from src.core.exceptions import InvalidRequestException
from src.models.endpoint_models import EndpointAPIKeyUpdate
async def _noop_invalidate_models_list_cache() -> None:
return None
_fake_models_service_module = types.ModuleType("src.api.base.models_service")
setattr(
_fake_models_service_module, "invalidate_models_list_cache", _noop_invalidate_models_list_cache
)
sys.modules.setdefault("src.api.base.models_service", _fake_models_service_module)
from src.services.provider_keys import key_command_service as command_module
from src.services.provider_keys import key_side_effects as side_effects_module
class _NoQueryDB:
def query(self, *args: Any, **kwargs: Any) -> Any: # pragma: no cover - 防御断言
_ = args, kwargs
raise AssertionError("unexpected query call")
class _FakeQuery:
def __init__(self, key: Any) -> None:
self._key = key
def filter(self, *args: Any, **kwargs: Any) -> _FakeQuery:
_ = args, kwargs
return self
def first(self) -> Any:
return self._key
class _FakeClearOAuthDB:
def __init__(self, key: Any) -> None:
self._key = key
self.commit_count = 0
def query(self, model: Any) -> _FakeQuery:
_ = model
return _FakeQuery(self._key)
def commit(self) -> None:
self.commit_count += 1
def _build_key(**overrides: Any) -> SimpleNamespace:
base: dict[str, Any] = {
"auto_fetch_models": False,
"allowed_models": None,
"model_include_patterns": None,
"model_exclude_patterns": None,
"provider_id": "provider-1",
"auth_type": "api_key",
}
base.update(overrides)
return SimpleNamespace(**base)
def test_prepare_update_payload_auth_type_null_ignored() -> None:
key = _build_key(auth_type="api_key")
key_data = EndpointAPIKeyUpdate.model_validate({"auth_type": None})
prepared = command_module._prepare_update_key_payload(
db=cast(Any, _NoQueryDB()),
key=cast(Any, key),
key_id="key-1",
key_data=key_data,
)
assert "auth_type" not in prepared.update_data
def test_prepare_update_payload_rejects_empty_api_key() -> None:
key = _build_key(auth_type="api_key")
key_data = EndpointAPIKeyUpdate.model_validate({"api_key": " "})
with pytest.raises(InvalidRequestException, match="api_key 不能为空"):
command_module._prepare_update_key_payload(
db=cast(Any, _NoQueryDB()),
key=cast(Any, key),
key_id="key-1",
key_data=key_data,
)
def test_prepare_update_payload_encrypts_empty_auth_config_dict(
monkeypatch: pytest.MonkeyPatch,
) -> None:
key = _build_key(auth_type="vertex_ai")
key_data = EndpointAPIKeyUpdate.model_validate({"auth_config": {}})
monkeypatch.setattr(command_module.crypto_service, "encrypt", lambda raw: f"ENC:{raw}")
prepared = command_module._prepare_update_key_payload(
db=cast(Any, _NoQueryDB()),
key=cast(Any, key),
key_id="key-1",
key_data=key_data,
)
assert prepared.update_data["auth_config"] == "ENC:{}"
def test_prepare_update_payload_vertex_to_oauth_clears_auth_config(
monkeypatch: pytest.MonkeyPatch,
) -> None:
key = _build_key(auth_type="vertex_ai")
key_data = EndpointAPIKeyUpdate.model_validate({"auth_type": "oauth"})
monkeypatch.setattr(command_module.crypto_service, "encrypt", lambda raw: f"ENC:{raw}")
prepared = command_module._prepare_update_key_payload(
db=cast(Any, _NoQueryDB()),
key=cast(Any, key),
key_id="key-1",
key_data=key_data,
)
assert prepared.update_data["auth_config"] is None
assert prepared.update_data["api_key"] == "ENC:__placeholder__"
def test_clear_oauth_invalid_response_invalidates_caches(
monkeypatch: pytest.MonkeyPatch,
) -> None:
cache_calls: list[tuple[str, str | None]] = []
async def _fake_invalidate_key_cache(key_id: str) -> None:
cache_calls.append(("key", key_id))
async def _fake_invalidate_models_cache() -> None:
cache_calls.append(("models", None))
fake_provider_cache_module = types.ModuleType("src.services.cache.provider_cache")
class _FakeProviderCacheService:
@staticmethod
async def invalidate_provider_api_key_cache(key_id: str) -> None:
await _fake_invalidate_key_cache(key_id)
setattr(fake_provider_cache_module, "ProviderCacheService", _FakeProviderCacheService)
monkeypatch.setitem(
sys.modules, "src.services.cache.provider_cache", fake_provider_cache_module
)
fake_models_service_module = types.ModuleType("src.api.base.models_service")
setattr(
fake_models_service_module, "invalidate_models_list_cache", _fake_invalidate_models_cache
)
monkeypatch.setitem(sys.modules, "src.api.base.models_service", fake_models_service_module)
key = SimpleNamespace(
oauth_invalid_at=datetime.now(timezone.utc),
oauth_invalid_reason="forbidden",
is_active=False,
)
db = _FakeClearOAuthDB(key=key)
result = command_module.clear_oauth_invalid_response(cast(Any, db), key_id="key-1")
assert result["message"] == "已清除 OAuth 失效标记Key 已自动启用"
assert key.oauth_invalid_at is None
assert key.oauth_invalid_reason is None
assert key.is_active is True
assert db.commit_count == 1
assert cache_calls == [("key", "key-1"), ("models", None)]
@pytest.mark.asyncio
async def test_run_delete_key_side_effects_not_skip_disassociate(
monkeypatch: pytest.MonkeyPatch,
) -> None:
captured: dict[str, Any] = {}
async def _fake_on_key_allowed_models_changed(**kwargs: Any) -> None:
captured.update(kwargs)
from src.services.model import global_model as global_model_module
monkeypatch.setattr(
global_model_module,
"on_key_allowed_models_changed",
_fake_on_key_allowed_models_changed,
)
await side_effects_module.run_delete_key_side_effects(
db=cast(Any, object()),
provider_id="provider-1",
deleted_key_allowed_models=None,
)
assert captured["provider_id"] == "provider-1"
assert "skip_disassociate" not in captured

View File

@@ -0,0 +1,155 @@
from __future__ import annotations
from types import SimpleNamespace
from typing import Any, cast
import pytest
from src.core.exceptions import NotFoundException
from src.services.provider_keys import key_query_service as query_service_module
from src.services.provider_keys.key_query_service import (
get_keys_grouped_by_format,
list_provider_keys_responses,
)
class _FakeQuery:
def __init__(
self,
*,
first_result: Any = None,
all_result: list[Any] | None = None,
) -> None:
self._first_result = first_result
self._all_result = all_result or []
def join(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def filter(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def order_by(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def offset(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def limit(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def first(self) -> Any:
return self._first_result
def all(self) -> list[Any]:
return self._all_result
class _FakeGroupedDB:
def __init__(
self,
*,
key_provider_rows: list[tuple[Any, Any]],
endpoint_rows: list[tuple[str, str, str]],
) -> None:
self._key_provider_rows = key_provider_rows
self._endpoint_rows = endpoint_rows
def query(self, *models: Any) -> _FakeQuery:
if len(models) == 2:
return _FakeQuery(all_result=self._key_provider_rows)
if len(models) == 3:
return _FakeQuery(all_result=self._endpoint_rows)
raise AssertionError(f"unexpected query models: {models}")
class _FakeListDB:
def __init__(self, *, provider: Any, keys: list[SimpleNamespace]) -> None:
self._provider = provider
self._keys = keys
def query(self, model: Any) -> _FakeQuery:
model_name = getattr(model, "__name__", "")
if model_name == "Provider":
return _FakeQuery(first_result=self._provider)
if model_name == "ProviderAPIKey":
return _FakeQuery(all_result=self._keys)
raise AssertionError(f"unexpected query model: {model}")
def test_get_keys_grouped_by_format_builds_expected_shape(
monkeypatch: pytest.MonkeyPatch,
) -> None:
provider = SimpleNamespace(id="p1", is_active=True, name="Provider-1")
key = SimpleNamespace(
id="k1",
name="Key-1",
api_formats=["openai:chat", "openai:cli"],
auth_type="api_key",
api_key="enc-key",
internal_priority=3,
global_priority_by_format={"openai:chat": 8},
rate_multipliers={"openai:chat": 1.1},
is_active=True,
capabilities={"cache_1h": True, "ctx_1m": False},
success_count=8,
request_count=10,
total_response_time_ms=800,
health_by_format={"openai:chat": {"health_score": 0.7}},
circuit_breaker_by_format={"openai:chat": {"open": True}},
)
db = _FakeGroupedDB(
key_provider_rows=[(key, provider)],
endpoint_rows=[
("p1", "openai:chat", "https://chat.example"),
("p1", "openai:cli", "https://cli.example"),
],
)
monkeypatch.setattr(
query_service_module.crypto_service, "decrypt", lambda _v: "sk-1234567890abcd"
)
monkeypatch.setattr(
query_service_module,
"get_capability",
lambda name: SimpleNamespace(short_name="缓存1h") if name == "cache_1h" else None,
)
result = get_keys_grouped_by_format(cast(Any, db))
assert set(result.keys()) == {"openai:chat", "openai:cli"}
chat_item = result["openai:chat"][0]
assert chat_item["id"] == "k1"
assert chat_item["provider_name"] == "Provider-1"
assert chat_item["endpoint_base_url"] == "https://chat.example"
assert chat_item["format_priority"] == 8
assert chat_item["circuit_breaker_open"] is True
assert chat_item["health_score"] == 0.7
assert chat_item["capabilities"] == ["缓存1h"]
assert chat_item["api_key_masked"].startswith("sk-12345")
assert chat_item["api_key_masked"].endswith("abcd")
cli_item = result["openai:cli"][0]
assert cli_item["endpoint_base_url"] == "https://cli.example"
assert cli_item["format_priority"] is None
assert cli_item["circuit_breaker_open"] is False
assert cli_item["health_score"] == 1.0
def test_list_provider_keys_responses_provider_not_found_raises() -> None:
db = _FakeListDB(provider=None, keys=[])
with pytest.raises(NotFoundException, match="Provider p1 不存在"):
list_provider_keys_responses(cast(Any, db), provider_id="p1", skip=0, limit=10)
def test_list_provider_keys_responses_uses_response_builder(
monkeypatch: pytest.MonkeyPatch,
) -> None:
provider = SimpleNamespace(id="p1")
keys = [SimpleNamespace(id="k1"), SimpleNamespace(id="k2")]
db = _FakeListDB(provider=provider, keys=keys)
monkeypatch.setattr(query_service_module, "build_key_response", lambda key: {"id": key.id})
result = list_provider_keys_responses(cast(Any, db), provider_id="p1", skip=0, limit=10)
assert result == [{"id": "k1"}, {"id": "k2"}]

View File

@@ -0,0 +1,699 @@
from __future__ import annotations
import json
import sys
import types
from types import SimpleNamespace
from typing import Any, cast
import pytest
from src.services.provider_keys.codex_usage_parser import (
CodexUsageParseError,
parse_codex_wham_usage_response,
)
from src.services.provider_keys.quota_refresh.antigravity_refresher import (
refresh_antigravity_key_quota,
)
from src.services.provider_keys.quota_refresh.codex_refresher import refresh_codex_key_quota
from src.services.provider_keys.quota_refresh.kiro_refresher import refresh_kiro_key_quota
class _FakeDB:
def __init__(self) -> None:
self.commit_count = 0
def commit(self) -> None:
self.commit_count += 1
class _FakeResponse:
def __init__(
self,
status_code: int,
payload: Any = None,
json_exc: Exception | None = None,
) -> None:
self.status_code = status_code
self._payload = payload
self._json_exc = json_exc
def json(self) -> Any:
if self._json_exc:
raise self._json_exc
return self._payload
class _FakeAsyncClient:
def __init__(self, response: _FakeResponse, **kwargs: Any) -> None:
self._response = response
self.kwargs = kwargs
self.last_url: str | None = None
self.last_headers: dict[str, str] | None = None
async def __aenter__(self) -> "_FakeAsyncClient":
return self
async def __aexit__(
self,
exc_type: type[BaseException] | None,
exc: BaseException | None,
tb: Any,
) -> bool:
_ = exc_type, exc, tb
return False
async def get(self, url: str, headers: dict[str, str]) -> _FakeResponse:
self.last_url = url
self.last_headers = headers
return self._response
def _install_module(monkeypatch: pytest.MonkeyPatch, name: str, attrs: dict[str, Any]) -> None:
module = types.ModuleType(name)
for key, value in attrs.items():
setattr(module, key, value)
monkeypatch.setitem(sys.modules, name, module)
@pytest.mark.asyncio
async def test_codex_refresher_endpoint_missing_returns_error() -> None:
key = SimpleNamespace(id="k1", name="K1")
provider = SimpleNamespace(proxy=None)
result = await refresh_codex_key_quota(
db=cast(Any, _FakeDB()),
provider=cast(Any, provider),
key=cast(Any, key),
endpoint=None,
codex_wham_usage_url="https://example.test",
metadata_updates={},
state_updates={},
)
assert result["status"] == "error"
assert "openai:cli" in result["message"]
@pytest.mark.asyncio
async def test_codex_refresher_http_non_200_returns_error(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from src.services.provider_keys.quota_refresh import codex_refresher as module
key = SimpleNamespace(
id="k1", name="K1", api_key="enc", auth_type="api_key", auth_config=None, proxy=None
)
provider = SimpleNamespace(proxy=None)
endpoint = SimpleNamespace()
async def _fake_auth_info(_endpoint: Any, _key: Any) -> Any:
return None
_install_module(
monkeypatch,
"src.services.proxy_node.resolver",
{
"resolve_effective_proxy": lambda provider_proxy, key_proxy: None,
"build_proxy_client_kwargs": lambda proxy, timeout: {"timeout": timeout},
},
)
monkeypatch.setattr(module, "get_provider_auth", _fake_auth_info)
monkeypatch.setattr(module.crypto_service, "decrypt", lambda _v: "sk-test")
response = _FakeResponse(status_code=503, payload={"x": 1})
monkeypatch.setattr(
module.httpx, "AsyncClient", lambda **kwargs: _FakeAsyncClient(response, **kwargs)
)
result = await refresh_codex_key_quota(
db=cast(Any, _FakeDB()),
provider=cast(Any, provider),
key=cast(Any, key),
endpoint=cast(Any, endpoint),
codex_wham_usage_url="https://example.test",
metadata_updates={},
state_updates={},
)
assert result["status"] == "error"
assert result["status_code"] == 503
@pytest.mark.asyncio
async def test_codex_refresher_success_updates_metadata(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from src.services.provider_keys.quota_refresh import codex_refresher as module
key = SimpleNamespace(
id="k1",
name="K1",
api_key="enc",
auth_type="api_key",
auth_config=None,
proxy=None,
oauth_invalid_at="old",
oauth_invalid_reason="old",
)
provider = SimpleNamespace(proxy=None)
endpoint = SimpleNamespace()
metadata_updates: dict[str, dict[str, Any]] = {}
state_updates: dict[str, dict[str, Any]] = {}
async def _fake_auth_info(_endpoint: Any, _key: Any) -> Any:
return None
_install_module(
monkeypatch,
"src.services.proxy_node.resolver",
{
"resolve_effective_proxy": lambda provider_proxy, key_proxy: None,
"build_proxy_client_kwargs": lambda proxy, timeout: {"timeout": timeout},
},
)
monkeypatch.setattr(module, "get_provider_auth", _fake_auth_info)
monkeypatch.setattr(module.crypto_service, "decrypt", lambda _v: "sk-test")
monkeypatch.setattr(
module, "parse_codex_wham_usage_response", lambda _data: {"used_percent": 10.0}
)
response = _FakeResponse(status_code=200, payload={"ok": True})
monkeypatch.setattr(
module.httpx, "AsyncClient", lambda **kwargs: _FakeAsyncClient(response, **kwargs)
)
result = await refresh_codex_key_quota(
db=cast(Any, _FakeDB()),
provider=cast(Any, provider),
key=cast(Any, key),
endpoint=cast(Any, endpoint),
codex_wham_usage_url="https://example.test",
metadata_updates=metadata_updates,
state_updates=state_updates,
)
assert result["status"] == "success"
assert metadata_updates == {"k1": {"codex": {"used_percent": 10.0}}}
assert state_updates == {"k1": {"oauth_invalid_at": None, "oauth_invalid_reason": None}}
@pytest.mark.asyncio
async def test_codex_refresher_parse_error_is_diagnostic(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from src.services.provider_keys.quota_refresh import codex_refresher as module
key = SimpleNamespace(
id="k1", name="K1", api_key="enc", auth_type="api_key", auth_config=None, proxy=None
)
provider = SimpleNamespace(proxy=None)
endpoint = SimpleNamespace()
async def _fake_auth_info(_endpoint: Any, _key: Any) -> Any:
return None
_install_module(
monkeypatch,
"src.services.proxy_node.resolver",
{
"resolve_effective_proxy": lambda provider_proxy, key_proxy: None,
"build_proxy_client_kwargs": lambda proxy, timeout: {"timeout": timeout},
},
)
monkeypatch.setattr(module, "get_provider_auth", _fake_auth_info)
monkeypatch.setattr(module.crypto_service, "decrypt", lambda _v: "sk-test")
monkeypatch.setattr(
module,
"parse_codex_wham_usage_response",
lambda _data: (_ for _ in ()).throw(
CodexUsageParseError("rate_limit.primary_window 类型错误")
),
)
response = _FakeResponse(status_code=200, payload={"ok": True})
monkeypatch.setattr(
module.httpx, "AsyncClient", lambda **kwargs: _FakeAsyncClient(response, **kwargs)
)
result = await refresh_codex_key_quota(
db=cast(Any, _FakeDB()),
provider=cast(Any, provider),
key=cast(Any, key),
endpoint=cast(Any, endpoint),
codex_wham_usage_url="https://example.test",
metadata_updates={},
state_updates={},
)
assert result["status"] == "error"
assert "响应结构异常" in result["message"]
assert "rate_limit.primary_window 类型错误" in result["message"]
@pytest.mark.asyncio
async def test_codex_refresher_oauth_missing_plan_type_adds_account_header(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from src.services.provider_keys.quota_refresh import codex_refresher as module
key = SimpleNamespace(
id="k1",
name="K1",
api_key="enc",
auth_type="oauth",
auth_config="enc-config",
proxy=None,
)
provider = SimpleNamespace(proxy=None)
endpoint = SimpleNamespace()
response = _FakeResponse(status_code=200, payload={"ok": True})
client_ref: dict[str, _FakeAsyncClient] = {}
async def _fake_auth_info(_endpoint: Any, _key: Any) -> Any:
return SimpleNamespace(auth_header="Authorization", auth_value="Bearer oauth-token")
_install_module(
monkeypatch,
"src.services.proxy_node.resolver",
{
"resolve_effective_proxy": lambda provider_proxy, key_proxy: None,
"build_proxy_client_kwargs": lambda proxy, timeout: {"timeout": timeout},
},
)
monkeypatch.setattr(module, "get_provider_auth", _fake_auth_info)
monkeypatch.setattr(
module, "parse_codex_wham_usage_response", lambda _data: {"used_percent": 1.0}
)
monkeypatch.setattr(
module.crypto_service, "decrypt", lambda _v: json.dumps({"account_id": "acc-1"})
)
def _client_factory(**kwargs: Any) -> _FakeAsyncClient:
client = _FakeAsyncClient(response, **kwargs)
client_ref["client"] = client
return client
monkeypatch.setattr(module.httpx, "AsyncClient", _client_factory)
result = await refresh_codex_key_quota(
db=cast(Any, _FakeDB()),
provider=cast(Any, provider),
key=cast(Any, key),
endpoint=cast(Any, endpoint),
codex_wham_usage_url="https://example.test",
metadata_updates={},
state_updates={},
)
assert result["status"] == "success"
assert client_ref["client"].last_headers is not None
assert client_ref["client"].last_headers.get("chatgpt-account-id") == "acc-1"
@pytest.mark.asyncio
async def test_codex_refresher_oauth_uppercase_free_does_not_add_account_header(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from src.services.provider_keys.quota_refresh import codex_refresher as module
key = SimpleNamespace(
id="k1",
name="K1",
api_key="enc",
auth_type="oauth",
auth_config="enc-config",
proxy=None,
)
provider = SimpleNamespace(proxy=None)
endpoint = SimpleNamespace()
response = _FakeResponse(status_code=200, payload={"ok": True})
client_ref: dict[str, _FakeAsyncClient] = {}
async def _fake_auth_info(_endpoint: Any, _key: Any) -> Any:
return SimpleNamespace(auth_header="Authorization", auth_value="Bearer oauth-token")
_install_module(
monkeypatch,
"src.services.proxy_node.resolver",
{
"resolve_effective_proxy": lambda provider_proxy, key_proxy: None,
"build_proxy_client_kwargs": lambda proxy, timeout: {"timeout": timeout},
},
)
monkeypatch.setattr(module, "get_provider_auth", _fake_auth_info)
monkeypatch.setattr(
module, "parse_codex_wham_usage_response", lambda _data: {"used_percent": 1.0}
)
monkeypatch.setattr(
module.crypto_service,
"decrypt",
lambda _v: json.dumps({"account_id": "acc-1", "plan_type": "FREE"}),
)
def _client_factory(**kwargs: Any) -> _FakeAsyncClient:
client = _FakeAsyncClient(response, **kwargs)
client_ref["client"] = client
return client
monkeypatch.setattr(module.httpx, "AsyncClient", _client_factory)
result = await refresh_codex_key_quota(
db=cast(Any, _FakeDB()),
provider=cast(Any, provider),
key=cast(Any, key),
endpoint=cast(Any, endpoint),
codex_wham_usage_url="https://example.test",
metadata_updates={},
state_updates={},
)
assert result["status"] == "success"
assert client_ref["client"].last_headers is not None
assert "chatgpt-account-id" not in client_ref["client"].last_headers
def test_parse_codex_usage_plan_type_case_insensitive_free_window_semantics() -> None:
parsed = parse_codex_wham_usage_response(
{
"plan_type": "FREE",
"rate_limit": {
"primary_window": {
"used_percent": "12.5",
"reset_after_seconds": "120",
"reset_at": "1700000000",
"limit_window_seconds": "604800",
}
},
}
)
assert parsed is not None
assert parsed["plan_type"] == "free"
assert parsed["primary_used_percent"] == 12.5
assert parsed["primary_window_minutes"] == 10080
assert "secondary_used_percent" not in parsed
def test_parse_codex_usage_missing_plan_type_infers_paid_windows() -> None:
parsed = parse_codex_wham_usage_response(
{
"rate_limit": {
"primary_window": {
"used_percent": 25,
"reset_after_seconds": 600,
"reset_at": 1700000000,
"limit_window_seconds": 18000,
},
"secondary_window": {
"used_percent": 80,
"reset_after_seconds": 3600,
"reset_at": 1700003600,
"limit_window_seconds": 604800,
},
}
}
)
assert parsed is not None
assert parsed["primary_used_percent"] == 80.0
assert parsed["secondary_used_percent"] == 25.0
assert parsed["primary_window_minutes"] == 10080
assert parsed["secondary_window_minutes"] == 300
def test_parse_codex_usage_invalid_type_raises_diagnostic_error() -> None:
with pytest.raises(CodexUsageParseError, match="rate_limit.primary_window 类型错误"):
parse_codex_wham_usage_response({"rate_limit": {"primary_window": []}})
@pytest.mark.asyncio
async def test_antigravity_refresher_forbidden_collects_updates_without_commit(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from src.services.provider_keys.quota_refresh import antigravity_refresher as module
class _Forbidden(Exception):
def __init__(self, reason: str) -> None:
super().__init__(reason)
self.reason = reason
self.message = reason
async def _fake_auth_info(_endpoint: Any, _key: Any) -> Any:
return SimpleNamespace(auth_value="Bearer tk", decrypted_auth_config={"pid": "p1"})
async def _fetch_models_for_key(_ctx: Any, timeout_seconds: float) -> Any:
_ = timeout_seconds
raise _Forbidden("forbidden-by-test")
class _UpstreamModelsFetchContext: # noqa: D101
def __init__(self, **kwargs: Any) -> None:
self.kwargs = kwargs
_install_module(
monkeypatch,
"src.services.model.upstream_fetcher",
{
"UpstreamModelsFetchContext": _UpstreamModelsFetchContext,
"fetch_models_for_key": _fetch_models_for_key,
},
)
_install_module(
monkeypatch,
"src.services.provider.adapters.antigravity.client",
{"AntigravityAccountForbiddenException": _Forbidden},
)
_install_module(
monkeypatch,
"src.services.proxy_node.resolver",
{"resolve_effective_proxy": lambda provider_proxy, key_proxy: None},
)
monkeypatch.setattr(module, "get_provider_auth", _fake_auth_info)
db = _FakeDB()
provider = SimpleNamespace(proxy=None)
key = SimpleNamespace(
id="k1",
name="K1",
proxy=None,
is_active=True,
oauth_invalid_at=None,
oauth_invalid_reason=None,
upstream_metadata={},
)
endpoint = SimpleNamespace()
metadata_updates: dict[str, dict[str, Any]] = {}
state_updates: dict[str, dict[str, Any]] = {}
result = await refresh_antigravity_key_quota(
db=cast(Any, db),
provider=cast(Any, provider),
key=cast(Any, key),
endpoint=cast(Any, endpoint),
codex_wham_usage_url="https://example.test",
metadata_updates=metadata_updates,
state_updates=state_updates,
)
assert result["status"] == "forbidden"
assert result["auto_disabled"] is True
assert key.is_active is True
assert key.oauth_invalid_reason is None
assert state_updates["k1"]["is_active"] is False
assert state_updates["k1"]["oauth_invalid_reason"].startswith("账户访问被禁止")
assert metadata_updates["k1"]["antigravity"]["is_forbidden"] is True
assert db.commit_count == 0
@pytest.mark.asyncio
async def test_antigravity_refresher_success_resets_forbidden_flag(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from src.services.provider_keys.quota_refresh import antigravity_refresher as module
async def _fake_auth_info(_endpoint: Any, _key: Any) -> Any:
return SimpleNamespace(auth_value="Bearer tk", decrypted_auth_config={})
async def _fetch_models_for_key(_ctx: Any, timeout_seconds: float) -> Any:
_ = timeout_seconds
return (
[],
[],
True,
{"antigravity": {"is_forbidden": True, "forbidden_reason": "x", "forbidden_at": 1}},
)
class _UpstreamModelsFetchContext: # noqa: D101
def __init__(self, **kwargs: Any) -> None:
self.kwargs = kwargs
_install_module(
monkeypatch,
"src.services.model.upstream_fetcher",
{
"UpstreamModelsFetchContext": _UpstreamModelsFetchContext,
"fetch_models_for_key": _fetch_models_for_key,
},
)
_install_module(
monkeypatch,
"src.services.provider.adapters.antigravity.client",
{"AntigravityAccountForbiddenException": RuntimeError},
)
_install_module(
monkeypatch,
"src.services.proxy_node.resolver",
{"resolve_effective_proxy": lambda provider_proxy, key_proxy: None},
)
monkeypatch.setattr(module, "get_provider_auth", _fake_auth_info)
metadata_updates: dict[str, dict[str, Any]] = {}
state_updates: dict[str, dict[str, Any]] = {}
result = await refresh_antigravity_key_quota(
db=cast(Any, _FakeDB()),
provider=cast(Any, SimpleNamespace(proxy=None)),
key=cast(
Any,
SimpleNamespace(
id="k1",
name="K1",
proxy=None,
oauth_invalid_at="old",
oauth_invalid_reason="old",
),
),
endpoint=cast(Any, SimpleNamespace()),
codex_wham_usage_url="https://example.test",
metadata_updates=metadata_updates,
state_updates=state_updates,
)
assert result["status"] == "success"
assert metadata_updates["k1"]["antigravity"]["is_forbidden"] is False
assert metadata_updates["k1"]["antigravity"]["forbidden_reason"] is None
assert metadata_updates["k1"]["antigravity"]["forbidden_at"] is None
assert state_updates["k1"]["oauth_invalid_at"] is None
assert state_updates["k1"]["oauth_invalid_reason"] is None
@pytest.mark.asyncio
async def test_kiro_refresher_runtime_401_marks_key_invalid(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from src.services.provider_keys.quota_refresh import kiro_refresher as module
class _Banned(Exception):
pass
async def _fetch_limits(auth_config: dict[str, Any], proxy_config: object) -> Any:
_ = auth_config, proxy_config
raise RuntimeError("401 token expired")
def _parse_usage(_usage: Any) -> dict[str, Any]:
return {"quota": 1}
_install_module(
monkeypatch,
"src.services.provider.adapters.kiro.usage",
{
"KiroAccountBannedException": _Banned,
"fetch_kiro_usage_limits": _fetch_limits,
"parse_kiro_usage_response": _parse_usage,
},
)
_install_module(
monkeypatch,
"src.services.proxy_node.resolver",
{"resolve_effective_proxy": lambda provider_proxy, key_proxy: None},
)
monkeypatch.setattr(module.crypto_service, "decrypt", lambda _v: "{}")
db = _FakeDB()
key = SimpleNamespace(
id="k1",
name="K1",
auth_config="enc",
proxy=None,
is_active=True,
oauth_invalid_at=None,
oauth_invalid_reason=None,
upstream_metadata={},
)
state_updates: dict[str, dict[str, Any]] = {}
result = await refresh_kiro_key_quota(
db=cast(Any, db),
provider=cast(Any, SimpleNamespace(proxy=None)),
key=cast(Any, key),
endpoint=None,
codex_wham_usage_url="https://example.test",
metadata_updates={},
state_updates=state_updates,
)
assert result["status"] == "error"
assert "401" in result["message"]
assert key.is_active is True
assert key.oauth_invalid_reason is None
assert state_updates["k1"]["is_active"] is False
assert state_updates["k1"]["oauth_invalid_reason"] == "Kiro Token 无效或已过期"
assert db.commit_count == 0
@pytest.mark.asyncio
async def test_kiro_refresher_success_updates_metadata_and_auth_config(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from src.services.provider_keys.quota_refresh import kiro_refresher as module
class _Banned(Exception):
pass
async def _fetch_limits(auth_config: dict[str, Any], proxy_config: object) -> Any:
_ = auth_config, proxy_config
return {"usage_data": {"x": 1}, "updated_auth_config": {"token": "new"}}
def _parse_usage(_usage: Any) -> dict[str, Any]:
return {"quota": 1}
_install_module(
monkeypatch,
"src.services.provider.adapters.kiro.usage",
{
"KiroAccountBannedException": _Banned,
"fetch_kiro_usage_limits": _fetch_limits,
"parse_kiro_usage_response": _parse_usage,
},
)
_install_module(
monkeypatch,
"src.services.proxy_node.resolver",
{"resolve_effective_proxy": lambda provider_proxy, key_proxy: None},
)
monkeypatch.setattr(module.crypto_service, "decrypt", lambda _v: json.dumps({"seed": 1}))
monkeypatch.setattr(module.crypto_service, "encrypt", lambda raw: f"ENC:{raw}")
key = SimpleNamespace(
id="k1",
name="K1",
auth_config="enc",
proxy=None,
oauth_invalid_at="old",
oauth_invalid_reason="old",
upstream_metadata={},
)
metadata_updates: dict[str, dict[str, Any]] = {}
state_updates: dict[str, dict[str, Any]] = {}
result = await refresh_kiro_key_quota(
db=cast(Any, _FakeDB()),
provider=cast(Any, SimpleNamespace(proxy=None)),
key=cast(Any, key),
endpoint=None,
codex_wham_usage_url="https://example.test",
metadata_updates=metadata_updates,
state_updates=state_updates,
)
assert result["status"] == "success"
assert metadata_updates["k1"]["kiro"]["is_banned"] is False
assert metadata_updates["k1"]["kiro"]["quota"] == 1
assert key.auth_config == "enc"
assert state_updates["k1"]["oauth_invalid_at"] is None
assert state_updates["k1"]["oauth_invalid_reason"] is None
assert state_updates["k1"]["auth_config"].startswith("ENC:")

View File

@@ -0,0 +1,300 @@
from __future__ import annotations
from types import SimpleNamespace
from typing import Any, cast
import pytest
from src.core.exceptions import InvalidRequestException
from src.core.provider_types import ProviderType
from src.services.provider_keys import key_quota_service as quota_service_module
from src.services.provider_keys.key_quota_service import (
_resolve_quota_refresh_handler,
_select_refresh_endpoint,
refresh_provider_quota_for_provider,
)
from src.services.provider_keys.quota_refresh.antigravity_refresher import (
refresh_antigravity_key_quota,
)
from src.services.provider_keys.quota_refresh.codex_refresher import refresh_codex_key_quota
from src.services.provider_keys.quota_refresh.kiro_refresher import refresh_kiro_key_quota
def _provider_with_endpoints(*endpoints: tuple[str, bool]) -> SimpleNamespace:
eps = [SimpleNamespace(api_format=fmt, is_active=active) for fmt, active in endpoints]
return SimpleNamespace(endpoints=eps)
class _FakeQuery:
def __init__(
self,
*,
first_result: Any = None,
all_result: list[Any] | None = None,
) -> None:
self._first_result = first_result
self._all_result = all_result or []
def filter(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def first(self) -> Any:
return self._first_result
def all(self) -> list[Any]:
return self._all_result
class _FakeDB:
def __init__(self, *, provider: Any, keys: list[SimpleNamespace]) -> None:
self._provider = provider
self._keys = keys
self.added: list[object] = []
self.commit_count = 0
def query(self, model: Any) -> _FakeQuery:
model_name = getattr(model, "__name__", "")
if model_name == "Provider":
return _FakeQuery(first_result=self._provider)
if model_name == "ProviderAPIKey":
return _FakeQuery(all_result=self._keys)
raise AssertionError(f"unexpected query model: {model}")
def add(self, obj: object) -> None:
self.added.append(obj)
def commit(self) -> None:
self.commit_count += 1
def test_select_refresh_endpoint_codex() -> None:
provider = _provider_with_endpoints(("openai:chat", True), ("openai:cli", True))
endpoint = _select_refresh_endpoint(cast(Any, provider), ProviderType.CODEX)
assert endpoint is not None
assert endpoint.api_format == "openai:cli"
def test_select_refresh_endpoint_codex_normalized_api_format() -> None:
provider = _provider_with_endpoints(("openai:chat", True), (" OpenAI:CLI ", True))
endpoint = _select_refresh_endpoint(cast(Any, provider), ProviderType.CODEX)
assert endpoint is not None
assert endpoint.api_format == " OpenAI:CLI "
def test_select_refresh_endpoint_antigravity_prefers_chat() -> None:
provider = _provider_with_endpoints(("gemini:cli", True), ("gemini:chat", True))
endpoint = _select_refresh_endpoint(cast(Any, provider), ProviderType.ANTIGRAVITY)
assert endpoint is not None
assert endpoint.api_format == "gemini:chat"
def test_select_refresh_endpoint_antigravity_fallback_cli() -> None:
provider = _provider_with_endpoints(("gemini:chat", False), ("gemini:cli", True))
endpoint = _select_refresh_endpoint(cast(Any, provider), ProviderType.ANTIGRAVITY)
assert endpoint is not None
assert endpoint.api_format == "gemini:cli"
def test_select_refresh_endpoint_kiro_returns_none() -> None:
provider = _provider_with_endpoints(("openai:cli", True))
endpoint = _select_refresh_endpoint(cast(Any, provider), ProviderType.KIRO)
assert endpoint is None
def test_select_refresh_endpoint_codex_missing_raises() -> None:
provider = _provider_with_endpoints(("openai:chat", True))
with pytest.raises(InvalidRequestException, match="找不到有效的 openai:cli 端点"):
_select_refresh_endpoint(cast(Any, provider), ProviderType.CODEX)
def test_select_refresh_endpoint_antigravity_missing_raises() -> None:
provider = _provider_with_endpoints(("gemini:chat", False), ("gemini:cli", False))
with pytest.raises(InvalidRequestException, match="找不到有效的 gemini:chat/gemini:cli 端点"):
_select_refresh_endpoint(cast(Any, provider), ProviderType.ANTIGRAVITY)
def test_resolve_quota_refresh_handler() -> None:
assert _resolve_quota_refresh_handler(ProviderType.CODEX) is refresh_codex_key_quota
assert _resolve_quota_refresh_handler(ProviderType.ANTIGRAVITY) is refresh_antigravity_key_quota
assert _resolve_quota_refresh_handler(ProviderType.KIRO) is refresh_kiro_key_quota
def test_resolve_quota_refresh_handler_unsupported_raises() -> None:
with pytest.raises(
InvalidRequestException,
match="仅支持 Codex / Antigravity / Kiro 类型的 Provider 刷新限额",
):
_resolve_quota_refresh_handler("unknown")
@pytest.mark.asyncio
async def test_refresh_provider_quota_no_active_keys_returns_empty() -> None:
provider = SimpleNamespace(
id="p1",
provider_type=ProviderType.CODEX,
endpoints=[SimpleNamespace(api_format="openai:cli", is_active=True)],
)
db = _FakeDB(provider=provider, keys=[])
result = await refresh_provider_quota_for_provider(
db=cast(Any, db),
provider_id="p1",
codex_wham_usage_url="https://example.test/wham/usage",
)
assert result == {
"success": 0,
"failed": 0,
"total": 0,
"results": [],
"message": "没有活跃的 Key",
}
@pytest.mark.asyncio
async def test_refresh_provider_quota_aggregates_and_merges_metadata(
monkeypatch: pytest.MonkeyPatch,
) -> None:
provider = SimpleNamespace(
id="p1",
provider_type=ProviderType.CODEX,
endpoints=[SimpleNamespace(api_format="openai:cli", is_active=True)],
)
key1 = SimpleNamespace(id="k1", name="K1", upstream_metadata={"old": True})
key2 = SimpleNamespace(id="k2", name="K2", upstream_metadata={})
db = _FakeDB(provider=provider, keys=[key1, key2])
async def _fake_handler(**kwargs: Any) -> dict[str, Any]:
key = kwargs["key"]
metadata_updates = kwargs["metadata_updates"]
if key.id == "k1":
metadata_updates[key.id] = {"codex": {"used": 10}}
return {"key_id": key.id, "key_name": key.name, "status": "success"}
return {"key_id": key.id, "key_name": key.name, "status": "error", "message": "boom"}
monkeypatch.setattr(
quota_service_module,
"_select_refresh_endpoint",
lambda provider, provider_type: provider.endpoints[0],
)
monkeypatch.setattr(
quota_service_module, "_resolve_quota_refresh_handler", lambda _: _fake_handler
)
monkeypatch.setattr(
quota_service_module,
"merge_upstream_metadata",
lambda current, updates: {**(current or {}), **updates},
)
result = await refresh_provider_quota_for_provider(
db=cast(Any, db),
provider_id="p1",
codex_wham_usage_url="https://example.test/wham/usage",
)
assert result["success"] == 1
assert result["failed"] == 1
assert result["total"] == 2
assert len(result["results"]) == 2
assert key1.upstream_metadata == {"old": True, "codex": {"used": 10}}
assert db.commit_count == 1
assert db.added == [key1]
@pytest.mark.asyncio
async def test_refresh_provider_quota_applies_state_updates_and_single_commit(
monkeypatch: pytest.MonkeyPatch,
) -> None:
provider = SimpleNamespace(
id="p1",
provider_type=ProviderType.CODEX,
endpoints=[SimpleNamespace(api_format="openai:cli", is_active=True)],
)
key1 = SimpleNamespace(
id="k1",
name="K1",
upstream_metadata={},
is_active=True,
oauth_invalid_at="old",
oauth_invalid_reason="old-reason",
)
key2 = SimpleNamespace(
id="k2",
name="K2",
upstream_metadata={},
is_active=True,
oauth_invalid_at=None,
oauth_invalid_reason=None,
)
db = _FakeDB(provider=provider, keys=[key1, key2])
async def _fake_handler(**kwargs: Any) -> dict[str, Any]:
key = kwargs["key"]
state_updates = kwargs["state_updates"]
if key.id == "k1":
state_updates[key.id] = {"oauth_invalid_at": None, "oauth_invalid_reason": None}
return {"key_id": key.id, "key_name": key.name, "status": "success"}
state_updates[key.id] = {"is_active": False, "oauth_invalid_reason": "401"}
return {"key_id": key.id, "key_name": key.name, "status": "error", "message": "401"}
monkeypatch.setattr(
quota_service_module,
"_select_refresh_endpoint",
lambda provider, provider_type: provider.endpoints[0],
)
monkeypatch.setattr(
quota_service_module, "_resolve_quota_refresh_handler", lambda _: _fake_handler
)
result = await refresh_provider_quota_for_provider(
db=cast(Any, db),
provider_id="p1",
codex_wham_usage_url="https://example.test/wham/usage",
)
assert result["success"] == 1
assert result["failed"] == 1
assert key1.oauth_invalid_at is None
assert key1.oauth_invalid_reason is None
assert key2.is_active is False
assert key2.oauth_invalid_reason == "401"
assert db.commit_count == 1
assert db.added == [key1, key2]
@pytest.mark.asyncio
async def test_refresh_provider_quota_handler_exception_returns_error(
monkeypatch: pytest.MonkeyPatch,
) -> None:
provider = SimpleNamespace(
id="p1",
provider_type=ProviderType.CODEX,
endpoints=[SimpleNamespace(api_format="openai:cli", is_active=True)],
)
key = SimpleNamespace(id="k1", name="K1", upstream_metadata={})
db = _FakeDB(provider=provider, keys=[key])
async def _boom_handler(**kwargs: Any) -> dict[str, Any]:
_ = kwargs
raise RuntimeError("unit-test boom")
monkeypatch.setattr(
quota_service_module,
"_select_refresh_endpoint",
lambda provider, provider_type: provider.endpoints[0],
)
monkeypatch.setattr(
quota_service_module, "_resolve_quota_refresh_handler", lambda _: _boom_handler
)
result = await refresh_provider_quota_for_provider(
db=cast(Any, db),
provider_id="p1",
codex_wham_usage_url="https://example.test/wham/usage",
)
assert result["success"] == 0
assert result["failed"] == 1
assert result["total"] == 1
assert result["results"][0]["status"] == "error"
assert "unit-test boom" in result["results"][0]["message"]