mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
fix(router): route Claude Code subagents through session router
This commit is contained in:
parent
3dac3f7a36
commit
1cd99a036e
2 changed files with 208 additions and 1 deletions
|
|
@ -353,6 +353,8 @@ _PreRoutingStrategyT = TypeVar("_PreRoutingStrategyT")
|
|||
|
||||
_ALIAS_PARAMS_NEVER_FORWARDED: Final = frozenset({"model", "api_base", "api_key", "api_version"})
|
||||
_ALIAS_MARKER_FORWARDED_PARAMS_KWARG: Final = "_alias_marker_forwarded_params"
|
||||
_CLAUDE_CODE_SESSION_ID_RE: Final = re.compile(r"^[a-zA-Z0-9_\-]{8,}$")
|
||||
_CLAUDE_CODE_SESSION_ROUTER_TTL_SECONDS: Final = 3600
|
||||
|
||||
|
||||
def _stream_chunks_have_generated_content(chunks: Sequence[ModelResponseStream]) -> bool:
|
||||
|
|
@ -12546,6 +12548,77 @@ class Router:
|
|||
return None
|
||||
return candidates[0]
|
||||
|
||||
@staticmethod
|
||||
def _request_header(request_kwargs: Mapping[str, object], header_name: str) -> str | None:
|
||||
proxy_server_request: Final = request_kwargs.get("proxy_server_request")
|
||||
if not isinstance(proxy_server_request, Mapping):
|
||||
return None
|
||||
headers: Final = proxy_server_request.get("headers")
|
||||
if not isinstance(headers, Mapping):
|
||||
return None
|
||||
return next(
|
||||
(
|
||||
value
|
||||
for key, value in headers.items()
|
||||
if isinstance(key, str) and key.lower() == header_name and isinstance(value, str)
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
def _claude_code_session_router_cache_key(self, request_kwargs: Mapping[str, object]) -> str | None:
|
||||
session_id: Final = self._request_header(request_kwargs, "x-claude-code-session-id")
|
||||
if session_id is None or _CLAUDE_CODE_SESSION_ID_RE.fullmatch(session_id) is None:
|
||||
return None
|
||||
metadata_name: Final = "litellm_metadata" if "litellm_metadata" in request_kwargs else "metadata"
|
||||
metadata: Final = request_kwargs.get(metadata_name)
|
||||
if not isinstance(metadata, Mapping):
|
||||
return None
|
||||
caller_scope: Final = metadata.get("user_api_key_hash")
|
||||
if not isinstance(caller_scope, str) or not caller_scope:
|
||||
return None
|
||||
return f"claude_code_session_router:v1:{caller_scope}:{session_id}"
|
||||
|
||||
async def _resolve_claude_code_session_router(
|
||||
self,
|
||||
model: str,
|
||||
registered_model_name: str,
|
||||
request_kwargs: Mapping[str, object],
|
||||
) -> str:
|
||||
cache_key: Final = self._claude_code_session_router_cache_key(request_kwargs)
|
||||
if cache_key is None or not isinstance(request_kwargs, dict):
|
||||
return registered_model_name
|
||||
|
||||
agent_id: Final = self._request_header(request_kwargs, "x-claude-code-agent-id")
|
||||
if agent_id is not None:
|
||||
bound_model: Final = await self.cache.async_get_cache(key=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
|
||||
if self._select_pre_routing_strategy(bound_registered_model, request_kwargs) is None:
|
||||
await self.cache.async_delete_cache(key=cache_key)
|
||||
return registered_model_name
|
||||
await self.cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value=bound_model,
|
||||
ttl=_CLAUDE_CODE_SESSION_ROUTER_TTL_SECONDS,
|
||||
)
|
||||
self._stamp_or_clear_metadata_key(request_kwargs, "model_group", bound_model)
|
||||
return bound_registered_model
|
||||
|
||||
if self._request_header(request_kwargs, "x-app") != "cli":
|
||||
return registered_model_name
|
||||
if request_kwargs.get("fallback_depth") not in (None, 0):
|
||||
return registered_model_name
|
||||
if self._select_pre_routing_strategy(registered_model_name, request_kwargs) is None:
|
||||
await self.cache.async_delete_cache(key=cache_key)
|
||||
return registered_model_name
|
||||
await self.cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value=model,
|
||||
ttl=_CLAUDE_CODE_SESSION_ROUTER_TTL_SECONDS,
|
||||
)
|
||||
return registered_model_name
|
||||
|
||||
async def async_pre_routing_hook(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -12565,7 +12638,12 @@ class Router:
|
|||
the alias, since spend metadata is stamped before routing and the response carries the tier
|
||||
group the strategy picked.
|
||||
"""
|
||||
registered_model_name: Final = self._get_model_from_alias(model=model) or model
|
||||
requested_registered_model_name: Final = self._get_model_from_alias(model=model) or model
|
||||
registered_model_name: Final = await self._resolve_claude_code_session_router(
|
||||
model=model,
|
||||
registered_model_name=requested_registered_model_name,
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
|
||||
#########################################################
|
||||
# Run the routing-plugin pipeline, if any plugins are configured.
|
||||
|
|
|
|||
|
|
@ -8316,6 +8316,135 @@ class TestConsumedRequestTagsStamp:
|
|||
assert CONSUMED_REQUEST_TAGS_METADATA_KEY not in request_kwargs["metadata"]
|
||||
|
||||
|
||||
class TestClaudeCodeSubagentSessionRouterBinding:
|
||||
class _RewriteStrategy:
|
||||
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",
|
||||
messages=messages,
|
||||
routing_decision={
|
||||
"router_model_name": "smart-router",
|
||||
"router_type": "complexity",
|
||||
"routed_model": "cheap-model",
|
||||
"cause": "heuristic_scorer",
|
||||
},
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _router(cls) -> "litellm.Router":
|
||||
from litellm.types.router import TaggedPreRoutingStrategy
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "cheap-model",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini", "mock_response": "cheap response"},
|
||||
},
|
||||
{
|
||||
"model_name": "expensive-model",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "mock_response": "expensive response"},
|
||||
},
|
||||
]
|
||||
)
|
||||
router.complexity_routers = {
|
||||
"smart-router": [TaggedPreRoutingStrategy(tags=(), strategy=cls._RewriteStrategy())]
|
||||
}
|
||||
return router
|
||||
|
||||
@staticmethod
|
||||
def _request_kwargs(
|
||||
*,
|
||||
key_hash: str = "key-hash-a",
|
||||
app: str = "cli",
|
||||
agent_id: str | None = None,
|
||||
fallback_depth: int | None = None,
|
||||
) -> dict:
|
||||
headers = {
|
||||
"X-Claude-Code-Session-Id": "session-1234",
|
||||
"x-app": app,
|
||||
**({"x-claude-code-agent-id": agent_id} if agent_id is not None else {}),
|
||||
}
|
||||
return {
|
||||
"metadata": {"user_api_key_hash": key_hash},
|
||||
"proxy_server_request": {"headers": headers},
|
||||
**({"fallback_depth": fallback_depth} if fallback_depth is not None else {}),
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_subagent_concrete_model_uses_the_main_sessions_router(self):
|
||||
router = self._router()
|
||||
|
||||
await router.acompletion(
|
||||
model="smart-router",
|
||||
messages=[{"role": "user", "content": "main turn"}],
|
||||
**self._request_kwargs(),
|
||||
)
|
||||
subagent_kwargs = self._request_kwargs(agent_id="agent-1234")
|
||||
|
||||
response = await router.acompletion(
|
||||
model="expensive-model",
|
||||
messages=[{"role": "user", "content": "subagent turn"}],
|
||||
**subagent_kwargs,
|
||||
)
|
||||
|
||||
assert response.choices[0].message.content == "cheap response"
|
||||
assert subagent_kwargs["metadata"]["model_group"] == "smart-router"
|
||||
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):
|
||||
router = self._router()
|
||||
|
||||
await router.async_pre_routing_hook(model="smart-router", request_kwargs=self._request_kwargs())
|
||||
await router.async_pre_routing_hook(model="expensive-model", request_kwargs=self._request_kwargs())
|
||||
|
||||
response = await router.async_pre_routing_hook(
|
||||
model="expensive-model",
|
||||
request_kwargs=self._request_kwargs(agent_id="agent-1234"),
|
||||
)
|
||||
|
||||
assert response is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_background_and_fallback_requests_do_not_clear_the_session_router(self):
|
||||
router = self._router()
|
||||
|
||||
await router.async_pre_routing_hook(model="smart-router", request_kwargs=self._request_kwargs())
|
||||
await router.async_pre_routing_hook(
|
||||
model="expensive-model",
|
||||
request_kwargs=self._request_kwargs(app="cli-bg"),
|
||||
)
|
||||
await router.async_pre_routing_hook(
|
||||
model="expensive-model",
|
||||
request_kwargs=self._request_kwargs(fallback_depth=1),
|
||||
)
|
||||
|
||||
response = await router.async_pre_routing_hook(
|
||||
model="expensive-model",
|
||||
request_kwargs=self._request_kwargs(agent_id="agent-1234"),
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert response.model == "cheap-model"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_router_binding_is_scoped_to_the_authenticated_key(self):
|
||||
router = self._router()
|
||||
|
||||
await router.async_pre_routing_hook(model="smart-router", request_kwargs=self._request_kwargs())
|
||||
|
||||
response = await router.async_pre_routing_hook(
|
||||
model="expensive-model",
|
||||
request_kwargs=self._request_kwargs(key_hash="key-hash-b", agent_id="agent-1234"),
|
||||
)
|
||||
|
||||
assert response is None
|
||||
|
||||
|
||||
class TestAutoRouterMaxInputCharsWiring:
|
||||
"""`auto_router_max_input_chars` on the deployment has to reach the AutoRouter that embeds prompts.
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue