diff --git a/litellm/router.py b/litellm/router.py index 249039fb8ba..67af54b8047 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -7661,14 +7661,20 @@ class Router: if self.optional_callbacks is None: self.optional_callbacks = [] + router_ref: Final = weakref.ref(self) + + def session_affinity_group_ttls_provider() -> Mapping[str, int]: + router: Final = router_ref() + if router is None: + return MappingProxyType({}) + return router._get_complexity_router_session_affinity_group_ttls() + existing_affinity_callback: DeploymentAffinityCheck | None = next( (callback for callback in self.optional_callbacks if isinstance(callback, DeploymentAffinityCheck)), None, ) if existing_affinity_callback is not None: - existing_affinity_callback.session_affinity_group_ttls = ( - self._get_complexity_router_session_affinity_group_ttls - ) + existing_affinity_callback.session_affinity_group_ttls = session_affinity_group_ttls_provider return affinity_callback: Final = DeploymentAffinityCheck( @@ -7678,7 +7684,7 @@ class Router: enable_responses_api_affinity=False, enable_session_id_affinity=False, model_group_affinity_config=self.model_group_affinity_config, - session_affinity_group_ttls=self._get_complexity_router_session_affinity_group_ttls, + session_affinity_group_ttls=session_affinity_group_ttls_provider, ) self.optional_callbacks.append(affinity_callback) litellm.logging_callback_manager.add_litellm_callback(affinity_callback) diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index b3f1e929741..2a269affe32 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -21,6 +21,7 @@ from litellm import Router from litellm._logging import verbose_router_logger from litellm.caching.dual_cache import DualCache from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY +from litellm.router import _iter_complexity_router_session_affinity_groups from litellm.router_strategy.complexity_router.complexity_router import ( _CLASSIFICATION_CURRENT_MESSAGE_ONLY, _CLASSIFICATION_WITH_CONVERSATION, @@ -80,6 +81,29 @@ def complexity_router(mock_router_instance, basic_config): ) +def test_router_complexity_session_affinity_provider_registers_callback(): + parent_router = Router(model_list=[]) + complexity_router = ComplexityRouter( + model_name="test-complexity-router", + litellm_router_instance=parent_router, + complexity_router_config={ + "tiers": {"SIMPLE": "target-group"}, + "session_affinity": True, + "session_affinity_ttl_seconds": 17, + }, + ) + parent_router.complexity_routers["test-complexity-router"] = [ + TaggedPreRoutingStrategy(tags=(), strategy=complexity_router) + ] + + assert tuple(_iter_complexity_router_session_affinity_groups(complexity_router)) == (("target-group", 17),) + assert dict(parent_router._get_complexity_router_session_affinity_group_ttls()) == {"target-group": 17} + + parent_router._ensure_deployment_affinity_check() + + assert parent_router.optional_callbacks + + class TestDimensionScore: """Test the DimensionScore class.""" @@ -4306,7 +4330,6 @@ class TestRoutingDecisionContents: # The score is still recorded, but the cause is what says it did not decide. assert decision["score"] < decision["tier_boundaries"]["complex_reasoning"] - @pytest.mark.asyncio async def test_an_unrenamed_router_writes_no_tier_label(self, complexity_router): """Renaming is opt-in, so a deployment that never renamed must gain no new key. @@ -5761,7 +5784,9 @@ class TestCustomClassifierSystemPrompt: @pytest.mark.asyncio async def test_custom_prompt_is_sent_verbatim_as_the_system_role(self, mock_router_instance, llm_classifier_config): - custom = "Classify the data sensitivity: SIMPLE=public, MEDIUM=internal, COMPLEX=confidential, REASONING=regulated." + custom = ( + "Classify the data sensitivity: SIMPLE=public, MEDIUM=internal, COMPLEX=confidential, REASONING=regulated." + ) router = ComplexityRouter( model_name="test-complexity-router", litellm_router_instance=mock_router_instance, diff --git a/tests/test_litellm/router_utils/pre_call_checks/test_session_id_affinity.py b/tests/test_litellm/router_utils/pre_call_checks/test_session_id_affinity.py index 37521701434..abf9412c9fd 100644 --- a/tests/test_litellm/router_utils/pre_call_checks/test_session_id_affinity.py +++ b/tests/test_litellm/router_utils/pre_call_checks/test_session_id_affinity.py @@ -1,6 +1,8 @@ import asyncio +import gc import os import sys +import weakref from unittest.mock import AsyncMock, patch import pytest @@ -426,6 +428,26 @@ def test_complexity_router_plugins_do_not_enable_deployment_affinity(): assert not any(isinstance(callback, DeploymentAffinityCheck) for callback in router.optional_callbacks or []) +def test_complexity_router_affinity_provider_does_not_keep_router_alive(): + router = _complexity_router() + affinity_callback = next( + callback for callback in router.optional_callbacks or [] if isinstance(callback, DeploymentAffinityCheck) + ) + router.optional_callbacks = [] + router.discard() + router_ref = weakref.ref(router) + + try: + del router + gc.collect() + + assert router_ref() is None + assert affinity_callback.session_affinity_group_ttls is not None + assert affinity_callback.session_affinity_group_ttls() == {} + finally: + litellm.logging_callback_manager.remove_callback_from_all_lists(affinity_callback) + + @pytest.mark.asyncio async def test_session_affinity_without_stable_model_map_key_falls_back_to_normal_selection(): callback = DeploymentAffinityCheck(