refactor: 全局适配 ApiFamily/EndpointKind 结构化标识体系

将新的 (ApiFamily, EndpointKind) / `family:kind` 签名体系应用到整个代码库:
- API Handlers: 所有 adapter/handler 使用新的签名格式
- Services: provider, model, usage, cache, auth 等服务层适配
- Database: ProviderEndpoint 新增 api_family/endpoint_kind 字段
- Frontend: Provider 管理、Usage 表格等组件适配
- Tests: 更新所有相关测试用例
This commit is contained in:
fawney19
2026-02-01 17:28:00 +08:00
parent c246ccfc91
commit 7b66505634
219 changed files with 4732 additions and 2545 deletions

View File

@@ -85,7 +85,7 @@ class TestConvertSseLineBasic:
def test_empty_line_returns_empty_list(self) -> None:
handler = MockCliHandler()
ctx = StreamContext(model="test", api_format="OPENAI")
ctx = StreamContext(model="test", api_format="openai:chat")
result = handler._convert_sse_line(ctx, "", [])
@@ -93,7 +93,7 @@ class TestConvertSseLineBasic:
def test_whitespace_line_returns_line(self) -> None:
handler = MockCliHandler()
ctx = StreamContext(model="test", api_format="OPENAI")
ctx = StreamContext(model="test", api_format="openai:chat")
result = handler._convert_sse_line(ctx, " ", [])
@@ -101,7 +101,7 @@ class TestConvertSseLineBasic:
def test_done_marker_returns_as_is(self) -> None:
handler = MockCliHandler()
ctx = StreamContext(model="test", api_format="OPENAI")
ctx = StreamContext(model="test", api_format="openai:chat")
result = handler._convert_sse_line(ctx, "data: [DONE]", [])
@@ -109,7 +109,7 @@ class TestConvertSseLineBasic:
def test_non_data_line_passthrough(self) -> None:
handler = MockCliHandler()
ctx = StreamContext(model="test", api_format="OPENAI")
ctx = StreamContext(model="test", api_format="openai:chat")
result = handler._convert_sse_line(ctx, "event: message_start", [])
@@ -117,7 +117,7 @@ class TestConvertSseLineBasic:
def test_invalid_json_passthrough(self) -> None:
handler = MockCliHandler()
ctx = StreamContext(model="test", api_format="OPENAI")
ctx = StreamContext(model="test", api_format="openai:chat")
result = handler._convert_sse_line(ctx, "data: {invalid json}", [])
@@ -130,9 +130,9 @@ class TestConvertSseLineWithMockConverter:
def test_same_format_returns_original(self) -> None:
"""同格式无需转换"""
handler = MockCliHandler()
ctx = StreamContext(model="test", api_format="OPENAI")
ctx.provider_api_format = "OPENAI"
ctx.client_api_format = "OPENAI"
ctx = StreamContext(model="test", api_format="openai:chat")
ctx.provider_api_format = "openai:chat"
ctx.client_api_format = "openai:chat"
chunk = {"choices": [{"delta": {"content": "hello"}}]}
line = f"data: {json.dumps(chunk)}"
@@ -144,14 +144,14 @@ class TestConvertSseLineWithMockConverter:
def test_state_initialization(self) -> None:
"""测试状态自动初始化
流式转换状态应使用用户请求的原始模型名ctx.model
而非映射后的模型名ctx.mapped_model确保返回给客户端的响应使用原始模型名。
"""
handler = MockCliHandler()
ctx = StreamContext(model="gpt-4", api_format="OPENAI")
ctx.provider_api_format = "OPENAI"
ctx.client_api_format = "OPENAI"
ctx = StreamContext(model="gpt-4", api_format="openai:chat")
ctx.provider_api_format = "openai:chat"
ctx.client_api_format = "openai:chat"
ctx.mapped_model = "claude-3-5-sonnet" # 映射后的模型名(发给上游的)
ctx.request_id = "req_123"
@@ -181,9 +181,9 @@ class TestConvertSseLineOneInManyOut:
def test_openai_to_claude_conversion(self) -> None:
"""测试 OpenAI -> Claude 流式转换"""
handler = MockCliHandler()
ctx = StreamContext(model="gpt-4", api_format="OPENAI")
ctx.provider_api_format = "OPENAI"
ctx.client_api_format = "CLAUDE"
ctx = StreamContext(model="gpt-4", api_format="openai:chat")
ctx.provider_api_format = "openai:chat"
ctx.client_api_format = "claude:chat"
ctx.mapped_model = "claude-3-5-sonnet"
ctx.request_id = "req_test"
@@ -205,9 +205,9 @@ class TestConvertSseLineOneInManyOut:
def test_claude_to_openai_conversion(self) -> None:
"""测试 Claude -> OpenAI 流式转换"""
handler = MockCliHandler()
ctx = StreamContext(model="claude-3-5-sonnet", api_format="CLAUDE")
ctx.provider_api_format = "CLAUDE"
ctx.client_api_format = "OPENAI"
ctx = StreamContext(model="claude-3-5-sonnet", api_format="claude:chat")
ctx.provider_api_format = "claude:chat"
ctx.client_api_format = "openai:chat"
ctx.mapped_model = "gpt-4"
ctx.request_id = "msg_test"
@@ -225,9 +225,9 @@ class TestConvertSseLineOneInManyOut:
def test_multiple_chunks_state_persistence(self) -> None:
"""测试多个 chunk 之间状态持久化"""
handler = MockCliHandler()
ctx = StreamContext(model="gpt-4", api_format="OPENAI")
ctx.provider_api_format = "OPENAI"
ctx.client_api_format = "CLAUDE"
ctx = StreamContext(model="gpt-4", api_format="openai:chat")
ctx.provider_api_format = "openai:chat"
ctx.client_api_format = "claude:chat"
ctx.mapped_model = "claude-3-5-sonnet"
# 第一个 chunk
@@ -249,7 +249,7 @@ class TestStreamContextIntegration:
def test_stream_conversion_state_reset_on_retry(self) -> None:
"""测试重试时重置流式转换状态"""
ctx = StreamContext(model="test", api_format="OPENAI")
ctx = StreamContext(model="test", api_format="openai:chat")
ctx.stream_conversion_state = StreamState(model="test", message_id="123")
ctx.reset_for_retry()
@@ -258,7 +258,7 @@ class TestStreamContextIntegration:
def test_stream_conversion_state_field_exists(self) -> None:
"""测试 StreamContext 有 stream_conversion_state 字段"""
ctx = StreamContext(model="test", api_format="OPENAI")
ctx = StreamContext(model="test", api_format="openai:chat")
assert hasattr(ctx, "stream_conversion_state")
assert ctx.stream_conversion_state is None

View File

@@ -3,7 +3,7 @@ from src.api.handlers.base.stream_context import StreamContext
def test_collected_text_append_and_property() -> None:
ctx = StreamContext(model="test-model", api_format="OPENAI")
ctx = StreamContext(model="test-model", api_format="openai:chat")
assert ctx.collected_text == ""
ctx.append_text("hello")
@@ -13,7 +13,7 @@ def test_collected_text_append_and_property() -> None:
def test_reset_for_retry_clears_state() -> None:
ctx = StreamContext(model="test-model", api_format="OPENAI")
ctx = StreamContext(model="test-model", api_format="openai:chat")
ctx.append_text("x")
ctx.update_usage(input_tokens=10, output_tokens=5)
ctx.parsed_chunks.append({"type": "chunk"})

View File

@@ -1,6 +1,11 @@
from typing import Any
from src.api.handlers.base.response_parser import ParsedChunk, ParsedResponse, ResponseParser, StreamStats
from src.api.handlers.base.response_parser import (
ParsedChunk,
ParsedResponse,
ResponseParser,
StreamStats,
)
from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.base.stream_processor import StreamProcessor
from src.utils.sse_parser import SSEEventParser
@@ -21,7 +26,7 @@ class DummyParser(ResponseParser):
def test_process_line_strips_newlines_and_finalizes_event() -> None:
ctx = StreamContext(model="test-model", api_format="OPENAI")
ctx = StreamContext(model="test-model", api_format="openai:chat")
processor = StreamProcessor(request_id="test-request", default_parser=DummyParser())
sse_parser = SSEEventParser()
@@ -29,4 +34,3 @@ def test_process_line_strips_newlines_and_finalizes_event() -> None:
processor._process_line(ctx, sse_parser, "\n")
assert ctx.has_completion is True

View File

@@ -33,9 +33,9 @@ async def _empty_async_iter():
async def test_create_response_stream_converts_claude_to_openai() -> None:
register_default_normalizers()
ctx = StreamContext(model="test-model", api_format="OPENAI")
ctx.client_api_format = "OPENAI"
ctx.provider_api_format = "CLAUDE"
ctx = StreamContext(model="test-model", api_format="openai:chat")
ctx.client_api_format = "openai:chat"
ctx.provider_api_format = "claude:chat"
ctx.needs_conversion = True
processor = StreamProcessor(request_id="test-request", default_parser=DummyParser())

View File

@@ -28,9 +28,9 @@ async def _iter_bytes(chunks: list[bytes]) -> AsyncIterator[bytes]:
async def test_stream_processor_converts_gemini_json_lines_without_data_prefix() -> None:
register_default_normalizers()
ctx = StreamContext(model="gemini-test", api_format="OPENAI")
ctx.provider_api_format = "GEMINI"
ctx.client_api_format = "OPENAI"
ctx = StreamContext(model="gemini-test", api_format="openai:chat")
ctx.provider_api_format = "gemini:chat"
ctx.client_api_format = "openai:chat"
ctx.needs_conversion = True
ctx.request_id = "req_test"
ctx.mapped_model = "gemini-test"
@@ -62,7 +62,7 @@ async def test_stream_processor_converts_gemini_json_lines_without_data_prefix()
processor = StreamProcessor(
request_id="req_test",
default_parser=get_parser_for_format("OPENAI"),
default_parser=get_parser_for_format("openai:chat"),
)
out = b""

View File

@@ -17,4 +17,3 @@ def test_is_done_event_false_when_no_candidates_or_reason() -> None:
parser = GeminiStreamParser()
assert parser.is_done_event({}) is False
assert parser.is_done_event({"candidates": [{}]}) is False

View File

@@ -7,10 +7,10 @@ API Pipeline 测试
- 审计日志记录
"""
import pytest
from datetime import datetime, timezone
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from src.api.base.pipeline import ApiRequestPipeline
@@ -142,7 +142,11 @@ class TestPipelineAuditLogging:
) as mock_log:
with patch("time.time", return_value=1001.0):
pipeline._record_audit_event(
mock_context, mock_adapter, success=False, status_code=500, error="Internal error"
mock_context,
mock_adapter,
success=False,
status_code=500,
error="Internal error",
)
mock_log.assert_called_once()

View File

@@ -0,0 +1,210 @@
from types import SimpleNamespace
from typing import Any
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
from src.config.settings import config
from src.services.task.orchestrator import (
AllCandidatesFailedError,
AsyncTaskOrchestrator,
UpstreamClientRequestError,
)
def _make_candidate(
*,
provider_id: str = "p1",
provider_name: str = "prov",
endpoint_id: str = "e1",
key_id: str = "k1",
key_name: str = "key",
auth_type: str = "api_key",
priority: int = 0,
is_cached: bool = False,
is_skipped: bool = False,
skip_reason: str | None = None,
needs_conversion: bool = False,
) -> SimpleNamespace:
provider = SimpleNamespace(id=provider_id, name=provider_name)
endpoint = SimpleNamespace(id=endpoint_id)
key = SimpleNamespace(id=key_id, name=key_name, auth_type=auth_type, priority=priority)
return SimpleNamespace(
provider=provider,
endpoint=endpoint,
key=key,
is_cached=is_cached,
is_skipped=is_skipped,
skip_reason=skip_reason,
needs_conversion=needs_conversion,
)
@pytest.mark.asyncio
async def test_submit_with_failover_skips_http_500_then_succeeds(
monkeypatch: pytest.MonkeyPatch,
) -> None:
db = MagicMock()
orch = AsyncTaskOrchestrator(db)
# bypass init
orch._candidate_resolver = SimpleNamespace(
fetch_candidates=AsyncMock(
return_value=(
[
_make_candidate(provider_id="p1", endpoint_id="e1", key_id="k1"),
_make_candidate(provider_id="p2", endpoint_id="e2", key_id="k2"),
],
"gm1",
)
)
)
orch._error_classifier = SimpleNamespace(is_client_error=lambda _text: False)
orch._ensure_initialized = AsyncMock(return_value=None)
responses = [
httpx.Response(500, text='{"error": {"message": "server"}}'),
httpx.Response(200, json={"id": "task-123"}),
]
submit = AsyncMock(side_effect=responses)
outcome = await orch.submit_with_failover(
api_format="openai:video",
model_name="sora",
affinity_key="a1",
user_api_key=MagicMock(),
request_id="req-1",
task_type="video",
submit_func=submit,
extract_external_task_id=lambda payload: payload.get("id"),
supported_auth_types={"api_key"},
allow_format_conversion=False,
max_candidates=10,
)
assert outcome.external_task_id == "task-123"
assert outcome.candidate.provider.id == "p2"
assert outcome.candidate_keys[1]["selected"] is True
assert submit.await_count == 2
@pytest.mark.asyncio
async def test_submit_with_failover_stops_on_client_error(monkeypatch: pytest.MonkeyPatch) -> None:
db = MagicMock()
orch = AsyncTaskOrchestrator(db)
orch._candidate_resolver = SimpleNamespace(
fetch_candidates=AsyncMock(return_value=([_make_candidate()], "gm1"))
)
orch._error_classifier = SimpleNamespace(is_client_error=lambda _text: True)
orch._ensure_initialized = AsyncMock(return_value=None)
response = httpx.Response(
400,
json={"error": {"type": "invalid_request_error", "message": "bad request"}},
)
submit = AsyncMock(return_value=response)
with pytest.raises(UpstreamClientRequestError):
await orch.submit_with_failover(
api_format="openai:video",
model_name="sora",
affinity_key="a1",
user_api_key=MagicMock(),
request_id="req-2",
task_type="video",
submit_func=submit,
extract_external_task_id=lambda payload: payload.get("id"),
supported_auth_types={"api_key"},
allow_format_conversion=False,
max_candidates=10,
)
@pytest.mark.asyncio
async def test_submit_with_failover_no_eligible_candidates_due_to_auth_type(
monkeypatch: pytest.MonkeyPatch,
) -> None:
db = MagicMock()
orch = AsyncTaskOrchestrator(db)
orch._candidate_resolver = SimpleNamespace(
fetch_candidates=AsyncMock(return_value=([_make_candidate(auth_type="vertex_ai")], "gm1"))
)
orch._error_classifier = SimpleNamespace(is_client_error=lambda _text: False)
orch._ensure_initialized = AsyncMock(return_value=None)
with pytest.raises(AllCandidatesFailedError) as excinfo:
await orch.submit_with_failover(
api_format="openai:video",
model_name="sora",
affinity_key="a1",
user_api_key=MagicMock(),
request_id="req-3",
task_type="video",
submit_func=AsyncMock(),
extract_external_task_id=lambda payload: payload.get("id"),
supported_auth_types={"api_key"},
allow_format_conversion=False,
max_candidates=10,
)
assert excinfo.value.reason == "no_eligible_candidates"
@pytest.mark.asyncio
async def test_submit_with_failover_filters_missing_billing_rule(
monkeypatch: pytest.MonkeyPatch,
) -> None:
db = MagicMock()
orch = AsyncTaskOrchestrator(db)
orch._candidate_resolver = SimpleNamespace(
fetch_candidates=AsyncMock(
return_value=(
[
_make_candidate(provider_id="p1", endpoint_id="e1", key_id="k1"),
_make_candidate(provider_id="p2", endpoint_id="e2", key_id="k2"),
],
"gm1",
)
)
)
orch._error_classifier = SimpleNamespace(is_client_error=lambda _text: False)
orch._ensure_initialized = AsyncMock(return_value=None)
# enable require_rule
old = config.billing_require_rule
try:
config.billing_require_rule = True
def _find_rule(
_db: Any, *, provider_id: str, model_name: str, task_type: str
) -> object | None:
return None if provider_id == "p1" else object()
monkeypatch.setattr(
"src.services.task.orchestrator.BillingRuleService.find_rule", _find_rule
)
submit = AsyncMock(return_value=httpx.Response(200, json={"id": "task-999"}))
outcome = await orch.submit_with_failover(
api_format="openai:video",
model_name="sora",
affinity_key="a1",
user_api_key=MagicMock(),
request_id="req-4",
task_type="video",
submit_func=submit,
extract_external_task_id=lambda payload: payload.get("id"),
supported_auth_types={"api_key"},
allow_format_conversion=False,
max_candidates=10,
)
assert outcome.candidate.provider.id == "p2"
assert outcome.external_task_id == "task-999"
# only called once because p1 skipped
assert submit.await_count == 1
finally:
config.billing_require_rule = old

View File

@@ -7,20 +7,20 @@
- API Key 认证
"""
import pytest
from datetime import datetime, timedelta, timezone
from unittest.mock import AsyncMock, MagicMock, patch
import jwt
import pytest
from src.services.auth.service import (
AuthService,
JWT_SECRET_KEY,
JWT_ALGORITHM,
JWT_EXPIRATION_HOURS,
)
from src.core.enums import AuthSource
from src.models.database import UserRole
from src.services.auth.service import (
JWT_ALGORITHM,
JWT_EXPIRATION_HOURS,
JWT_SECRET_KEY,
AuthService,
)
class TestJWTTokenCreation:

View File

@@ -10,9 +10,10 @@
期望Provider B 应该被跳过,因为它不支持 haiku 模型
"""
import pytest
from unittest.mock import MagicMock, patch
import pytest
from src.models.database import GlobalModel, Model, Provider
from src.services.cache.aware_scheduler import CacheAwareScheduler

View File

@@ -94,9 +94,7 @@ class TestThinkingErrorPatterns:
error = '{"error": {"message": "expected redacted_thinking, found text"}}'
assert classifier._is_thinking_error(error) is True
def test_detect_expected_redacted_thinking_backticks(
self, classifier: ErrorClassifier
) -> None:
def test_detect_expected_redacted_thinking_backticks(self, classifier: ErrorClassifier) -> None:
"""检测带反引号的 expected redacted_thinking 错误"""
error = '{"error": {"message": "expected `redacted_thinking`, found `text`"}}'
assert classifier._is_thinking_error(error) is True

View File

@@ -1,7 +1,7 @@
import pytest
from unittest.mock import AsyncMock, MagicMock
from src.core.api_format import APIFormat
import pytest
from src.core.api_format.conversion import register_default_normalizers
from src.services.cache.aware_scheduler import CacheAwareScheduler
@@ -18,9 +18,11 @@ def _mock_key(key_id: str, api_formats: list[str]) -> MagicMock:
def _mock_endpoint(api_format: str, config: dict | None = None) -> MagicMock:
endpoint = MagicMock()
endpoint.id = f"ep_{api_format.lower()}"
endpoint.id = f"ep_{api_format.lower().replace(':', '_')}"
endpoint.is_active = True
endpoint.api_format = api_format
endpoint.api_family = api_format.split(":", 1)[0]
endpoint.endpoint_kind = api_format.split(":", 1)[1]
endpoint.format_acceptance_config = config
return endpoint
@@ -37,16 +39,16 @@ async def test_build_candidates_blocks_cross_format_when_global_switch_off() ->
provider.name = "p1"
provider.endpoints = [
_mock_endpoint(
"OPENAI",
{"enabled": True, "accept_formats": ["CLAUDE"], "stream_conversion": True},
"openai:chat",
{"enabled": True, "accept_formats": ["claude:chat"], "stream_conversion": True},
)
]
provider.api_keys = [_mock_key("k1", ["OPENAI"])]
provider.api_keys = [_mock_key("k1", ["openai:chat"])]
candidates = await scheduler._build_candidates(
db=MagicMock(),
providers=[provider],
client_format=APIFormat.CLAUDE,
client_format="claude:chat",
model_name="dummy-model",
affinity_key=None,
global_conversion_enabled=False,
@@ -67,16 +69,16 @@ async def test_build_candidates_includes_cross_format_when_enabled() -> None:
provider.name = "p1"
provider.endpoints = [
_mock_endpoint(
"OPENAI",
{"enabled": True, "accept_formats": ["CLAUDE"], "stream_conversion": True},
"openai:chat",
{"enabled": True, "accept_formats": ["claude:chat"], "stream_conversion": True},
)
]
provider.api_keys = [_mock_key("k1", ["OPENAI"])]
provider.api_keys = [_mock_key("k1", ["openai:chat"])]
candidates = await scheduler._build_candidates(
db=MagicMock(),
providers=[provider],
client_format=APIFormat.CLAUDE,
client_format="claude:chat",
model_name="dummy-model",
affinity_key=None,
global_conversion_enabled=True,
@@ -84,7 +86,7 @@ async def test_build_candidates_includes_cross_format_when_enabled() -> None:
assert len(candidates) == 1
assert candidates[0].needs_conversion is True
assert candidates[0].provider_api_format == "OPENAI"
assert candidates[0].provider_api_format == "openai:chat"
@pytest.mark.asyncio
@@ -100,20 +102,20 @@ async def test_exact_matches_rank_before_convertible() -> None:
# 故意把 OPENAI 放在 endpoints[0],验证排序仍然是 CLAUDEexact在前
provider.endpoints = [
_mock_endpoint(
"OPENAI",
{"enabled": True, "accept_formats": ["CLAUDE"], "stream_conversion": True},
"openai:chat",
{"enabled": True, "accept_formats": ["claude:chat"], "stream_conversion": True},
),
_mock_endpoint("CLAUDE", None),
_mock_endpoint("claude:chat", None),
]
provider.api_keys = [
_mock_key("k_openai", ["OPENAI"]),
_mock_key("k_claude", ["CLAUDE"]),
_mock_key("k_openai", ["openai:chat"]),
_mock_key("k_claude", ["claude:chat"]),
]
candidates = await scheduler._build_candidates(
db=MagicMock(),
providers=[provider],
client_format=APIFormat.CLAUDE,
client_format="claude:chat",
model_name="dummy-model",
affinity_key=None,
global_conversion_enabled=True,
@@ -121,6 +123,6 @@ async def test_exact_matches_rank_before_convertible() -> None:
assert len(candidates) == 2
assert candidates[0].needs_conversion is False
assert candidates[0].provider_api_format == "CLAUDE"
assert candidates[0].provider_api_format == "claude:chat"
assert candidates[1].needs_conversion is True
assert candidates[1].provider_api_format == "OPENAI"
assert candidates[1].provider_api_format == "openai:chat"

View File

@@ -3,7 +3,6 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from src.core.api_format import APIFormat
from src.services.orchestration.error_classifier import ErrorClassifier
from src.services.request.executor import RequestExecutor
@@ -30,7 +29,9 @@ async def test_executor_records_health_by_provider_format() -> None:
endpoint = MagicMock()
endpoint.id = "e1"
endpoint.api_format = "OPENAI"
endpoint.api_format = "openai:chat"
endpoint.api_family = "openai"
endpoint.endpoint_kind = "chat"
key = MagicMock()
key.id = "k1"
@@ -48,16 +49,19 @@ async def test_executor_records_health_by_provider_format() -> None:
async def request_func(_provider, _endpoint, _key, _candidate): # noqa: ANN001
return {"ok": True}
with patch("src.services.request.executor.RequestCandidateService.mark_candidate_started"), patch(
"src.services.request.executor.RequestCandidateService.mark_candidate_success"
), patch("src.services.request.executor.get_adaptive_reservation_manager") as mock_res_mgr, patch(
"src.services.request.executor.health_monitor.record_success"
) as record_success:
with (
patch("src.services.request.executor.RequestCandidateService.mark_candidate_started"),
patch("src.services.request.executor.RequestCandidateService.mark_candidate_success"),
patch("src.services.request.executor.get_adaptive_reservation_manager") as mock_res_mgr,
patch("src.services.request.executor.health_monitor.record_success") as record_success,
):
mock_res_mgr.return_value.calculate_reservation.return_value = MagicMock(
ratio=0.0, phase="stable", confidence=1.0
)
executor = RequestExecutor(db=db, concurrency_manager=concurrency_manager, adaptive_manager=adaptive_manager)
executor = RequestExecutor(
db=db, concurrency_manager=concurrency_manager, adaptive_manager=adaptive_manager
)
await executor.execute(
candidate=candidate,
candidate_id="c1",
@@ -65,13 +69,13 @@ async def test_executor_records_health_by_provider_format() -> None:
user_api_key=MagicMock(user_id="u1", id="ak1"),
request_func=request_func,
request_id="r1",
api_format=APIFormat.CLAUDE, # client_format
api_format="claude:chat", # client_format
model_name="m",
is_stream=False,
)
record_success.assert_called()
assert record_success.call_args.kwargs["api_format"] == "OPENAI"
assert record_success.call_args.kwargs["api_format"] == "openai:chat"
@pytest.mark.asyncio
@@ -84,19 +88,23 @@ async def test_error_classifier_records_failure_by_provider_format() -> None:
endpoint = MagicMock()
endpoint.id = "e1"
endpoint.api_format = "OPENAI"
endpoint.api_format = "openai:chat"
endpoint.api_family = "openai"
endpoint.endpoint_kind = "chat"
key = MagicMock()
key.id = "k1"
with patch("src.services.orchestration.error_classifier.health_monitor.record_failure") as record_failure:
with patch(
"src.services.orchestration.error_classifier.health_monitor.record_failure"
) as record_failure:
await classifier.handle_retriable_error(
error=RuntimeError("boom"),
provider=provider,
endpoint=endpoint,
key=key,
affinity_key="aff",
api_format=APIFormat.CLAUDE, # client_format
api_format="claude:chat", # client_format
global_model_id="gm1",
captured_key_concurrent=None,
elapsed_ms=None,
@@ -106,5 +114,4 @@ async def test_error_classifier_records_failure_by_provider_format() -> None:
)
record_failure.assert_called()
assert record_failure.call_args.kwargs["api_format"] == "OPENAI"
assert record_failure.call_args.kwargs["api_format"] == "openai:chat"

View File

@@ -15,7 +15,6 @@ from loguru import logger
from src.services.model.availability import ModelAvailabilityQuery
# ============================================================================
# 测试辅助类
# ============================================================================
@@ -63,12 +62,12 @@ class TestBaseActiveModelsSourceCode:
"""应使用内连接 GlobalModel排除 global_model_id=NULL"""
source = inspect.getsource(ModelAvailabilityQuery.base_active_models)
assert "join(Model.global_model)" in source or ".join(Model.global_model)" in source, (
"base_active_models 应使用 join(Model.global_model) 内连接"
)
assert "outerjoin" not in source.lower(), (
"base_active_models 不应使用 outerjoin会返回 global_model_id=NULL 的记录)"
)
assert (
"join(Model.global_model)" in source or ".join(Model.global_model)" in source
), "base_active_models 应使用 join(Model.global_model) 内连接"
assert (
"outerjoin" not in source.lower()
), "base_active_models 不应使用 outerjoin会返回 global_model_id=NULL 的记录)"
def test_filters_is_available_true_or_null(self) -> None:
"""应过滤 is_available = True 或 NULL"""
@@ -76,9 +75,9 @@ class TestBaseActiveModelsSourceCode:
assert "or_(" in source, "base_active_models 应使用 or_() 处理 is_available"
assert "is_available" in source, "base_active_models 应包含 is_available 条件"
assert "is_(True)" in source and "is_(None)" in source, (
"base_active_models 应同时检查 is_available = True 和 NULL"
)
assert (
"is_(True)" in source and "is_(None)" in source
), "base_active_models 应同时检查 is_available = True 和 NULL"
def test_filters_all_is_active_fields(self) -> None:
"""应过滤 Model/Provider/GlobalModel 的 is_active"""
@@ -97,7 +96,7 @@ class TestGetProviderKeyRules:
result = ModelAvailabilityQuery.get_provider_key_rules(
FakeSession([]),
provider_ids=set(),
api_formats=["OPENAI"],
api_formats=["openai:chat"],
provider_to_endpoint_formats={},
)
assert result == {}
@@ -114,13 +113,13 @@ class TestGetProviderKeyRules:
try:
# (key_id, provider_id, allowed_models, api_formats)
data = [("key-1", "provider-1", "invalid-string-type", ["OPENAI"])]
data = [("key-1", "provider-1", "invalid-string-type", ["openai:chat"])]
result = ModelAvailabilityQuery.get_provider_key_rules(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
api_formats=["openai:chat"],
provider_to_endpoint_formats={"provider-1": {"openai:chat"}},
)
assert "provider-1" not in result or len(result.get("provider-1", [])) == 0
@@ -143,13 +142,13 @@ class TestGetProviderKeyRules:
try:
# api_formats 是字符串而非列表
data = [("key-1", "provider-1", None, "OPENAI")]
data = [("key-1", "provider-1", None, "openai:chat")]
result = ModelAvailabilityQuery.get_provider_key_rules(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
api_formats=["openai:chat"],
provider_to_endpoint_formats={"provider-1": {"openai:chat"}},
)
assert "provider-1" not in result or len(result.get("provider-1", [])) == 0
@@ -160,27 +159,32 @@ class TestGetProviderKeyRules:
logger.remove(handler_id)
def test_skips_key_with_none_api_formats(self) -> None:
"""api_formats 为 None 时应跳过该 Key(不打日志)"""
"""api_formats 为 None 时应使用 endpoint 支持的格式(不打日志)"""
data = [("key-1", "provider-1", None, None)]
result = ModelAvailabilityQuery.get_provider_key_rules(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
api_formats=["openai:chat"],
provider_to_endpoint_formats={"provider-1": {"openai:chat"}},
)
assert "provider-1" not in result or len(result.get("provider-1", [])) == 0
# api_formats=None 表示使用 endpoint_formats
assert "provider-1" in result
rules = result["provider-1"]
assert len(rules) == 1
allowed_models, usable_formats = rules[0]
assert "openai:chat" in usable_formats
def test_allowed_models_none_means_no_restriction(self) -> None:
"""allowed_models = None 表示不限制模型"""
data = [("key-1", "provider-1", None, ["OPENAI"])]
data = [("key-1", "provider-1", None, ["openai:chat"])]
result = ModelAvailabilityQuery.get_provider_key_rules(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
api_formats=["openai:chat"],
provider_to_endpoint_formats={"provider-1": {"openai:chat"}},
)
assert "provider-1" in result
@@ -188,17 +192,17 @@ class TestGetProviderKeyRules:
assert len(rules) == 1
allowed_models, usable_formats = rules[0]
assert allowed_models is None # None = 不限制
assert "OPENAI" in usable_formats
assert "openai:chat" in usable_formats
def test_allowed_models_list_is_preserved(self) -> None:
"""allowed_models 为有效列表时应正常返回"""
data = [("key-1", "provider-1", ["claude-3-opus", "claude-3-sonnet"], ["OPENAI"])]
data = [("key-1", "provider-1", ["claude-3-opus", "claude-3-sonnet"], ["openai:chat"])]
result = ModelAvailabilityQuery.get_provider_key_rules(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
api_formats=["openai:chat"],
provider_to_endpoint_formats={"provider-1": {"openai:chat"}},
)
assert "provider-1" in result
@@ -210,13 +214,13 @@ class TestGetProviderKeyRules:
def test_skips_key_when_format_intersection_empty(self) -> None:
"""格式交集为空时不包含该 Key"""
# Key 支持 CLAUDE但请求的是 OPENAI
data = [("key-1", "provider-1", None, ["CLAUDE"])]
data = [("key-1", "provider-1", None, ["claude:chat"])]
result = ModelAvailabilityQuery.get_provider_key_rules(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
api_formats=["openai:chat"],
provider_to_endpoint_formats={"provider-1": {"openai:chat"}},
)
assert "provider-1" not in result or len(result.get("provider-1", [])) == 0
@@ -224,15 +228,15 @@ class TestGetProviderKeyRules:
def test_multiple_keys_same_provider(self) -> None:
"""同一 Provider 下多个 Key 应合并规则"""
data = [
("key-1", "provider-1", None, ["OPENAI"]),
("key-2", "provider-1", ["claude-3-opus"], ["OPENAI"]),
("key-1", "provider-1", None, ["openai:chat"]),
("key-2", "provider-1", ["claude-3-opus"], ["openai:chat"]),
]
result = ModelAvailabilityQuery.get_provider_key_rules(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
api_formats=["openai:chat"],
provider_to_endpoint_formats={"provider-1": {"openai:chat"}},
)
assert "provider-1" in result
@@ -248,7 +252,7 @@ class TestGetProvidersWithActiveKeys:
result = ModelAvailabilityQuery.get_providers_with_active_keys(
FakeSession([]),
provider_ids=set(),
api_formats=["OPENAI"],
api_formats=["openai:chat"],
provider_to_endpoint_formats={},
)
assert result == set()
@@ -265,13 +269,13 @@ class TestGetProvidersWithActiveKeys:
try:
# (provider_id, api_formats) - api_formats 是字符串
data = [("provider-1", "OPENAI")]
data = [("provider-1", "openai:chat")]
result = ModelAvailabilityQuery.get_providers_with_active_keys(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
api_formats=["openai:chat"],
provider_to_endpoint_formats={"provider-1": {"openai:chat"}},
)
assert "provider-1" not in result
@@ -282,27 +286,28 @@ class TestGetProvidersWithActiveKeys:
logger.remove(handler_id)
def test_skips_key_with_none_api_formats(self) -> None:
"""api_formats 为 None 时应跳过该 Key不打日志"""
"""api_formats 为 None 时应使用 endpoint 支持的格式(返回该 Provider"""
data = [("provider-1", None)]
result = ModelAvailabilityQuery.get_providers_with_active_keys(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
api_formats=["openai:chat"],
provider_to_endpoint_formats={"provider-1": {"openai:chat"}},
)
assert "provider-1" not in result
# api_formats=None 表示使用 endpoint_formats所以应该返回
assert "provider-1" in result
def test_returns_provider_when_format_matches(self) -> None:
"""格式匹配时应返回该 Provider"""
data = [("provider-1", ["OPENAI", "CLAUDE"])]
data = [("provider-1", ["openai:chat", "claude:chat"])]
result = ModelAvailabilityQuery.get_providers_with_active_keys(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
api_formats=["openai:chat"],
provider_to_endpoint_formats={"provider-1": {"openai:chat"}},
)
assert "provider-1" in result
@@ -310,25 +315,25 @@ class TestGetProvidersWithActiveKeys:
def test_skips_provider_when_format_intersection_empty(self) -> None:
"""格式交集为空时不返回该 Provider"""
# Key 支持 GEMINI但请求的是 OPENAI且端点支持 OPENAI
data = [("provider-1", ["GEMINI"])]
data = [("provider-1", ["gemini:chat"])]
result = ModelAvailabilityQuery.get_providers_with_active_keys(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
api_formats=["openai:chat"],
provider_to_endpoint_formats={"provider-1": {"openai:chat"}},
)
assert "provider-1" not in result
def test_skips_provider_not_in_endpoint_formats(self) -> None:
"""Provider 不在 endpoint_formats 中时应跳过"""
data = [("provider-1", ["OPENAI"])]
data = [("provider-1", ["openai:chat"])]
result = ModelAvailabilityQuery.get_providers_with_active_keys(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
api_formats=["openai:chat"],
provider_to_endpoint_formats={}, # 空,无匹配端点
)
@@ -336,13 +341,13 @@ class TestGetProvidersWithActiveKeys:
def test_case_insensitive_format_matching(self) -> None:
"""格式匹配应忽略大小写"""
data = [("provider-1", ["openai"])] # 小写
data = [("provider-1", ["openai:chat"])] # 小写
result = ModelAvailabilityQuery.get_providers_with_active_keys(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"], # 大写
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
api_formats=["OPENAI:CHAT"], # 大写
provider_to_endpoint_formats={"provider-1": {"openai:chat"}},
)
assert "provider-1" in result
@@ -394,9 +399,9 @@ class TestFindModelByIdNoFallback:
source = inspect.getsource(models_service.find_model_by_id)
assert "provider_model_name" not in source, (
"find_model_by_id 中不应出现 provider_model_name已删除回退逻辑"
)
assert (
"provider_model_name" not in source
), "find_model_by_id 中不应出现 provider_model_name已删除回退逻辑"
class TestGetAvailableModelIdsWithMappings:
@@ -408,9 +413,9 @@ class TestGetAvailableModelIdsWithMappings:
source = inspect.getsource(models_service._get_available_model_ids_for_format)
assert "check_model_allowed_with_mappings" in source, (
"_get_available_model_ids_for_format 应使用 check_model_allowed_with_mappings"
)
assert (
"check_model_allowed_with_mappings" in source
), "_get_available_model_ids_for_format 应使用 check_model_allowed_with_mappings"
def test_source_code_extracts_model_mappings_from_config(self) -> None:
"""静态验证:应从 global_model.config 提取 model_mappings"""
@@ -418,12 +423,12 @@ class TestGetAvailableModelIdsWithMappings:
source = inspect.getsource(models_service._get_available_model_ids_for_format)
assert "model_mappings" in source, (
"_get_available_model_ids_for_format 应提取 model_mappings"
)
assert ".config" in source or "config" in source, (
"_get_available_model_ids_for_format 应从 config 获取 model_mappings"
)
assert (
"model_mappings" in source
), "_get_available_model_ids_for_format 应提取 model_mappings"
assert (
".config" in source or "config" in source
), "_get_available_model_ids_for_format 应从 config 获取 model_mappings"
class TestModelMappingsIntegration:
@@ -521,18 +526,18 @@ class TestAccessRestrictionsFromApiKeyAndUser:
api_key = MagicMock()
api_key.allowed_providers = ["provider-a"]
api_key.allowed_models = ["model-a"]
api_key.allowed_api_formats = ["OPENAI"]
api_key.allowed_api_formats = ["openai:chat"]
user = MagicMock()
user.allowed_providers = ["provider-b"]
user.allowed_models = ["model-b"]
user.allowed_api_formats = ["CLAUDE"]
user.allowed_api_formats = ["claude:chat"]
result = AccessRestrictions.from_api_key_and_user(api_key, user)
assert result.allowed_providers == ["provider-a"]
assert result.allowed_models == ["model-a"]
assert result.allowed_api_formats == ["OPENAI"]
assert result.allowed_api_formats == ["openai:chat"]
def test_user_restrictions_used_when_api_key_has_none(self) -> None:
"""API Key 无限制时使用 User 的限制"""
@@ -548,13 +553,13 @@ class TestAccessRestrictionsFromApiKeyAndUser:
user = MagicMock()
user.allowed_providers = ["provider-b"]
user.allowed_models = ["model-b"]
user.allowed_api_formats = ["CLAUDE"]
user.allowed_api_formats = ["claude:chat"]
result = AccessRestrictions.from_api_key_and_user(api_key, user)
assert result.allowed_providers == ["provider-b"]
assert result.allowed_models == ["model-b"]
assert result.allowed_api_formats == ["CLAUDE"]
assert result.allowed_api_formats == ["claude:chat"]
def test_partial_api_key_restrictions(self) -> None:
"""API Key 部分限制时,其余字段从 User 获取"""
@@ -570,13 +575,13 @@ class TestAccessRestrictionsFromApiKeyAndUser:
user = MagicMock()
user.allowed_providers = ["provider-b"]
user.allowed_models = ["model-b"]
user.allowed_api_formats = ["CLAUDE"]
user.allowed_api_formats = ["claude:chat"]
result = AccessRestrictions.from_api_key_and_user(api_key, user)
assert result.allowed_providers == ["provider-a"] # 来自 API Key
assert result.allowed_models == ["model-b"] # 来自 User
assert result.allowed_api_formats == ["CLAUDE"] # 来自 User
assert result.allowed_api_formats == ["claude:chat"] # 来自 User
def test_only_api_key_provided(self) -> None:
"""只提供 API Key 时使用其限制"""
@@ -587,13 +592,13 @@ class TestAccessRestrictionsFromApiKeyAndUser:
api_key = MagicMock()
api_key.allowed_providers = ["provider-a"]
api_key.allowed_models = ["model-a"]
api_key.allowed_api_formats = ["OPENAI"]
api_key.allowed_api_formats = ["openai:chat"]
result = AccessRestrictions.from_api_key_and_user(api_key, None)
assert result.allowed_providers == ["provider-a"]
assert result.allowed_models == ["model-a"]
assert result.allowed_api_formats == ["OPENAI"]
assert result.allowed_api_formats == ["openai:chat"]
def test_only_user_provided(self) -> None:
"""只提供 User 时使用其限制"""
@@ -604,13 +609,13 @@ class TestAccessRestrictionsFromApiKeyAndUser:
user = MagicMock()
user.allowed_providers = ["provider-b"]
user.allowed_models = ["model-b"]
user.allowed_api_formats = ["CLAUDE"]
user.allowed_api_formats = ["claude:chat"]
result = AccessRestrictions.from_api_key_and_user(None, user)
assert result.allowed_providers == ["provider-b"]
assert result.allowed_models == ["model-b"]
assert result.allowed_api_formats == ["CLAUDE"]
assert result.allowed_api_formats == ["claude:chat"]
class TestAccessRestrictionsIsApiFormatAllowed:
@@ -622,27 +627,27 @@ class TestAccessRestrictionsIsApiFormatAllowed:
restrictions = AccessRestrictions(allowed_api_formats=None)
assert restrictions.is_api_format_allowed("OPENAI") is True
assert restrictions.is_api_format_allowed("CLAUDE") is True
assert restrictions.is_api_format_allowed("GEMINI") is True
assert restrictions.is_api_format_allowed("openai:chat") is True
assert restrictions.is_api_format_allowed("claude:chat") is True
assert restrictions.is_api_format_allowed("gemini:chat") is True
def test_format_in_allowed_list(self) -> None:
"""格式在允许列表中时返回 True"""
from src.api.base.models_service import AccessRestrictions
restrictions = AccessRestrictions(allowed_api_formats=["OPENAI", "CLAUDE"])
restrictions = AccessRestrictions(allowed_api_formats=["openai:chat", "claude:chat"])
assert restrictions.is_api_format_allowed("OPENAI") is True
assert restrictions.is_api_format_allowed("CLAUDE") is True
assert restrictions.is_api_format_allowed("openai:chat") is True
assert restrictions.is_api_format_allowed("claude:chat") is True
def test_format_not_in_allowed_list(self) -> None:
"""格式不在允许列表中时返回 False"""
from src.api.base.models_service import AccessRestrictions
restrictions = AccessRestrictions(allowed_api_formats=["OPENAI"])
restrictions = AccessRestrictions(allowed_api_formats=["openai:chat"])
assert restrictions.is_api_format_allowed("CLAUDE") is False
assert restrictions.is_api_format_allowed("GEMINI") is False
assert restrictions.is_api_format_allowed("claude:chat") is False
assert restrictions.is_api_format_allowed("gemini:chat") is False
def test_empty_allowed_list_blocks_all(self) -> None:
"""空允许列表阻止所有格式"""
@@ -650,8 +655,8 @@ class TestAccessRestrictionsIsApiFormatAllowed:
restrictions = AccessRestrictions(allowed_api_formats=[])
assert restrictions.is_api_format_allowed("OPENAI") is False
assert restrictions.is_api_format_allowed("CLAUDE") is False
assert restrictions.is_api_format_allowed("openai:chat") is False
assert restrictions.is_api_format_allowed("claude:chat") is False
class TestAccessRestrictionsIsModelAllowed:
@@ -720,5 +725,3 @@ class TestAccessRestrictionsIsModelAllowed:
restrictions = AccessRestrictions(allowed_models=[])
assert restrictions.is_model_allowed("claude-3-opus", "provider-a") is False

View File

@@ -36,7 +36,9 @@ class _FakeSession:
# 如果 direct match 命中,不应再走 provider_model_name 分支
if entities == (Model, GlobalModel):
raise AssertionError("provider_model_name query should not run when direct match exists")
raise AssertionError(
"provider_model_name query should not run when direct match exists"
)
raise AssertionError(f"Unexpected query entities: {entities}")
@@ -64,6 +66,7 @@ async def test_resolve_global_model_prefers_direct_match(monkeypatch) -> None:
)
db = _FakeSession(direct_match=global_model)
resolved = await ModelCacheService.resolve_global_model_by_name_or_mapping(db, global_model.name)
resolved = await ModelCacheService.resolve_global_model_by_name_or_mapping(
db, global_model.name
)
assert resolved is global_model

View File

@@ -1,6 +1,5 @@
from dataclasses import dataclass
from src.services.provider.transport import build_provider_url
@@ -14,7 +13,7 @@ class _DummyEndpoint:
def test_gemini_stream_adds_alt_sse_and_drops_key_query_param() -> None:
endpoint = _DummyEndpoint(
base_url="https://generativelanguage.googleapis.com",
api_format="GEMINI",
api_format="gemini:chat",
)
url = build_provider_url(
@@ -34,7 +33,7 @@ def test_gemini_stream_adds_alt_sse_and_drops_key_query_param() -> None:
def test_gemini_stream_does_not_override_existing_alt() -> None:
endpoint = _DummyEndpoint(
base_url="https://generativelanguage.googleapis.com",
api_format="GEMINI",
api_format="gemini:chat",
)
url = build_provider_url(
@@ -51,7 +50,7 @@ def test_gemini_stream_does_not_override_existing_alt() -> None:
def test_gemini_non_stream_does_not_add_alt() -> None:
endpoint = _DummyEndpoint(
base_url="https://generativelanguage.googleapis.com",
api_format="GEMINI",
api_format="gemini:chat",
)
url = build_provider_url(
@@ -62,4 +61,3 @@ def test_gemini_non_stream_does_not_add_alt() -> None:
assert url.endswith("/v1beta/models/gemini-1.5-pro:generateContent")
assert "alt=" not in url

View File

@@ -1,11 +1,16 @@
import asyncio
import json
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from redis.exceptions import ResponseError
from src.config.settings import config
from src.services.usage.consumer_streams import (
UsageQueueConsumer,
_consumer_name,
ensure_usage_stream_group,
)
from src.services.usage.events import (
UsageEvent,
UsageEventType,
@@ -16,11 +21,6 @@ from src.services.usage.telemetry_writer import (
DbTelemetryWriter,
QueueTelemetryWriter,
)
from src.services.usage.consumer_streams import (
UsageQueueConsumer,
ensure_usage_stream_group,
_consumer_name,
)
class DummyRedis:
@@ -58,9 +58,7 @@ async def test_queue_writer_publishes_event(monkeypatch):
async def _get_redis_client(require_redis=False):
return dummy
monkeypatch.setattr(
"src.services.usage.telemetry_writer.get_redis_client", _get_redis_client
)
monkeypatch.setattr("src.services.usage.telemetry_writer.get_redis_client", _get_redis_client)
old_stream_key = config.usage_queue_stream_key
old_maxlen = config.usage_queue_stream_maxlen
@@ -252,9 +250,7 @@ async def test_queue_writer_record_failure(monkeypatch):
async def _get_redis_client(require_redis=False):
return dummy
monkeypatch.setattr(
"src.services.usage.telemetry_writer.get_redis_client", _get_redis_client
)
monkeypatch.setattr("src.services.usage.telemetry_writer.get_redis_client", _get_redis_client)
old_maxlen = config.usage_queue_stream_maxlen
try:
@@ -291,9 +287,7 @@ async def test_queue_writer_record_cancelled(monkeypatch):
async def _get_redis_client(require_redis=False):
return dummy
monkeypatch.setattr(
"src.services.usage.telemetry_writer.get_redis_client", _get_redis_client
)
monkeypatch.setattr("src.services.usage.telemetry_writer.get_redis_client", _get_redis_client)
writer = QueueTelemetryWriter(
request_id="req-cancel",
@@ -314,9 +308,7 @@ async def test_queue_writer_include_headers_bodies(monkeypatch):
async def _get_redis_client(require_redis=False):
return dummy
monkeypatch.setattr(
"src.services.usage.telemetry_writer.get_redis_client", _get_redis_client
)
monkeypatch.setattr("src.services.usage.telemetry_writer.get_redis_client", _get_redis_client)
old_include_headers = config.usage_queue_include_headers
old_include_bodies = config.usage_queue_include_bodies
@@ -384,9 +376,7 @@ async def test_queue_writer_body_truncation(monkeypatch):
async def _get_redis_client(require_redis=False):
return dummy
monkeypatch.setattr(
"src.services.usage.telemetry_writer.get_redis_client", _get_redis_client
)
monkeypatch.setattr("src.services.usage.telemetry_writer.get_redis_client", _get_redis_client)
old_include_bodies = config.usage_queue_include_bodies
old_max_bytes = config.usage_queue_body_max_bytes
@@ -418,12 +408,11 @@ async def test_queue_writer_body_truncation(monkeypatch):
@pytest.mark.asyncio
async def test_queue_writer_redis_unavailable(monkeypatch):
"""测试 Redis 不可用时抛出异常"""
async def _get_redis_client(require_redis=False):
return None
monkeypatch.setattr(
"src.services.usage.telemetry_writer.get_redis_client", _get_redis_client
)
monkeypatch.setattr("src.services.usage.telemetry_writer.get_redis_client", _get_redis_client)
writer = QueueTelemetryWriter(
request_id="req-no-redis",
@@ -443,9 +432,7 @@ async def test_queue_writer_xadd_error(monkeypatch):
async def _get_redis_client(require_redis=False):
return dummy
monkeypatch.setattr(
"src.services.usage.telemetry_writer.get_redis_client", _get_redis_client
)
monkeypatch.setattr("src.services.usage.telemetry_writer.get_redis_client", _get_redis_client)
writer = QueueTelemetryWriter(
request_id="req-xadd-fail",
@@ -537,6 +524,7 @@ def test_consumer_name():
assert ":" in name
# 应包含 PID
import os
assert str(os.getpid()) in name
@@ -548,9 +536,7 @@ async def test_ensure_stream_group_creates_group(monkeypatch):
async def _get_redis_client(require_redis=False):
return mock_redis
monkeypatch.setattr(
"src.services.usage.consumer_streams.get_redis_client", _get_redis_client
)
monkeypatch.setattr("src.services.usage.consumer_streams.get_redis_client", _get_redis_client)
await ensure_usage_stream_group()
@@ -570,9 +556,7 @@ async def test_ensure_stream_group_handles_busygroup(monkeypatch):
async def _get_redis_client(require_redis=False):
return mock_redis
monkeypatch.setattr(
"src.services.usage.consumer_streams.get_redis_client", _get_redis_client
)
monkeypatch.setattr("src.services.usage.consumer_streams.get_redis_client", _get_redis_client)
# 不应抛出异常
await ensure_usage_stream_group()
@@ -587,9 +571,7 @@ async def test_ensure_stream_group_raises_other_errors(monkeypatch):
async def _get_redis_client(require_redis=False):
return mock_redis
monkeypatch.setattr(
"src.services.usage.consumer_streams.get_redis_client", _get_redis_client
)
monkeypatch.setattr("src.services.usage.consumer_streams.get_redis_client", _get_redis_client)
with pytest.raises(ResponseError, match="SOME OTHER ERROR"):
await ensure_usage_stream_group()
@@ -598,12 +580,11 @@ async def test_ensure_stream_group_raises_other_errors(monkeypatch):
@pytest.mark.asyncio
async def test_ensure_stream_group_no_redis(monkeypatch):
"""测试 Redis 不可用时直接返回"""
async def _get_redis_client(require_redis=False):
return None
monkeypatch.setattr(
"src.services.usage.consumer_streams.get_redis_client", _get_redis_client
)
monkeypatch.setattr("src.services.usage.consumer_streams.get_redis_client", _get_redis_client)
# 不应抛出异常
await ensure_usage_stream_group()

View File

@@ -7,8 +7,9 @@ UsageService 测试
- 用量统计查询
"""
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from unittest.mock import MagicMock, patch, AsyncMock
from src.services.usage.service import UsageService
@@ -26,7 +27,15 @@ class TestCostCalculation:
output_price_per_1m=15.0,
)
input_cost, output_cost, cache_creation_cost, cache_read_cost, cache_cost, request_cost, total_cost = result
(
input_cost,
output_cost,
cache_creation_cost,
cache_read_cost,
cache_cost,
request_cost,
total_cost,
) = result
# 1000 tokens * $3 / 1M = $0.003
assert abs(input_cost - 0.003) < 0.0001
@@ -268,7 +277,7 @@ class TestHelperMethods:
async def test_get_rate_multiplier_from_provider_api_key(self) -> None:
"""测试从 ProviderAPIKey 获取费率倍数"""
mock_provider_api_key = MagicMock()
mock_provider_api_key.rate_multipliers = {"CLAUDE": 0.8}
mock_provider_api_key.rate_multipliers = {"claude:chat": 0.8}
mock_endpoint = MagicMock()
mock_endpoint.provider_id = "provider-123"
@@ -285,7 +294,7 @@ class TestHelperMethods:
]
rate_multiplier, is_free_tier = await UsageService._get_rate_multiplier_and_free_tier(
mock_db, provider_api_key_id="pak-123", provider_id=None, api_format="CLAUDE"
mock_db, provider_api_key_id="pak-123", provider_id=None, api_format="claude:chat"
)
assert rate_multiplier == 0.8

View File

@@ -1,10 +1,9 @@
from src.core.api_format import APIFormat
from src.core.api_format import (
CORE_REDACT_HEADERS,
HeaderBuilder,
build_upstream_headers,
detect_capabilities,
extract_client_api_key,
build_upstream_headers_for_endpoint,
detect_capabilities_for_endpoint,
extract_client_api_key_for_endpoint,
filter_response_headers,
get_header_value,
normalize_headers,
@@ -36,25 +35,25 @@ class TestGetHeaderValue:
class TestExtractClientApiKey:
def test_bearer_token(self) -> None:
headers = {"authorization": "Bearer test-token"}
assert extract_client_api_key(headers, APIFormat.OPENAI) == "test-token"
assert extract_client_api_key_for_endpoint(headers, "openai:chat") == "test-token"
def test_bearer_token_requires_prefix(self) -> None:
headers = {"Authorization": "test-token"}
assert extract_client_api_key(headers, APIFormat.OPENAI) is None
assert extract_client_api_key_for_endpoint(headers, "openai:chat") is None
def test_header_auth(self) -> None:
headers = {"X-API-Key": "abc"}
assert extract_client_api_key(headers, APIFormat.CLAUDE) == "abc"
assert extract_client_api_key_for_endpoint(headers, "claude:chat") == "abc"
class TestDetectCapabilities:
def test_claude_context_1m(self) -> None:
headers = {"Anthropic-Beta": "context-1m,foo"}
assert detect_capabilities(headers, APIFormat.CLAUDE) == {"context_1m": True}
assert detect_capabilities_for_endpoint(headers, "claude:chat") == {"context_1m": True}
def test_non_claude_noop(self) -> None:
headers = {"Anthropic-Beta": "context-1m"}
assert detect_capabilities(headers, APIFormat.OPENAI) == {}
assert detect_capabilities_for_endpoint(headers, "openai:chat") == {}
class TestHeaderBuilder:
@@ -77,14 +76,14 @@ class TestHeaderBuilder:
class TestBuildUpstreamHeaders:
def test_priority_and_drop_headers(self) -> None:
result = build_upstream_headers(
result = build_upstream_headers_for_endpoint(
{
"Host": "example.com",
"X-Api-Key": "client",
"User-Agent": "ua",
"Content-Type": "text/plain",
},
APIFormat.OPENAI,
"openai:chat",
"provider",
endpoint_headers={"Authorization": "bad", "X-Endpoint": "1", "Content-Type": "bad"},
extra_headers={"User-Agent": "extra", "X-Extra": "1"},
@@ -98,9 +97,9 @@ class TestBuildUpstreamHeaders:
assert result["X-Extra"] == "1"
def test_drop_headers_empty_does_not_fallback(self) -> None:
result = build_upstream_headers(
result = build_upstream_headers_for_endpoint(
{"Host": "example.com"},
APIFormat.OPENAI,
"openai:chat",
"provider",
drop_headers=frozenset(),
)
@@ -108,9 +107,9 @@ class TestBuildUpstreamHeaders:
assert result["Authorization"] == "Bearer provider"
def test_no_duplicate_on_case_variants(self) -> None:
result = build_upstream_headers(
result = build_upstream_headers_for_endpoint(
{"user-agent": "a"},
APIFormat.OPENAI,
"openai:chat",
"provider",
extra_headers={"User-Agent": "b"},
)
@@ -118,7 +117,7 @@ class TestBuildUpstreamHeaders:
assert result["User-Agent"] == "b"
def test_default_content_type(self) -> None:
result = build_upstream_headers({}, APIFormat.OPENAI, "provider")
result = build_upstream_headers_for_endpoint({}, "openai:chat", "provider")
assert result["Content-Type"] == "application/json"
@@ -149,4 +148,3 @@ class TestCapabilityResolverHeaderParsing:
request_headers={"x-require-capability": "context_1m"}
)
assert reqs == {"context_1m": True}