mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +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,
|
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(
|
async def _resolve_claude_code_session_router(
|
||||||
self,
|
self,
|
||||||
model: str,
|
model: str,
|
||||||
|
|
@ -12646,7 +12652,7 @@ class Router:
|
||||||
|
|
||||||
agent_id: Final = self._request_header(request_kwargs, "x-claude-code-agent-id")
|
agent_id: Final = self._request_header(request_kwargs, "x-claude-code-agent-id")
|
||||||
if agent_id is not None:
|
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):
|
if not isinstance(bound_model, str):
|
||||||
return registered_model_name
|
return registered_model_name
|
||||||
bound_registered_model: Final = self._get_model_from_alias(model=bound_model) or bound_model
|
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":
|
if self._request_header(request_kwargs, "x-app") != "cli":
|
||||||
return registered_model_name
|
return registered_model_name
|
||||||
if self._select_pre_routing_strategy(registered_model_name, request_kwargs) is None:
|
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
|
return registered_model_name
|
||||||
await self._claude_code_session_router_cache.async_set_cache(
|
await self._claude_code_session_router_cache.async_set_cache(
|
||||||
key=cache_key,
|
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
|
"_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
|
"_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
|
"_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 TestClaudeCodeSubagentSessionRouterBinding:
|
||||||
class _RewriteStrategy:
|
class _RewriteStrategy:
|
||||||
|
def __init__(self, routed_model: str = "cheap-model") -> None:
|
||||||
|
self.routed_model = routed_model
|
||||||
|
|
||||||
async def async_pre_routing_hook(
|
async def async_pre_routing_hook(
|
||||||
self, model, request_kwargs, messages=None, input=None, specific_deployment=False
|
self, model, request_kwargs, messages=None, input=None, specific_deployment=False
|
||||||
):
|
):
|
||||||
from litellm.types.router import PreRoutingHookResponse
|
from litellm.types.router import PreRoutingHookResponse
|
||||||
|
|
||||||
return PreRoutingHookResponse(
|
return PreRoutingHookResponse(
|
||||||
model="cheap-model",
|
model=self.routed_model,
|
||||||
messages=messages,
|
messages=messages,
|
||||||
routing_decision={
|
routing_decision={
|
||||||
"router_model_name": "smart-router",
|
"router_model_name": "smart-router",
|
||||||
"router_type": "complexity",
|
"router_type": "complexity",
|
||||||
"routed_model": "cheap-model",
|
"routed_model": self.routed_model,
|
||||||
"cause": "heuristic_scorer",
|
"cause": "heuristic_scorer",
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
@ -8422,7 +8425,8 @@ class TestClaudeCodeSubagentSessionRouterBinding:
|
||||||
num_retries=0,
|
num_retries=0,
|
||||||
)
|
)
|
||||||
router.complexity_routers = {
|
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
|
return router
|
||||||
|
|
||||||
|
|
@ -8467,7 +8471,7 @@ class TestClaudeCodeSubagentSessionRouterBinding:
|
||||||
assert subagent_kwargs["metadata"]["routing_decision"]["router_model_name"] == "smart-router"
|
assert subagent_kwargs["metadata"]["routing_decision"]["router_model_name"] == "smart-router"
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@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()
|
router = self._router()
|
||||||
|
|
||||||
await router.async_pre_routing_hook(model="smart-router", request_kwargs=self._request_kwargs())
|
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"),
|
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
|
@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
|
from litellm.caching.caching import RedisCache
|
||||||
|
|
||||||
router = self._router()
|
router = self._router()
|
||||||
|
del router.complexity_routers["smart-router"]
|
||||||
redis_cache = MagicMock(spec=RedisCache)
|
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"))
|
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)
|
router._update_redis_cache(cache=redis_cache)
|
||||||
|
|
||||||
response = await router.async_pre_routing_hook(
|
response = await router.async_pre_routing_hook(
|
||||||
model="expensive-model",
|
model="expensive-model",
|
||||||
request_kwargs=self._request_kwargs(),
|
request_kwargs=self._request_kwargs(agent_id="agent-1234"),
|
||||||
)
|
)
|
||||||
|
|
||||||
assert response is None
|
assert response is None
|
||||||
redis_cache.async_delete_cache.assert_awaited_once()
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_no_pre_routing_strategies_means_no_session_cache_traffic(self):
|
async def test_no_pre_routing_strategies_means_no_session_cache_traffic(self):
|
||||||
from litellm.caching.caching import RedisCache
|
from litellm.caching.caching import RedisCache
|
||||||
|
|
||||||
router = self._router()
|
router = self._router()
|
||||||
router.complexity_routers = {}
|
router.complexity_routers.clear()
|
||||||
redis_cache = MagicMock(spec=RedisCache)
|
redis_cache = MagicMock(spec=RedisCache)
|
||||||
redis_cache.async_get_cache = AsyncMock(return_value=None)
|
redis_cache.async_get_cache = AsyncMock(return_value=None)
|
||||||
redis_cache.async_set_cache = AsyncMock()
|
redis_cache.async_set_cache = AsyncMock()
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue