"""LLM 错误分类:httpx 超时族必须可重试(瞬时网络故障不能直接终结 run)。 线上实测:兼容模式流式响应中途 200 OK 后停发数据 ~17s,抛 ``httpx.ReadTimeout``。 旧名单只认 ``ReadError`` / ``RemoteProtocolError``,ReadTimeout 被归为 generic (不可重试),1 次尝试即兜底报错、整个 run 终止。 """ from __future__ import annotations import asyncio from types import SimpleNamespace import httpx from langchain_core.messages import AIMessage from deerflow.agents.middlewares.llm_error_handling_middleware import LLM_ERROR_MARKER, LLMErrorHandlingMiddleware def _middleware() -> LLMErrorHandlingMiddleware: config = SimpleNamespace( circuit_breaker=SimpleNamespace(failure_threshold=3, recovery_timeout_sec=30), ) middleware = LLMErrorHandlingMiddleware(app_config=config) middleware.retry_base_delay_ms = 0 middleware.retry_cap_delay_ms = 0 return middleware def test_httpx_read_timeout_is_transient() -> None: assert _middleware()._classify_error(httpx.ReadTimeout("read timed out")) == (True, "transient") def test_httpx_timeout_family_is_transient() -> None: middleware = _middleware() for exc in ( httpx.ConnectTimeout("connect timed out"), httpx.WriteTimeout("write timed out"), httpx.PoolTimeout("pool timed out"), httpx.ConnectError("connection refused"), TimeoutError(), ): assert middleware._classify_error(exc) == (True, "transient"), exc def test_value_error_stays_generic() -> None: assert _middleware()._classify_error(ValueError("boom")) == (False, "generic") def test_read_timeout_is_retried_and_recovers() -> None: """第一次 ReadTimeout、第二次成功:必须重试并返回真实结果。""" middleware = _middleware() calls = {"count": 0} async def handler(_request): calls["count"] += 1 if calls["count"] == 1: raise httpx.ReadTimeout("read timed out") return AIMessage(content="ok") result = asyncio.run(middleware.awrap_model_call(SimpleNamespace(), handler)) assert calls["count"] == 2 assert result.content == "ok" def test_persistent_read_timeout_falls_back_with_transient_reason() -> None: """连续超时到重试上限:兜底消息 reason=transient(而非 generic)。""" middleware = _middleware() async def handler(_request): raise httpx.ReadTimeout("read timed out") result = asyncio.run(middleware.awrap_model_call(SimpleNamespace(), handler)) assert isinstance(result, AIMessage) assert result.additional_kwargs[LLM_ERROR_MARKER] == "transient"