refactor(proxy): isolate stream timing prefetch

This commit is contained in:
Yucheng Zhu 2026-08-29 13:40:49 -07:00
parent 034d5fe62b
commit 75ebc3f3e4
7 changed files with 32 additions and 18 deletions

View file

@ -74,7 +74,7 @@ async def prefetch_proxy_stream_for_timing(completion_response: object) -> None:
not is_proxy_stream_header_prefetch.get()
or is_internal_call.get()
or not isinstance(completion_response, CustomStreamWrapper)
or completion_response.custom_llm_provider != "vertex_ai_beta"
or not completion_response.prefetch_for_proxy_stream_headers
or completion_response.completion_stream is not None
or completion_response.make_call is None
):

View file

@ -235,9 +235,11 @@ class CustomStreamWrapper:
stream_options=None,
make_call: Callable | None = None,
_response_headers: dict | httpx.Headers | None = None,
prefetch_for_proxy_stream_headers: bool = False,
):
self.model = model
self.make_call = make_call
self.prefetch_for_proxy_stream_headers = prefetch_for_proxy_stream_headers
self.custom_llm_provider = custom_llm_provider
self.logging_obj: LiteLLMLoggingObject = logging_obj
self.completion_stream = completion_stream

View file

@ -2744,6 +2744,7 @@ class VertexLLM(VertexBase):
model=model,
custom_llm_provider="vertex_ai_beta",
logging_obj=logging_obj,
prefetch_for_proxy_stream_headers=True,
)
return streaming_response
@ -3026,6 +3027,7 @@ class VertexLLM(VertexBase):
model=model,
custom_llm_provider="vertex_ai_beta",
logging_obj=logging_obj,
prefetch_for_proxy_stream_headers=True,
)
return streaming_response

View file

@ -410,6 +410,17 @@ async def _cancel_pending_gather_tasks(tasks: list["asyncio.Task[Any]"]) -> None
pass
async def _close_unowned_stream(response: object) -> None:
aclose: Final = getattr(response, "aclose", None)
if aclose is None:
return
with anyio.CancelScope(shield=True):
try:
await aclose()
except Exception as exc: # noqa: BLE001
verbose_proxy_logger.debug("Error closing unowned response stream: %s", exc)
@lru_cache(maxsize=512)
def _litellm_model_supports_stream_options(litellm_model: str) -> bool:
try:
@ -2320,8 +2331,6 @@ class ProxyBaseLLMRequestProcessing:
await _cancel_pending_gather_tasks(tasks)
response = responses[1]
response_ownership_transferred = False
response_requires_cleanup = False
_exception_raised = False
try:
@ -2350,7 +2359,6 @@ class ProxyBaseLLMRequestProcessing:
if self._is_streaming_request(
data=self.data, is_streaming_request=is_streaming_request
) or self._is_streaming_response(response): # use generate_responses to stream responses
response_requires_cleanup = self._is_streaming_response(response)
custom_headers: Final = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=user_api_key_dict,
call_id=logging_obj.litellm_call_id,
@ -2434,7 +2442,6 @@ class ProxyBaseLLMRequestProcessing:
if route_type == "allm_passthrough_route":
# Check if response is an async generator
if self._is_streaming_response(response):
response_ownership_transferred = True
if asyncio.iscoroutine(response):
generator = await response
else:
@ -2505,7 +2512,6 @@ class ProxyBaseLLMRequestProcessing:
headers=custom_headers,
request=request,
)
response_ownership_transferred = True
return anthropic_stream_response
# Non-streaming response - fall through to normal response handling
elif select_data_generator:
@ -2538,7 +2544,6 @@ class ProxyBaseLLMRequestProcessing:
headers=custom_headers,
request=request,
)
response_ownership_transferred = True
return responses_stream_response
### CALL HOOKS ### - modify outgoing data
@ -2580,18 +2585,16 @@ class ProxyBaseLLMRequestProcessing:
user_api_key_dict=user_api_key_dict,
response=response,
)
except asyncio.CancelledError:
if self._is_streaming_response(response):
await _close_unowned_stream(response)
raise
except Exception:
_exception_raised = True
if self._is_streaming_response(response):
await _close_unowned_stream(response)
raise
finally:
if response_requires_cleanup and not response_ownership_transferred:
aclose = getattr(response, "aclose", None)
if aclose is not None:
with anyio.CancelScope(shield=True):
try:
await aclose()
except Exception as exc: # noqa: BLE001
verbose_proxy_logger.debug("Error closing unowned response stream: %s", exc)
ProxyBaseLLMRequestProcessing._flush_deferred_async_logging(
logging_obj=logging_obj,
exception_raised=_exception_raised,

View file

@ -281,6 +281,7 @@ class TestPrefetchedStreamTiming:
logging_obj=logging_obj,
custom_llm_provider="vertex_ai_beta",
make_call=AsyncMock(return_value=empty_stream()),
prefetch_for_proxy_stream_headers=True,
)
token = is_proxy_stream_header_prefetch.set(True)
@ -293,14 +294,14 @@ class TestPrefetchedStreamTiming:
assert dict(logging_obj.set_response_timing_metrics.call_args.args[0])["litellm_overhead_time_ms"] > 0
@pytest.mark.asyncio
async def test_prefetch_does_not_open_non_vertex_stream(self):
async def test_prefetch_does_not_open_stream_without_capability(self):
logging_obj = _messages_logging_obj()
logging_obj.start_time = datetime.datetime.now() - datetime.timedelta(milliseconds=1000)
stream = litellm.CustomStreamWrapper(
completion_stream=None,
model="other-model",
model="gemini-3.5-flash",
logging_obj=logging_obj,
custom_llm_provider="openai",
custom_llm_provider="vertex_ai_beta",
make_call=AsyncMock(return_value=object()),
)

View file

@ -74,6 +74,7 @@ def _deferred_stream(provider: str = "vertex_ai_beta") -> litellm.CustomStreamWr
logging_obj=logging_obj,
custom_llm_provider=provider,
make_call=AsyncMock(return_value=_EmptyAsyncStream()),
prefetch_for_proxy_stream_headers=provider == "vertex_ai_beta",
)

View file

@ -96,6 +96,7 @@ async def test_async_responses_bridge_keeps_sdk_deferred_gemini_stream_lazy():
model="gemini-3.5-flash",
logging_obj=logging_obj,
custom_llm_provider="vertex_ai_beta",
prefetch_for_proxy_stream_headers=True,
make_call=AsyncMock(return_value=_EmptyAsyncStream()),
)
@ -123,6 +124,7 @@ async def test_async_responses_bridge_prefetches_deferred_gemini_stream_for_prox
model="gemini-3.5-flash",
logging_obj=logging_obj,
custom_llm_provider="vertex_ai_beta",
prefetch_for_proxy_stream_headers=True,
make_call=AsyncMock(return_value=_EmptyAsyncStream()),
)
token = is_proxy_stream_header_prefetch.set(True)
@ -153,6 +155,7 @@ async def test_async_responses_bridge_does_not_prefetch_already_connected_gemini
model="gemini-3.5-flash",
logging_obj=logging_obj,
custom_llm_provider="vertex_ai_beta",
prefetch_for_proxy_stream_headers=True,
make_call=AsyncMock(return_value=_EmptyAsyncStream()),
)
token = is_proxy_stream_header_prefetch.set(True)
@ -182,6 +185,7 @@ async def test_async_responses_bridge_closes_prefetched_gemini_stream_on_early_c
model="gemini-3.5-flash",
logging_obj=logging_obj,
custom_llm_provider="vertex_ai_beta",
prefetch_for_proxy_stream_headers=True,
make_call=AsyncMock(return_value=upstream),
)
token = is_proxy_stream_header_prefetch.set(True)
@ -217,6 +221,7 @@ async def test_async_responses_bridge_propagates_initial_fetch_failure():
model="gemini-3.5-flash",
logging_obj=logging_obj,
custom_llm_provider="vertex_ai_beta",
prefetch_for_proxy_stream_headers=True,
make_call=AsyncMock(side_effect=expected_error),
)
token = is_proxy_stream_header_prefetch.set(True)