mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 8b8a42c74e into b781d157d7
This commit is contained in:
commit
ca07687140
3 changed files with 230 additions and 6 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
162
tests/test_litellm/proxy/test_user_config_router_lifetime.py
Normal file
162
tests/test_litellm/proxy/test_user_config_router_lifetime.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue