diff --git a/packages/llm_gateway/tests/test_fallback_chain.py b/packages/llm_gateway/tests/test_fallback_chain.py index ca739ff..5a43f69 100644 --- a/packages/llm_gateway/tests/test_fallback_chain.py +++ b/packages/llm_gateway/tests/test_fallback_chain.py @@ -7,7 +7,10 @@ from __future__ import annotations +import uuid + import pytest +import structlog from fakes_resilience import ( AuthError, FakeLedger, @@ -18,9 +21,22 @@ from fakes_resilience import ( transient, ) from ww_llm_gateway.gateway import CircuitBreaker, Gateway -from ww_llm_gateway.types import LlmRequest +from ww_llm_gateway.types import LlmRequest, Scope from ww_shared import AppError, ErrorCode +_SCOPE = Scope(user_id=uuid.UUID(int=1)) + + +def _failing_gw() -> Gateway: + """单 provider 链、每次 transient 失败:耗尽重试后记 `llm_provider_failed` 再抛。""" + primary = ScriptedAdapter("deepseek", failures=[transient() for _ in range(10)]) + return Gateway( + {"deepseek": primary}, + FakeLedger(), + chain_resolver=chain_resolver(chain(("deepseek", "deepseek-chat"))), + max_retries=1, + ) + async def test_primary_success_no_fallback(req: LlmRequest) -> None: primary = ScriptedAdapter("deepseek", text="主模型") @@ -238,3 +254,41 @@ async def test_open_circuit_skips_provider(req: LlmRequest) -> None: assert resp.served_by.provider == "openai" assert resp.served_by.fell_back is True assert primary.complete_calls == 0 + + +# CR-H4 收尾:`llm_provider_failed`(失败/回退日志)也须条件透传 request_id—— +# 设了则带上、未设则不发键(否则 None 覆盖 merge_contextvars 供的 id,回退 sync/SSE 追踪)。 + + +async def test_run_failure_log_carries_request_id_conditionally() -> None: + with structlog.testing.capture_logs() as logs: + with pytest.raises(AppError): + await _failing_gw().run( + LlmRequest(tier="writer", input="x", scope=_SCOPE, request_id="rid-run") + ) + failed = next(e for e in logs if e["event"] == "llm_provider_failed") + assert failed["request_id"] == "rid-run" + + with structlog.testing.capture_logs() as logs: + with pytest.raises(AppError): + await _failing_gw().run(LlmRequest(tier="writer", input="x", scope=_SCOPE)) + failed = next(e for e in logs if e["event"] == "llm_provider_failed") + assert "request_id" not in failed + + +async def test_stream_failure_log_carries_request_id_conditionally() -> None: + with structlog.testing.capture_logs() as logs: + with pytest.raises(AppError): + async for _ in _failing_gw().stream( + LlmRequest(tier="writer", input="x", scope=_SCOPE, request_id="rid-stream") + ): + pass + failed = next(e for e in logs if e["event"] == "llm_provider_failed") + assert failed["request_id"] == "rid-stream" + + with structlog.testing.capture_logs() as logs: + with pytest.raises(AppError): + async for _ in _failing_gw().stream(LlmRequest(tier="writer", input="x", scope=_SCOPE)): + pass + failed = next(e for e in logs if e["event"] == "llm_provider_failed") + assert "request_id" not in failed