mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
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:
@@ -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
|
||||
|
||||
@@ -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"})
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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""
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
210
tests/services/test_async_task_orchestrator.py
Normal file
210
tests/services/test_async_task_orchestrator.py
Normal 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
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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],验证排序仍然是 CLAUDE(exact)在前
|
||||
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"
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user