Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 46 additions & 2 deletions unstract/sdk1/src/unstract/sdk1/llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,10 +7,12 @@
from functools import cache, lru_cache
from typing import Any, NoReturn, cast

import httpx
import litellm

# from litellm import get_supported_openai_params
from litellm import get_max_tokens
from litellm.llms.custom_httpx.http_handler import HTTPHandler
from unstract.sdk1.adapters.constants import Common
from unstract.sdk1.adapters.llm1 import adapters
from unstract.sdk1.constants import Common as SdkCommon
Expand Down Expand Up @@ -95,6 +97,46 @@ def _inject_mock_response(completion_kwargs: dict[str, object]) -> None:
completion_kwargs["mock_response"] = mock


@lru_cache(maxsize=8)
def _gemini_stream_client(timeout: float) -> HTTPHandler:
"""One shared HTTP client per timeout value, so connections are pooled.

Adapters use a handful of timeouts, so a few clients cover them. The cap
keeps a worker that sees many distinct values from holding a pool for
each; an evicted client is not closed here, because a call in flight may
still be using it, and is released once the last reference goes.
"""
return HTTPHandler(timeout=httpx.Timeout(timeout))


def _with_gemini_stream_timeout(
completion_kwargs: dict[str, object],
) -> dict[str, object]:
"""Make a streamed Gemini call honour the adapter's ``timeout``.

LiteLLM's Gemini handler (``gemini/*`` and ``vertex_ai/*gemini*``) drops
``timeout`` on the sync streaming path: it streams on
``litellm.module_level_client``, whose deadline is
``litellm.request_timeout`` (6000 s by default), so a stalled stream could
sit for 100 minutes per attempt. Passing our own client is the only way the
handler applies a timeout there. Other providers forward ``timeout`` to the
request themselves and are left untouched.
"""
model = str(completion_kwargs.get("model", ""))
is_gemini = model.startswith("gemini/") or (
model.startswith("vertex_ai/") and "gemini" in model
)
timeout = completion_kwargs.get("timeout")
if (
not is_gemini
or "client" in completion_kwargs
or not isinstance(timeout, int | float)
or timeout <= 0
):
return completion_kwargs
return {**completion_kwargs, "client": _gemini_stream_client(float(timeout))}


# Drop unsupported params rather than raising errors.
# Set once at module level instead of per-call to avoid repeated
# global mutation in concurrent environments.
Expand Down Expand Up @@ -561,12 +603,13 @@ def _complete_via_stream(
thinking blocks, ``finish_reason``, usage (including cache tokens)
and the provider response headers.
"""
stream_kwargs = _with_gemini_stream_timeout(completion_kwargs)
chunks = collect_with_retry(
lambda: litellm.completion(
messages=messages,
stream=True,
stream_options={"include_usage": True},
**completion_kwargs,
**stream_kwargs,
),
max_retries=max_retries,
retry_predicate=is_retryable_litellm_error,
Expand Down Expand Up @@ -829,13 +872,14 @@ def stream_complete(
max_retries = pop_litellm_retry_kwargs(
completion_kwargs, self._get_adapter_info()
)
stream_kwargs = _with_gemini_stream_timeout(completion_kwargs)
has_yielded_content = False
for chunk in iter_with_retry(
lambda: litellm.completion(
messages=messages,
stream=True,
stream_options={"include_usage": True},
**completion_kwargs,
**stream_kwargs,
),
max_retries=max_retries,
retry_predicate=is_retryable_litellm_error,
Expand Down
214 changes: 214 additions & 0 deletions unstract/sdk1/tests/test_gemini_stream_timeout.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,214 @@
"""Streamed Gemini calls honour the adapter's ``timeout``.

LiteLLM's Gemini handler (``gemini/*`` and ``vertex_ai/*gemini*``) drops the
per-call ``timeout`` on the sync streaming path and streams on
``litellm.module_level_client``, whose deadline is ``litellm.request_timeout``
(6000 s by default). In production a hung ``vertex_ai/gemini-3.1-flash-lite``
stream was therefore bounded by 100 minutes per attempt, not the adapter's
600 s (UN-4223). ``LLM`` passes its own client for these models so the
adapter's timeout reaches the HTTP request.
"""

from __future__ import annotations

import json
from functools import lru_cache
from importlib import import_module
from typing import TYPE_CHECKING
from unittest.mock import patch

import httpx
import litellm
import pytest
from litellm.llms.custom_httpx.http_handler import HTTPHandler
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase

if TYPE_CHECKING:
from collections.abc import Iterator

GEMINI_ADAPTER_ID = "gemini|085f6c03-b57e-4594-85bb-40e2616c2736"
VERTEX_ADAPTER_ID = "vertexai|78fa17a5-a619-47d4-ac6e-3fc1698fdb55"
_REAL_COMPLETION = litellm.completion


@lru_cache(maxsize=1)
def _load_llm_module() -> object:
import sys
from types import ModuleType

sys.modules.setdefault("magic", ModuleType("magic"))
return import_module("unstract.sdk1.llm")


@pytest.fixture
def no_cost() -> Iterator[None]:
llm_module = _load_llm_module()
with patch.object(llm_module.litellm, "cost_per_token", return_value=(0.0, 0.0)):
yield


# ── Which calls get a client ─────────────────────────────────────────────────


@pytest.mark.parametrize(
"model", ["vertex_ai/gemini-3.1-flash-lite", "gemini/gemini-2.5-flash"]
)
def test_gemini_models_get_a_client_with_the_adapter_timeout(model: str) -> None:
llm_module = _load_llm_module()
kwargs = {"model": model, "timeout": 600}

result = llm_module._with_gemini_stream_timeout(kwargs)

client = result["client"]
assert isinstance(client, HTTPHandler)
assert client.client.timeout.read == 600
assert "client" not in kwargs # the caller's kwargs are not mutated


@pytest.mark.parametrize(
"model",
[
"anthropic/claude-sonnet-4-6",
# A Vertex partner model takes LiteLLM's partner route, which forwards
# ``timeout`` itself.
"vertex_ai/claude-sonnet-4-6",
"openai/gpt-4o",
],
)
def test_other_models_are_left_untouched(model: str) -> None:
llm_module = _load_llm_module()
kwargs = {"model": model, "timeout": 600}

assert llm_module._with_gemini_stream_timeout(kwargs) is kwargs


@pytest.mark.parametrize("timeout", [None, 0])
def test_no_client_without_a_usable_timeout(timeout: object) -> None:
llm_module = _load_llm_module()
kwargs = {"model": "gemini/gemini-2.5-flash", "timeout": timeout}

assert "client" not in llm_module._with_gemini_stream_timeout(kwargs)


def test_caller_supplied_client_is_kept() -> None:
llm_module = _load_llm_module()
own = HTTPHandler(timeout=30)
kwargs = {"model": "gemini/gemini-2.5-flash", "timeout": 600, "client": own}

assert llm_module._with_gemini_stream_timeout(kwargs)["client"] is own


def test_client_is_shared_per_timeout() -> None:
llm_module = _load_llm_module()
first = llm_module._with_gemini_stream_timeout(
{"model": "gemini/gemini-2.5-flash", "timeout": 600}
)
second = llm_module._with_gemini_stream_timeout(
{"model": "vertex_ai/gemini-2.5-pro", "timeout": 600.0}
)

assert first["client"] is second["client"]


# ── The timeout reaches the HTTP request ─────────────────────────────────────


def _gemini_sse_body(text: str) -> bytes:
event = {
"candidates": [
{
"content": {"parts": [{"text": text}], "role": "model"},
"finishReason": "STOP",
"index": 0,
}
],
"usageMetadata": {
"promptTokenCount": 1,
"candidatesTokenCount": 1,
"totalTokenCount": 2,
},
}
return f"data: {json.dumps(event)}\n\n".encode()


@pytest.fixture
def sent() -> Iterator[list[httpx.Request]]:
"""Capture outgoing HTTP requests and answer each with a Gemini stream."""
requests: list[httpx.Request] = []

def handle_request(
_transport: httpx.HTTPTransport, request: httpx.Request
) -> httpx.Response:
requests.append(request)
return httpx.Response(
200,
headers={"content-type": "text/event-stream"},
content=_gemini_sse_body("hello"),
)

# Vertex needs an OAuth token; skip the Google auth round trip.
with (
patch.object(httpx.HTTPTransport, "handle_request", handle_request),
patch.object(
VertexBase, "_ensure_access_token", return_value=("token", "test-project")
),
):
yield requests


def _gemini_llm(timeout: int) -> object:
return _load_llm_module().LLM(
adapter_id=GEMINI_ADAPTER_ID,
adapter_metadata={
"model": "gemini-2.5-flash",
"api_key": "test-key",
"timeout": timeout,
},
)


def _vertex_llm(timeout: int) -> object:
return _load_llm_module().LLM(
adapter_id=VERTEX_ADAPTER_ID,
adapter_metadata={
"model": "gemini-2.5-flash",
"json_credentials": "{}",
"project": "test-project",
"timeout": timeout,
},
)


# Without the client LLM passes, each of these requests goes out with
# LiteLLM's 6000 s ``request_timeout`` instead of the adapter's.


def test_gemini_complete_request_carries_the_adapter_timeout(
no_cost: None, sent: list[httpx.Request]
) -> None:
result = _gemini_llm(timeout=420).complete("hi")

assert result["response"].text == "hello"
assert len(sent) == 1
assert sent[0].extensions["timeout"]["read"] == 420


def test_vertex_complete_request_carries_the_adapter_timeout(
no_cost: None, sent: list[httpx.Request]
) -> None:
result = _vertex_llm(timeout=300).complete("hi")

assert result["response"].text == "hello"
assert len(sent) == 1
assert "aiplatform.googleapis.com" in str(sent[0].url)
assert sent[0].extensions["timeout"]["read"] == 300


def test_gemini_stream_complete_request_carries_the_adapter_timeout(
no_cost: None, sent: list[httpx.Request]
) -> None:
text = "".join(r.text for r in _gemini_llm(timeout=240).stream_complete("hi"))

assert text == "hello"
assert len(sent) == 1
assert sent[0].extensions["timeout"]["read"] == 240
Loading