fix(router): route Claude Code subagents through session router

This commit is contained in:
moe-berri 2026-09-01 17:29:58 -07:00
parent 3dac3f7a36
commit 1cd99a036e
2 changed files with 208 additions and 1 deletions

View file

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

View file

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