This commit is contained in:
Charan Rathore 2026-09-30 10:30:01 -04:00 • committed by GitHub
commit ca07687140
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 230 additions and 6 deletions

View file

@ -3456,9 +3456,13 @@ class ProxyBaseLLMRequestProcessing:
unwrapped inner iterator.
"""
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.proxy.route_llm_request import _RouterOwnedStream
from litellm.router_utils.add_retry_fallback_headers import HiddenParamsAsyncIteratorWrapper
unwrapped: Final = response._inner if isinstance(response, HiddenParamsAsyncIteratorWrapper) else response
owned_stream: Final = response._stream if isinstance(response, _RouterOwnedStream) else response
unwrapped: Final = (
owned_stream._inner if isinstance(owned_stream, HiddenParamsAsyncIteratorWrapper) else owned_stream
)
if isinstance(unwrapped, CustomStreamWrapper):
# Intentionally a live reference (not a copy) — mirrors

View file

@ -1,6 +1,7 @@
import asyncio
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final, Literal
from collections.abc import AsyncIterator, Mapping
from contextlib import suppress
from typing import TYPE_CHECKING, Any, Final, Literal, cast
import httpx
from fastapi import HTTPException, status
@ -35,6 +36,46 @@ else:
LitellmRouter = Any
class _RouterOwnedStream:
"""Keep per-request router callbacks until the stream ends or is closed."""
def __init__(self, stream: AsyncIterator[object], router: LitellmRouter) -> None:
self._stream = stream
self._router = router
self._discarded = False
def __getattr__(self, name: str) -> object:
return getattr(self._stream, name)
def __aiter__(self) -> "_RouterOwnedStream":
return self
def _discard(self) -> None:
if not self._discarded:
self._discarded = True
self._router.discard()
async def __anext__(self) -> object:
try:
return await anext(self._stream)
except BaseException:
# A failed or cancelled stream still owns an upstream connection.
# Preserve the original exception if closing also fails.
with suppress(BaseException):
await self.aclose()
raise
async def aclose(self) -> None:
if self._discarded:
return
try:
closer = getattr(self._stream, "aclose", None)
if closer is not None:
await closer()
finally:
self._discard()
def _route_user_config_request(data: dict, route_type: str):
"""Route a request using the user-provided router config."""
router_config: Final = data.pop("user_config")
@ -45,9 +86,26 @@ def _route_user_config_request(data: dict, route_type: str):
filtered_config: Final = {k: v for k, v in router_config.items() if k in valid_args}
user_router: Final = litellm.Router(**filtered_config)
ret_val: Final = getattr(user_router, f"{route_type}")(**data)
user_router.discard()
return ret_val
try:
call = getattr(user_router, route_type)(**data)
except BaseException:
user_router.discard()
raise
async def _run():
stream_returned = False
try:
response = await call
if isinstance(response, AsyncIterator):
wrapped = _RouterOwnedStream(cast("AsyncIterator[object]", response), user_router)
stream_returned = True
return wrapped
return response
finally:
if not stream_returned:
user_router.discard()
return _run()
def _is_a2a_agent_model(model_name: object) -> bool:

View file

@ -0,0 +1,162 @@
import asyncio
import pytest
@pytest.mark.asyncio
async def test_user_config_router_stays_live_until_await_and_discards_after(monkeypatch):
import litellm
from litellm.proxy.route_llm_request import _route_user_config_request
routers = []
class FakeRouter:
@staticmethod
def get_valid_args():
return ["model_list"]
def __init__(self, **kwargs):
self.discarded = False
routers.append(self)
async def acompletion(self, **kwargs):
assert not self.discarded
return "done"
def discard(self):
assert not self.discarded
self.discarded = True
monkeypatch.setattr(litellm, "Router", FakeRouter)
call = _route_user_config_request({"user_config": {"model_list": []}}, "acompletion")
assert not routers[0].discarded
assert await call == "done"
assert routers[0].discarded
@pytest.mark.asyncio
@pytest.mark.parametrize("failure", [ValueError("provider failed"), asyncio.CancelledError()])
async def test_user_config_router_discards_on_provider_error(monkeypatch, failure):
import litellm
from litellm.proxy.route_llm_request import _route_user_config_request
routers = []
class FakeRouter:
@staticmethod
def get_valid_args():
return []
def __init__(self):
self.discarded = False
routers.append(self)
async def acompletion(self):
assert not self.discarded
raise failure
def discard(self):
self.discarded = True
monkeypatch.setattr(litellm, "Router", FakeRouter)
with pytest.raises(type(failure)):
await _route_user_config_request({"user_config": {}}, "acompletion")
assert routers[0].discarded
@pytest.mark.asyncio
async def test_user_config_stream_keeps_router_until_exhaustion_or_close(monkeypatch):
import litellm
from litellm.proxy.route_llm_request import _route_user_config_request
routers = []
class FakeStream:
def __init__(self, router):
self.router = router
self.count = 0
self.closed = False
self._hidden_params = {"model_id": "test"}
def __aiter__(self):
return self
async def __anext__(self):
assert not self.router.discarded
self.count += 1
if self.count > 1:
raise StopAsyncIteration
return "chunk"
async def aclose(self):
assert not self.router.discarded
self.closed = True
class FakeRouter:
@staticmethod
def get_valid_args():
return []
def __init__(self):
self.discarded = False
self.discards = 0
routers.append(self)
async def acompletion(self, **kwargs):
return FakeStream(self)
def discard(self):
self.discarded = True
self.discards += 1
monkeypatch.setattr(litellm, "Router", FakeRouter)
for close_early in (False, True):
stream = await _route_user_config_request({"user_config": {}, "stream": True}, "acompletion")
router = routers[-1]
assert not router.discarded
assert stream._hidden_params == {"model_id": "test"}
assert await anext(stream) == "chunk"
assert not router.discarded
if close_early:
await stream.aclose()
assert stream._stream.closed
else:
with pytest.raises(StopAsyncIteration):
await anext(stream)
assert router.discards == 1
await stream.aclose()
assert router.discards == 1
@pytest.mark.asyncio
async def test_user_config_stream_discards_when_iteration_raises(monkeypatch):
import litellm
from litellm.proxy.route_llm_request import _route_user_config_request
routers = []
class FakeRouter:
@staticmethod
def get_valid_args():
return []
def __init__(self):
self.discards = 0
routers.append(self)
async def acompletion(self):
async def stream():
assert self.discards == 0
raise ValueError("stream failed")
yield "unreachable"
return stream()
def discard(self):
self.discards += 1
monkeypatch.setattr(litellm, "Router", FakeRouter)
stream = await _route_user_config_request({"user_config": {}}, "acompletion")
with pytest.raises(ValueError, match="stream failed"):
await anext(stream)
assert routers[0].discards == 1