From 7c026bdbd68486c3360a51854a6f7421a3ae9e42 Mon Sep 17 00:00:00 2001 From: Charan Rathore Date: Mon, 28 Sep 2026 10:45:22 +0530 Subject: [PATCH] fix(proxy): retain user-config router through call and stream --- litellm/proxy/common_request_processing.py | 6 +- litellm/proxy/route_llm_request.py | 68 +++++++- .../proxy/test_user_config_router_lifetime.py | 162 ++++++++++++++++++ 3 files changed, 230 insertions(+), 6 deletions(-) create mode 100644 tests/test_litellm/proxy/test_user_config_router_lifetime.py diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 15610da9aec..9322901e537 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -3453,9 +3453,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 diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 536c58df65a..9309de387f4 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -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): + 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: diff --git a/tests/test_litellm/proxy/test_user_config_router_lifetime.py b/tests/test_litellm/proxy/test_user_config_router_lifetime.py new file mode 100644 index 00000000000..28a59b00218 --- /dev/null +++ b/tests/test_litellm/proxy/test_user_config_router_lifetime.py @@ -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