fix(router): keep Claude Code session bindings across side calls and workers

This commit is contained in:
mateo-berri 2026-09-02 13:55:16 -07:00
parent 0f6d983c70
commit f6eff1bde0
3 changed files with 56 additions and 12 deletions

View file

@ -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,

View file

@ -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
]

View file

@ -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()