diff --git a/litellm/router.py b/litellm/router.py index e2a1b7a9f7a..6061baaaca1 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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) diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 253c265c2f8..f8aa4173f93 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -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"