mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
refactor(provider-keys): 拆分 keys 端点逻辑并补充配额刷新测试
将 admin keys 接口的创建、更新、查询、删除、导出与配额刷新逻辑迁移到 provider_keys 服务层 新增 auth_type 归一化、重复校验、响应构建与 quota_refresh(codex/kiro/antigravity)模块,并补充对应单元测试
This commit is contained in:
207
tests/services/test_provider_keys_key_command_service.py
Normal file
207
tests/services/test_provider_keys_key_command_service.py
Normal 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
|
||||
155
tests/services/test_provider_keys_query_service.py
Normal file
155
tests/services/test_provider_keys_query_service.py
Normal 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"}]
|
||||
699
tests/services/test_provider_keys_quota_refresh_strategies.py
Normal file
699
tests/services/test_provider_keys_quota_refresh_strategies.py
Normal 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:")
|
||||
300
tests/services/test_provider_keys_quota_service.py
Normal file
300
tests/services/test_provider_keys_quota_service.py
Normal 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"]
|
||||
Reference in New Issue
Block a user