diff --git a/libs/partners/fireworks/langchain_fireworks/chat_models.py b/libs/partners/fireworks/langchain_fireworks/chat_models.py index 3b10e157b9..ff38f7e3a7 100644 --- a/libs/partners/fireworks/langchain_fireworks/chat_models.py +++ b/libs/partners/fireworks/langchain_fireworks/chat_models.py @@ -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): diff --git a/libs/partners/fireworks/tests/unit_tests/test_chat_models.py b/libs/partners/fireworks/tests/unit_tests/test_chat_models.py index 2679d550b2..5ac6fefc26 100644 --- a/libs/partners/fireworks/tests/unit_tests/test_chat_models.py +++ b/libs/partners/fireworks/tests/unit_tests/test_chat_models.py @@ -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."""