diff --git a/litellm/router.py b/litellm/router.py index b1038ca6002..f4d912eb5be 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -12630,6 +12630,12 @@ class Router: e, ) + async def _get_claude_code_session_router_binding(self, cache_key: str) -> object: + session_cache: Final = self._claude_code_session_router_cache + if session_cache.redis_cache is None: + return await session_cache.async_get_cache(key=cache_key) + return await session_cache.redis_cache.async_get_cache(key=cache_key) + async def _resolve_claude_code_session_router( self, model: str, @@ -12646,7 +12652,7 @@ class Router: agent_id: Final = self._request_header(request_kwargs, "x-claude-code-agent-id") if agent_id is not None: - bound_model: Final = await self._claude_code_session_router_cache.async_get_cache(key=cache_key) + bound_model: Final = await self._get_claude_code_session_router_binding(cache_key) if not isinstance(bound_model, str): return registered_model_name bound_registered_model: Final = self._get_model_from_alias(model=bound_model) or bound_model @@ -12664,7 +12670,6 @@ class Router: if self._request_header(request_kwargs, "x-app") != "cli": return registered_model_name if self._select_pre_routing_strategy(registered_model_name, request_kwargs) is None: - await self._delete_claude_code_session_router_binding(cache_key) return registered_model_name await self._claude_code_session_router_cache.async_set_cache( key=cache_key, diff --git a/tests/code_coverage_tests/router_code_coverage.py b/tests/code_coverage_tests/router_code_coverage.py index 60b56b7fac6..0af29f069c6 100644 --- a/tests/code_coverage_tests/router_code_coverage.py +++ b/tests/code_coverage_tests/router_code_coverage.py @@ -86,6 +86,7 @@ ignored_function_names = [ "_claude_code_session_router_cache_key", # Tested through Claude Code session routing in test_router.py "_delete_claude_code_session_router_binding", # Tested through Redis cleanup failure in test_router.py "_resolve_claude_code_session_router", # Tested through Claude Code session routing in test_router.py + "_get_claude_code_session_router_binding", # Tested through the two-worker session routing test in test_router.py ] diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index ef25502a4f9..1e5767aa645 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -8383,18 +8383,21 @@ class TestConsumedRequestTagsStamp: class TestClaudeCodeSubagentSessionRouterBinding: class _RewriteStrategy: + def __init__(self, routed_model: str = "cheap-model") -> None: + self.routed_model = routed_model + async def async_pre_routing_hook( self, model, request_kwargs, messages=None, input=None, specific_deployment=False ): from litellm.types.router import PreRoutingHookResponse return PreRoutingHookResponse( - model="cheap-model", + model=self.routed_model, messages=messages, routing_decision={ "router_model_name": "smart-router", "router_type": "complexity", - "routed_model": "cheap-model", + "routed_model": self.routed_model, "cause": "heuristic_scorer", }, ) @@ -8422,7 +8425,8 @@ class TestClaudeCodeSubagentSessionRouterBinding: num_retries=0, ) router.complexity_routers = { - "smart-router": [TaggedPreRoutingStrategy(tags=(), strategy=cls._RewriteStrategy())] + "smart-router": (TaggedPreRoutingStrategy(tags=(), strategy=cls._RewriteStrategy()),), + "premium-router": (TaggedPreRoutingStrategy(tags=(), strategy=cls._RewriteStrategy("expensive-model")),), } return router @@ -8467,7 +8471,7 @@ class TestClaudeCodeSubagentSessionRouterBinding: assert subagent_kwargs["metadata"]["routing_decision"]["router_model_name"] == "smart-router" @pytest.mark.asyncio - async def test_main_direct_model_clears_the_session_router(self): + async def test_main_thread_side_calls_to_a_plain_model_keep_the_session_router(self): router = self._router() await router.async_pre_routing_hook(model="smart-router", request_kwargs=self._request_kwargs()) @@ -8478,33 +8482,67 @@ class TestClaudeCodeSubagentSessionRouterBinding: request_kwargs=self._request_kwargs(agent_id="agent-1234"), ) - assert response is None + assert response is not None + assert response.model == "cheap-model" @pytest.mark.asyncio - async def test_redis_cleanup_failure_does_not_reject_a_direct_model_request(self): + async def test_redis_cleanup_failure_does_not_reject_a_subagent_request(self): from litellm.caching.caching import RedisCache router = self._router() + del router.complexity_routers["smart-router"] redis_cache = MagicMock(spec=RedisCache) + redis_cache.async_get_cache = AsyncMock(return_value="smart-router") redis_cache.async_delete_cache = AsyncMock(side_effect=ConnectionError("redis unavailable")) - - await router.async_pre_routing_hook(model="smart-router", request_kwargs=self._request_kwargs()) router._update_redis_cache(cache=redis_cache) response = await router.async_pre_routing_hook( model="expensive-model", - request_kwargs=self._request_kwargs(), + request_kwargs=self._request_kwargs(agent_id="agent-1234"), ) assert response is None redis_cache.async_delete_cache.assert_awaited_once() + @pytest.mark.asyncio + async def test_subagents_follow_the_main_threads_latest_router_across_workers(self): + from types import SimpleNamespace + + from litellm.caching.caching import RedisCache + + shared_binding = SimpleNamespace(value=None) + shared_redis = MagicMock(spec=RedisCache) + shared_redis.async_get_cache = AsyncMock(side_effect=lambda key, **_: shared_binding.value) + shared_redis.async_set_cache = AsyncMock( + side_effect=lambda key, value, **_: setattr(shared_binding, "value", value) + ) + main_worker, subagent_worker = self._router(), self._router() + main_worker._update_redis_cache(cache=shared_redis) + subagent_worker._update_redis_cache(cache=shared_redis) + + await main_worker.async_pre_routing_hook(model="smart-router", request_kwargs=self._request_kwargs()) + first = await subagent_worker.async_pre_routing_hook( + model="expensive-model", + request_kwargs=self._request_kwargs(agent_id="agent-1234"), + ) + await main_worker.async_pre_routing_hook(model="premium-router", request_kwargs=self._request_kwargs()) + second = await subagent_worker.async_pre_routing_hook( + model="expensive-model", + request_kwargs=self._request_kwargs(agent_id="agent-1234"), + ) + + assert first is not None + assert first.model == "cheap-model" + assert second is not None + assert second.model == "expensive-model" + assert shared_binding.value == "premium-router" + @pytest.mark.asyncio async def test_no_pre_routing_strategies_means_no_session_cache_traffic(self): from litellm.caching.caching import RedisCache router = self._router() - router.complexity_routers = {} + router.complexity_routers.clear() redis_cache = MagicMock(spec=RedisCache) redis_cache.async_get_cache = AsyncMock(return_value=None) redis_cache.async_set_cache = AsyncMock()