refactor(router): keep request kwargs immutable in the priority path

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mateo 2026-08-13 20:03:12 +00:00
parent 642f941296
commit 05f314546e
2 changed files with 30 additions and 29 deletions

View file

@ -2140,7 +2140,7 @@ class Router:
kwargs["original_function"] = self._acompletion
self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs)
request_priority: Final = kwargs.pop("priority", None) or self.default_priority
request_priority: Final = kwargs.get("priority") or self.default_priority
start_time: Final = time.time()
_is_prompt_management_model: Final = self._is_prompt_management_model(model)
@ -2156,7 +2156,7 @@ class Router:
priority=request_priority,
original_function=self.async_function_with_fallbacks,
args=(),
kwargs=kwargs,
kwargs={key: value for key, value in kwargs.items() if key != "priority"},
)
else:
response = await self.async_function_with_fallbacks(**kwargs)

View file

@ -4,6 +4,9 @@ import json
import logging
import os
import sys
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -16,6 +19,7 @@ sys.path.insert(
import litellm
from litellm.exceptions import MidStreamFallbackError
from litellm.integrations.custom_logger import CustomLogger
from litellm.scheduler import FlowItem
def test_update_kwargs_does_not_mutate_defaults_and_merges_metadata():
@ -7968,39 +7972,41 @@ def _priority_router() -> litellm.Router:
)
@dataclass(slots=True)
class _QueueSpy:
"""Records the priority of every request handed to the scheduler queue"""
inner: Callable[..., Awaitable[None]]
priorities: tuple[int, ...] = ()
async def __call__(self, request: FlowItem) -> None:
self.priorities = (*self.priorities, request.priority)
await self.inner(request=request)
@pytest.mark.asyncio
async def test_acompletion_uses_default_priority_when_request_has_none():
"""default_priority must be forwarded to the scheduler instead of blowing up the request"""
router = _priority_router()
queued: list[int] = []
original_add_request = router.scheduler.add_request
router: Final = _priority_router()
spy: Final = _QueueSpy(inner=router.scheduler.add_request)
async def spy(request):
queued.append(request.priority)
return await original_add_request(request=request)
with patch.object(router.scheduler, "add_request", side_effect=spy):
with patch.object(router.scheduler, "add_request", new=spy):
response = await router.acompletion(
model="code",
messages=[{"role": "user", "content": "Hi"}],
mock_response="Hello",
)
assert queued == [10]
assert spy.priorities == (10,)
assert response._hidden_params["additional_headers"]["x-litellm-request-prioritization-used"] is True
@pytest.mark.asyncio
async def test_acompletion_request_priority_wins_over_default_priority():
router = _priority_router()
queued: list[int] = []
original_add_request = router.scheduler.add_request
router: Final = _priority_router()
spy: Final = _QueueSpy(inner=router.scheduler.add_request)
async def spy(request):
queued.append(request.priority)
return await original_add_request(request=request)
with patch.object(router.scheduler, "add_request", side_effect=spy):
with patch.object(router.scheduler, "add_request", new=spy):
response = await router.acompletion(
model="code",
messages=[{"role": "user", "content": "Hi"}],
@ -8008,22 +8014,17 @@ async def test_acompletion_request_priority_wins_over_default_priority():
mock_response="Hello",
)
assert queued == [5]
assert spy.priorities == (5,)
assert response.choices[0].message.content == "Hello"
@pytest.mark.asyncio
async def test_schedule_acompletion_queues_once_with_default_priority_configured():
"""schedule_acompletion must not re-enter the scheduler via acompletion's default_priority"""
router = _priority_router()
queued: list[int] = []
original_add_request = router.scheduler.add_request
router: Final = _priority_router()
spy: Final = _QueueSpy(inner=router.scheduler.add_request)
async def spy(request):
queued.append(request.priority)
return await original_add_request(request=request)
with patch.object(router.scheduler, "add_request", side_effect=spy):
with patch.object(router.scheduler, "add_request", new=spy):
response = await router.schedule_acompletion(
model="code",
messages=[{"role": "user", "content": "Hi"}],
@ -8031,5 +8032,5 @@ async def test_schedule_acompletion_queues_once_with_default_priority_configured
mock_response="Hello",
)
assert queued == [3]
assert spy.priorities == (3,)
assert response.choices[0].message.content == "Hello"