111 lines
3.8 KiB
Python
111 lines
3.8 KiB
Python
"""T1.1 网关核心单测:run / stream / 记账 / 路由 / 错误。"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import uuid
|
||
|
||
import pytest
|
||
import structlog
|
||
from fakes import FakeAdapter, FakeLedger, fake_route
|
||
from ww_llm_gateway.gateway import Gateway
|
||
from ww_llm_gateway.routing import resolve_route
|
||
from ww_llm_gateway.types import LlmRequest, Scope
|
||
from ww_shared import AppError, ErrorCode
|
||
|
||
|
||
async def test_run_returns_text_usage_and_served_by(req: LlmRequest) -> None:
|
||
adapter = FakeAdapter(text="第一章正文")
|
||
ledger = FakeLedger()
|
||
gw = Gateway({"deepseek": adapter}, ledger, resolver=fake_route)
|
||
|
||
resp = await gw.run(req)
|
||
|
||
assert resp.text == "第一章正文"
|
||
assert resp.served_by.provider == "deepseek"
|
||
assert resp.served_by.fell_back is False
|
||
assert resp.usage.input_tokens == 100
|
||
assert resp.usage.output_tokens == 50
|
||
assert adapter.complete_calls == 1
|
||
|
||
|
||
async def test_run_writes_exactly_one_ledger_record(req: LlmRequest) -> None:
|
||
ledger = FakeLedger()
|
||
gw = Gateway({"deepseek": FakeAdapter()}, ledger, resolver=fake_route)
|
||
|
||
await gw.run(req)
|
||
|
||
assert len(ledger.records) == 1
|
||
assert ledger.records[0].provider == "deepseek"
|
||
assert ledger.records[0].model == "deepseek-chat"
|
||
|
||
|
||
async def test_run_computes_cost_from_pricing(req: LlmRequest) -> None:
|
||
ledger = FakeLedger()
|
||
gw = Gateway({"deepseek": FakeAdapter()}, ledger, resolver=fake_route)
|
||
|
||
resp = await gw.run(req)
|
||
|
||
# ceil(100/1e6*27 + 50/1e6*110) == 1 cent
|
||
assert resp.usage.cost_minor == 1
|
||
assert resp.usage.currency == "USD"
|
||
|
||
|
||
async def test_stream_yields_deltas_and_records_once(req: LlmRequest) -> None:
|
||
ledger = FakeLedger()
|
||
gw = Gateway({"deepseek": FakeAdapter(deltas=["a", "b", "c"])}, ledger, resolver=fake_route)
|
||
|
||
collected = [d.text async for d in gw.stream(req)]
|
||
|
||
assert "".join(collected) == "abc"
|
||
assert len(ledger.records) == 1
|
||
assert ledger.records[0].output_tokens == 50
|
||
|
||
|
||
async def test_run_logs_request_id_from_request(scope: Scope) -> None:
|
||
# CR-H4:调用方设了 request_id 时,网关 llm_call 日志必须带上它(贯通 §9.3 追踪)。
|
||
ledger = FakeLedger()
|
||
gw = Gateway({"deepseek": FakeAdapter()}, ledger, resolver=fake_route)
|
||
req = LlmRequest(tier="writer", input="x", scope=scope, request_id="rid-xyz")
|
||
|
||
with structlog.testing.capture_logs() as logs:
|
||
await gw.run(req)
|
||
|
||
call = next(entry for entry in logs if entry["event"] == "llm_call")
|
||
assert call["request_id"] == "rid-xyz"
|
||
|
||
|
||
async def test_run_omits_request_id_when_unset(scope: Scope) -> None:
|
||
# 未设 request_id 时**不得**发出该键——否则 None 会覆盖 merge_contextvars 供的 id,
|
||
# 反倒回退了 sync/SSE 路径的追踪(条件透传锁定)。
|
||
ledger = FakeLedger()
|
||
gw = Gateway({"deepseek": FakeAdapter()}, ledger, resolver=fake_route)
|
||
req = LlmRequest(tier="writer", input="x", scope=scope)
|
||
|
||
with structlog.testing.capture_logs() as logs:
|
||
await gw.run(req)
|
||
|
||
call = next(entry for entry in logs if entry["event"] == "llm_call")
|
||
assert "request_id" not in call
|
||
|
||
|
||
async def test_unknown_provider_raises_llm_unavailable(req: LlmRequest) -> None:
|
||
gw = Gateway({}, FakeLedger(), resolver=fake_route)
|
||
|
||
with pytest.raises(AppError) as exc:
|
||
await gw.run(req)
|
||
|
||
assert exc.value.code == ErrorCode.LLM_UNAVAILABLE
|
||
|
||
|
||
def test_resolve_route_parses_tier_defaults() -> None:
|
||
route = resolve_route("writer")
|
||
assert route.provider == "deepseek"
|
||
assert route.model == "deepseek-chat"
|
||
|
||
|
||
def test_agent_passes_only_tier_never_model() -> None:
|
||
# 不变量 ②:LlmRequest 无 model 字段,只有 tier
|
||
req = LlmRequest(tier="analyst", input="x", scope=Scope(user_id=uuid.UUID(int=9)))
|
||
assert "model" not in LlmRequest.model_fields
|
||
assert req.tier == "analyst"
|