mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
refactor(proxy): isolate stream timing prefetch
This commit is contained in:
parent
034d5fe62b
commit
75ebc3f3e4
7 changed files with 32 additions and 18 deletions
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue