mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 09:25:14 +03:00
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:
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."""
|
||||
|
||||
|
||||
Reference in new issue
Block a user