mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(router): keep Claude Code session bindings across side calls and workers
This commit is contained in:
parent
0f6d983c70
commit
f6eff1bde0
3 changed files with 56 additions and 12 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue