fix(fireworks): classify mid-stream read timeouts (#40874)

`ChatFireworks` now reports mid-stream read timeouts as retryable
`ModelTimeoutError`s while remaining catchable as `httpx.ReadTimeout`.
Partial streams are not automatically replayed.

---

Fireworks stream read timeouts currently bypass LangChain's model-error
classification after the first chunk. Classify them as retryable
`ModelTimeoutError`s while preserving `httpx.ReadTimeout` compatibility,
the original cause, and request context when available.

- Covers synchronous and asynchronous stream consumption.
- Keeps setup retries unchanged. Does **not** replay partial streams:
doing so could duplicate text or corrupt tool-call arguments.
Whole-generation recovery remains the caller's responsibility.
- Built-in streaming retry and fallback behavior is unchanged:
`.with_retry()` does not retry streaming, and `.with_fallbacks()` only
switches models before the first chunk. Applications must explicitly
handle recovery after output has started.

Made by [Open SWE](https://github.com/langchain-ai/open-swe) · [view
thread](https://openswe.vercel.app/agents/c2ba7088-1448-5e4d-9541-ae6f25b513c7)
· openai:gpt-6-astra (medium)

Co-authored-by: Mason Daugherty <mdrxy@users.noreply.github.com>
Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
authored and GitHub committed 2026-09-28 14:18:14 -04:00
1 parent 42f04f460d
commit 60e57f5bc5
2 files changed
+127 -10

No files matched your search

@@ -653,6 +653,18 @@ class FireworksTimeoutError(APITimeoutError, ModelTimeoutError):
"""Fireworks timeout error classified as a LangChain model error."""
class FireworksReadTimeoutError(httpx.ReadTimeout, ModelTimeoutError):
"""Fireworks stream read timeout classified as a retryable model error."""
def _handle_stream_read_timeout(error: httpx.ReadTimeout) -> NoReturn:
try:
request = error.request
except RuntimeError:
request = None
raise FireworksReadTimeoutError(str(error), request=request) from error
def _handle_fireworks_invalid_request(e: BadRequestError) -> NoReturn:
"""Promote prompt-too-long errors to `FireworksContextOverflowError`."""
if "prompt is too long" in str(e):
@@ -815,13 +827,19 @@ async def _acompletion_with_retry(
def _prepend_chunk(first: Any, rest: Iterator[Any]) -> Iterator[Any]:
yield first
yield from rest
try:
yield from rest
except httpx.ReadTimeout as e:
_handle_stream_read_timeout(e)
async def _aprepend_chunk(first: Any, rest: AsyncIterator[Any]) -> AsyncIterator[Any]:
yield first
async for item in rest:
yield item
try:
async for item in rest:
yield item
except httpx.ReadTimeout as e:
_handle_stream_read_timeout(e)
class ChatFireworks(BaseChatModel):
@@ -5,8 +5,9 @@ from __future__ import annotations
import json
import logging
import os
from collections.abc import AsyncIterator, Iterator
from typing import Any
from unittest.mock import MagicMock
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
@@ -841,7 +842,11 @@ def test_completion_with_retry_exhausts_and_raises() -> None:
assert mock_client.create.call_count == 3
def test_completion_with_retry_streaming_retries_on_setup() -> None:
@pytest.mark.parametrize(
"error",
[_api_error(RateLimitError, "rate limited", 429), httpx.ReadTimeout("slow")],
)
def test_completion_with_retry_streaming_retries_on_setup(error: Exception) -> None:
"""Streaming errors raised during the first-chunk pull are retried."""
llm = _make_llm(max_retries=1)
@@ -852,8 +857,7 @@ def test_completion_with_retry_streaming_retries_on_setup() -> None:
if calls["n"] == 1:
def _failing_gen() -> Any:
msg = "rate limited"
raise _api_error(RateLimitError, msg, 429)
raise error
yield # pragma: no cover
return _failing_gen()
@@ -994,7 +998,13 @@ def test_chat_fireworks_invoke_routes_through_retry() -> None:
assert mock_client.create.call_count == 2
async def test_acompletion_with_retry_streaming_retries_on_setup() -> None:
@pytest.mark.parametrize(
"error",
[_api_error(RateLimitError, "rate limited", 429), httpx.ReadTimeout("slow")],
)
async def test_acompletion_with_retry_streaming_retries_on_setup(
error: Exception,
) -> None:
"""Async streaming errors during the first-chunk pull are retried."""
llm = _make_llm(max_retries=1)
calls = {"n": 0}
@@ -1004,8 +1014,7 @@ async def test_acompletion_with_retry_streaming_retries_on_setup() -> None:
if calls["n"] == 1:
async def _failing_agen() -> Any:
msg = "rate limited"
raise _api_error(RateLimitError, msg, 429)
raise error
yield # pragma: no cover
return _failing_agen()
@@ -1546,6 +1555,96 @@ class TestExtraHeaders:
assert "extra_headers" not in call_kwargs.get("extra_body", {})
@pytest.mark.parametrize("with_request", [False, True])
@pytest.mark.parametrize("async_mode", [False, True])
@pytest.mark.parametrize("tool_call", [False, True])
async def test_midstream_read_timeout_is_classified_without_replay(
*, with_request: bool, async_mode: bool, tool_call: bool
) -> None:
request = httpx.Request("POST", "https://api.fireworks.ai/inference/v1")
error = httpx.ReadTimeout(
"Timeout on reading data from socket",
request=request if with_request else None,
)
delta: dict[str, object] = (
{
"tool_calls": [
{
"index": 0,
"id": "call_1",
"type": "function",
"function": {"name": "lookup", "arguments": '{"query":'},
}
]
}
if tool_call
else {"content": "Hello"}
)
chunk: dict[str, object] = {"choices": [{"delta": delta, "index": 0}]}
def _stream() -> Iterator[dict[str, object]]:
yield chunk
raise error
async def _astream() -> AsyncIterator[dict[str, object]]:
yield chunk
raise error
model = _make_llm(max_retries=2)
received: list[AIMessageChunk] = []
with pytest.raises(httpx.ReadTimeout) as exc_info:
if async_mode:
model.async_client = MagicMock(create=AsyncMock(return_value=_astream()))
async for message in model.astream("Hello"):
received.append(message) # noqa: PERF401
else:
model.client = MagicMock(create=MagicMock(return_value=_stream()))
received.extend(model.stream("Hello"))
classified = exc_info.value
assert isinstance(classified, ModelTimeoutError)
assert classified.is_retryable
assert classified.__cause__ is error
assert str(classified) == str(error)
if with_request:
assert classified.request is request
else:
with pytest.raises(RuntimeError, match="request"):
_ = classified.request
assert len(received) == 1
if tool_call:
assert received[0].tool_call_chunks[0]["args"] == '{"query":'
else:
assert received[0].content == "Hello"
client = model.async_client if async_mode else model.client
client.create.assert_called_once()
@pytest.mark.parametrize("async_mode", [False, True])
async def test_midstream_non_timeout_propagates_unchanged(*, async_mode: bool) -> None:
error = ValueError("invalid stream")
def _stream() -> Iterator[dict[str, object]]:
yield {"choices": [{"delta": {"content": "Hello"}}]}
raise error
async def _astream() -> AsyncIterator[dict[str, object]]:
for chunk in _stream():
yield chunk
model = _make_llm(max_retries=2)
with pytest.raises(ValueError) as exc_info:
if async_mode:
model.async_client = MagicMock(create=AsyncMock(return_value=_astream()))
_ = [chunk async for chunk in model.astream("Hello")]
else:
model.client = MagicMock(create=MagicMock(return_value=_stream()))
list(model.stream("Hello"))
assert exc_info.value is error
client = model.async_client if async_mode else model.client
client.create.assert_called_once()
class TestStreamUsage:
"""Tests for the `stream_usage` field and `stream_options` plumbing."""