diff --git a/packages/llm_gateway/tests/test_fallback_chain.py b/packages/llm_gateway/tests/test_fallback_chain.py index 9e9e1ee..ca739ff 100644 --- a/packages/llm_gateway/tests/test_fallback_chain.py +++ b/packages/llm_gateway/tests/test_fallback_chain.py @@ -132,6 +132,31 @@ async def test_missing_adapter_in_chain_skipped(req: LlmRequest) -> None: assert resp.served_by.fell_back is True +async def test_stream_retries_before_first_token_then_succeeds(req: LlmRequest) -> None: + # CR-H6:首 token 前的瞬时失败应 per-provider 退避重试(对称 run 的 _complete_with_retry), + # 而非直接烧掉一个回退名额切备用。max_retries=2 → 最多 3 次尝试;前两次 transient、 + # 第三次成功 → 主 provider 服务,备用零调用。 + primary = ScriptedAdapter( + "deepseek", text="流式重试后成功", failures=[transient(), transient()] + ) + backup = ScriptedAdapter("openai", text="备用") + ledger = FakeLedger() + gw = Gateway( + {"deepseek": primary, "openai": backup}, + ledger, + chain_resolver=chain_resolver(chain(("deepseek", "deepseek-chat"), ("openai", "gpt-4o"))), + max_retries=2, + ) + + collected = [d.text async for d in gw.stream(req)] + + assert "".join(collected) == "流式重试后成功" + assert primary.stream_calls == 3 # 2 次失败 + 1 次成功(重试在首 token 前) + assert backup.stream_calls == 0 # 未切备用 + assert len(ledger.records) == 1 + assert ledger.records[0].provider == "deepseek" + + # ---- 熔断器 ---- diff --git a/packages/llm_gateway/ww_llm_gateway/gateway.py b/packages/llm_gateway/ww_llm_gateway/gateway.py index babc123..11270b0 100644 --- a/packages/llm_gateway/ww_llm_gateway/gateway.py +++ b/packages/llm_gateway/ww_llm_gateway/gateway.py @@ -28,7 +28,7 @@ from tenacity import ( ) from ww_shared import AppError, ErrorCode -from .adapters.base import ProviderAdapter, ProviderResult, ProviderUsage +from .adapters.base import ProviderAdapter, ProviderResult, ProviderUsage, StreamChunk from .errors import TransientProviderError from .ledger import LedgerSink from .pricing import cost_minor @@ -268,6 +268,39 @@ class Gateway: return await adapter.complete(req, model) raise AssertionError("unreachable: reraise=True") + async def _stream_with_retry( + self, adapter: ProviderAdapter, req: LlmRequest, model: str + ) -> AsyncIterator[StreamChunk]: + """流式版 per-provider 重试:**仅**覆盖到首个含文本的 chunk(§4.5)。 + + 对称 `_complete_with_retry`——首 token 前的瞬时失败退避重试;一旦吐出文本即 + 提交该 candidate,其后失败属中途失败(由 `stream()` 的 `started` 门直接上抛, + 不静默重连)。空流/纯 usage 提前结束(StopAsyncIteration)视为正常完成,不重试。 + record_failure 归外层循环(此助手不碰熔断,同 `_complete_with_retry`)。 + """ + retrying = self._retrying() + gen: AsyncIterator[StreamChunk] | None = None + prelude: list[StreamChunk] = [] + async for attempt in retrying: + with attempt: + candidate = adapter.stream(req, model) + buffered: list[StreamChunk] = [] + got_text = False + try: + while not got_text: + chunk = await candidate.__anext__() + buffered.append(chunk) + got_text = bool(chunk.text) + except StopAsyncIteration: + pass + gen = candidate + prelude = buffered + for chunk in prelude: + yield chunk + if gen is not None: + async for chunk in gen: + yield chunk + # ---- stream(流式)---- async def stream(self, req: LlmRequest) -> AsyncIterator[Delta]: @@ -282,7 +315,7 @@ class Gateway: # 流式回退:仅在「尚未吐出任何 token」前可切。首块产出后失败属中途失败, # 不静默重连(§4.5:已存部分留 draft,节点报错停在 write 前 checkpoint)。 try: - stream_iter = adapter.stream(req, route.model) + stream_iter = self._stream_with_retry(adapter, req, route.model) final = ProviderUsage(input_tokens=0, output_tokens=0) started = False async for chunk in stream_iter: