Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_decrease_anys_opus5_r4

This commit is contained in:
mateo-berri 2026-09-04 20:06:56 -07:00
commit 4b0fa87b9e
41 changed files with 2966 additions and 319 deletions

View file

@ -229,6 +229,45 @@ def _content_parts_contain_image(parts: Sequence[object]) -> bool:
return False
def anthropic_image_source_to_openai_url(image_source: Mapping[str, object]) -> str | None:
"""Data or remote URL for an Anthropic ``source`` block, in the form chat completions expects."""
source_type: Final = image_source.get("type")
if source_type == "base64":
media_type: Final = image_source.get("media_type") or "image/jpeg"
image_data: Final = image_source.get("data") or ""
return f"data:{media_type};base64,{image_data}" if image_data else None
if source_type == "url":
url: Final = image_source.get("url")
return url if isinstance(url, str) else ""
return None
def _image_part_url(part: Mapping[str, object]) -> str | None:
"""The image URL carried by one content part, whichever of the three dialects wrote it."""
part_type: Final = part.get("type")
if part_type == "image_url":
image_url: Final = part.get("image_url")
if isinstance(image_url, str):
return image_url
return image_url.get("url") if isinstance(image_url, Mapping) else None
if part_type == "input_image":
responses_url: Final = part.get("image_url")
return responses_url if isinstance(responses_url, str) else None
if part_type == "image":
source: Final = part.get("source")
return anthropic_image_source_to_openai_url(source) if isinstance(source, Mapping) else None
return None
def as_openai_image_part(part: Mapping[str, object]) -> ChatCompletionImageObject | None:
"""One image content part rewritten into chat-completions dialect, or None when it is not one.
Rebuilt rather than forwarded so no caller-controlled key beyond the URL rides along.
"""
url: Final = _image_part_url(part)
return {"type": "image_url", "image_url": {"url": url}} if url else None
def request_contains_image_content(messages: Sequence[Mapping[str, object]]) -> bool:
"""Whether any message carries an image content part, across the dialects that reach
pre-routing hooks untranslated: chat-completions ``image_url``, Responses ``input_image``,

View file

@ -99,6 +99,7 @@ def create_tool_name_mapping(
from openai.types.chat.chat_completion_chunk import Choice as OpenAIStreamingChoice
from litellm.litellm_core_utils.prompt_templates.common_utils import (
anthropic_image_source_to_openai_url,
parse_tool_call_arguments,
reasoning_content_from_thinking_blocks,
with_prompt_cache_breakpoint,
@ -1225,20 +1226,7 @@ class LiteLLMAnthropicMessagesAdapter:
"""
if not isinstance(image_source, dict):
return None
source_type: Final = image_source.get("type")
if source_type == "base64":
# Base64 image format
media_type: Final = image_source.get("media_type", "image/jpeg")
image_data: Final = image_source.get("data", "")
if image_data:
return f"data:{media_type};base64,{image_data}"
elif source_type == "url":
# URL-referenced image format
return image_source.get("url", "")
return None
return anthropic_image_source_to_openai_url(image_source)
def _tool_result_content(self, raw_content: object) -> ToolResultContent:
if isinstance(raw_content, str):

View file

@ -6,6 +6,7 @@ from openai.types.responses import ResponseReasoningItem
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
from litellm.llms.azure.chat.gpt_5_transformation import AzureOpenAIGPT5Config
from litellm.llms.azure.common_utils import BaseAzureLLM
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
from litellm.types.llms.openai import *
@ -29,6 +30,14 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
def custom_llm_provider(self) -> LlmProviders:
return LlmProviders.AZURE
@staticmethod
def _supports_reasoning_effort_none(model: str) -> bool:
return AzureOpenAIGPT5Config._supports_reasoning_effort_level(model, "none")
@staticmethod
def _effort_resolves_to_none(model: str, effort: str | None) -> bool:
return AzureOpenAIGPT5Config.effort_resolves_to_none(model, effort)
def get_supported_openai_params(self, model: str) -> list:
"""
Azure Responses API does not support context_management (compaction).

View file

@ -7157,6 +7157,53 @@
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": false
},
"azure/gpt-6-astra": {
"cache_creation_input_token_cost": 1.25e-05,
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-05,
"cache_read_input_token_cost": 1e-06,
"cache_read_input_token_cost_above_272k_tokens": 2e-06,
"input_cost_per_token": 1e-05,
"input_cost_per_token_above_272k_tokens": 2e-05,
"litellm_provider": "azure",
"max_input_tokens": 922000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 5e-05,
"output_cost_per_token_above_272k_tokens": 7.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_computer_use": true,
"supports_function_calling": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": false,
"supports_native_streaming": true,
"supports_none_reasoning_effort": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true,
"supports_xhigh_reasoning_effort": true
},
"azure/us/gpt-5.6": {
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_272k_tokens": 1.375e-05,
@ -7376,6 +7423,53 @@
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": false
},
"azure/us/gpt-6-astra": {
"cache_creation_input_token_cost": 1.375e-05,
"cache_creation_input_token_cost_above_272k_tokens": 2.75e-05,
"cache_read_input_token_cost": 1.1e-06,
"cache_read_input_token_cost_above_272k_tokens": 2.2e-06,
"input_cost_per_token": 1.1e-05,
"input_cost_per_token_above_272k_tokens": 2.2e-05,
"litellm_provider": "azure",
"max_input_tokens": 922000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 5.5e-05,
"output_cost_per_token_above_272k_tokens": 8.25e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_computer_use": true,
"supports_function_calling": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": false,
"supports_native_streaming": true,
"supports_none_reasoning_effort": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true,
"supports_xhigh_reasoning_effort": true
},
"azure/eu/gpt-5.6": {
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_272k_tokens": 1.375e-05,

View file

@ -741,6 +741,32 @@
"title": "AccessGroupInfo",
"type": "object"
},
"AccessGroupResource": {
"description": "A resource referenced by an access group. `name` is null when the id no longer resolves or has no alias.",
"properties": {
"id": {
"title": "Id",
"type": "string"
},
"name": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Name"
}
},
"required": [
"id",
"name"
],
"title": "AccessGroupResource",
"type": "object"
},
"AccessGroupResponse": {
"properties": {
"access_agent_ids": {
@ -750,6 +776,13 @@
"title": "Access Agent Ids",
"type": "array"
},
"access_agents": {
"items": {
"$ref": "#/components/schemas/AccessGroupResource"
},
"title": "Access Agents",
"type": "array"
},
"access_group_id": {
"title": "Access Group Id",
"type": "string"
@ -765,6 +798,13 @@
"title": "Access Mcp Server Ids",
"type": "array"
},
"access_mcp_servers": {
"items": {
"$ref": "#/components/schemas/AccessGroupResource"
},
"title": "Access Mcp Servers",
"type": "array"
},
"access_model_names": {
"items": {
"type": "string"
@ -779,6 +819,13 @@
"title": "Assigned Key Ids",
"type": "array"
},
"assigned_keys": {
"items": {
"$ref": "#/components/schemas/AccessGroupResource"
},
"title": "Assigned Keys",
"type": "array"
},
"assigned_team_ids": {
"items": {
"type": "string"
@ -786,6 +833,13 @@
"title": "Assigned Team Ids",
"type": "array"
},
"assigned_teams": {
"items": {
"$ref": "#/components/schemas/AccessGroupResource"
},
"title": "Assigned Teams",
"type": "array"
},
"created_at": {
"format": "date-time",
"title": "Created At",
@ -838,6 +892,10 @@
"access_agent_ids",
"assigned_team_ids",
"assigned_key_ids",
"access_mcp_servers",
"access_agents",
"assigned_teams",
"assigned_keys",
"created_at",
"updated_at"
],

View file

@ -1,16 +1,20 @@
from collections.abc import Mapping, Sequence
import asyncio
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final, Protocol
from fastapi import APIRouter, Depends, HTTPException, status
from litellm._logging import verbose_proxy_logger
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
from litellm.proxy._types import (
CommonProxyErrors,
LiteLLM_AccessGroupTable,
LitellmUserRoles,
UserAPIKeyAuth,
)
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.proxy.auth.auth_checks import (
_cache_access_object,
_cache_key_object,
@ -20,10 +24,16 @@ from litellm.proxy.auth.auth_checks import (
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.management_helpers.access_group_team_sync import invalidate_access_group_cache
from litellm.proxy.utils import get_prisma_client_or_throw
from litellm.proxy.management_helpers.resource_display_names import (
agent_display_names,
key_display_names,
mcp_server_display_names,
)
from litellm.proxy.utils import PrismaClient, get_prisma_client_or_throw
from litellm.repositories.table_repositories import AccessGroupRepository, TeamRepository
from litellm.types.access_group import (
AccessGroupCreateRequest,
AccessGroupResource,
AccessGroupResponse,
AccessGroupUpdateRequest,
)
@ -37,6 +47,12 @@ class _AccessGroupRecord(Protocol):
@property
def access_group_id(self) -> str: ...
@property
def access_mcp_server_ids(self) -> Sequence[str] | None: ...
@property
def access_agent_ids(self) -> Sequence[str] | None: ...
@property
def assigned_team_ids(self) -> Sequence[str] | None: ...
@ -50,6 +66,9 @@ class _TeamRecord(Protocol):
@property
def team_id(self) -> str: ...
@property
def team_alias(self) -> str | None: ...
@property
def access_group_ids(self) -> Sequence[str] | None: ...
@ -120,16 +139,75 @@ def _require_admin_view(user_api_key_dict: UserAPIKeyAuth) -> None:
)
@dataclass(frozen=True, slots=True)
class _ResourceNames:
mcp_servers: Mapping[str, str]
agents: Mapping[str, str]
teams: Mapping[str, str | None]
keys: Mapping[str, str]
def _label(ids: Sequence[str], names: Mapping[str, str | None]) -> tuple[AccessGroupResource, ...]:
return tuple(AccessGroupResource(id=resource_id, name=names.get(resource_id)) for resource_id in ids)
def _record_to_response(
record: _AccessGroupRecord, *, assigned_team_ids: Sequence[str] | None = None
record: _AccessGroupRecord, *, assigned_team_ids: Sequence[str], names: _ResourceNames
) -> AccessGroupResponse:
stored: Final = record.dict()
payload: Final = (
stored if assigned_team_ids is None else MappingProxyType({**stored, "assigned_team_ids": assigned_team_ids})
payload: Final = MappingProxyType(
{
**record.dict(),
"assigned_team_ids": assigned_team_ids,
"access_mcp_servers": _label(record.access_mcp_server_ids or (), names.mcp_servers),
"access_agents": _label(record.access_agent_ids or (), names.agents),
"assigned_teams": _label(assigned_team_ids, names.teams),
"assigned_keys": _label(record.assigned_key_ids or (), names.keys),
}
)
return AccessGroupResponse.model_validate(payload)
def _ids_across(
records: Sequence[_AccessGroupRecord], pick: Callable[[_AccessGroupRecord], Sequence[str] | None]
) -> tuple[str, ...]:
return tuple(dict.fromkeys(resource_id for record in records for resource_id in (pick(record) or ())))
async def _responses_for(
prisma_client: PrismaClient, records: Sequence[_AccessGroupRecord]
) -> tuple[AccessGroupResponse, ...]:
if not records:
return ()
teams: Final = await _teams_touching(TeamRepository(prisma_client).table, records)
mcp_servers, agents, keys = await asyncio.gather(
mcp_server_display_names(
prisma_client,
_ids_across(records, lambda record: record.access_mcp_server_ids),
global_mcp_server_manager.config_mcp_servers,
),
agent_display_names(
prisma_client, _ids_across(records, lambda record: record.access_agent_ids), global_agent_registry
),
key_display_names(prisma_client, _ids_across(records, lambda record: record.assigned_key_ids)),
)
names: Final = _ResourceNames(
mcp_servers=mcp_servers,
agents=agents,
teams=MappingProxyType({team.team_id: team.team_alias for team in teams}),
keys=keys,
)
attached: Final = _attached_team_ids_by_group(records, teams)
return tuple(
_record_to_response(record, assigned_team_ids=attached[record.access_group_id], names=names)
for record in records
)
async def _response_for(prisma_client: PrismaClient, record: _AccessGroupRecord) -> AccessGroupResponse:
(response,) = await _responses_for(prisma_client, (record,))
return response
def _attached_team_ids_by_group(
records: Sequence[_AccessGroupRecord], teams: Sequence[_TeamRecord]
) -> Mapping[str, tuple[str, ...]]:
@ -144,19 +222,21 @@ def _attached_team_ids_by_group(
return MappingProxyType({record.access_group_id: attached(record) for record in records})
async def _teams_touching(team_table: _TeamTable, records: Sequence[_AccessGroupRecord]) -> Sequence[_TeamRecord]:
"""Team rows listed on any of the groups or carrying any of them in access_group_ids."""
group_ids: Final = tuple(record.access_group_id for record in records)
stored_team_ids: Final = _ids_across(records, lambda record: record.assigned_team_ids)
carrying: Final = {"access_group_ids": {"hasSome": group_ids}} # mutable-ok: prisma where is a dict
listed: Final = {"team_id": {"in": stored_team_ids}} # mutable-ok: prisma where is a dict
return await team_table.find_many(where={"OR": (carrying, listed)}) # mutable-ok: prisma where is a dict
async def _attached_team_ids_for(
team_table: _TeamTable, records: Sequence[_AccessGroupRecord]
) -> Mapping[str, tuple[str, ...]]:
if not records:
return MappingProxyType({})
group_ids: Final = tuple(record.access_group_id for record in records)
stored_team_ids: Final = tuple(
dict.fromkeys(team_id for record in records for team_id in (record.assigned_team_ids or ()))
)
carrying: Final = {"access_group_ids": {"hasSome": group_ids}} # mutable-ok: prisma where is a dict
listed: Final = {"team_id": {"in": stored_team_ids}} # mutable-ok: prisma where is a dict
teams: Final = await team_table.find_many(where={"OR": (carrying, listed)}) # mutable-ok: prisma where is a dict
return _attached_team_ids_by_group(records, teams)
return _attached_team_ids_by_group(records, await _teams_touching(team_table, records))
async def _require_teams_exist(tx: _AccessGroupTx, team_ids: Sequence[str]) -> None:
@ -425,7 +505,7 @@ async def create_access_group(
proxy_logging_obj,
)
return _record_to_response(record)
return await _response_for(prisma_client, record)
@router.get(
@ -434,14 +514,13 @@ async def create_access_group(
)
async def list_access_groups(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> list[AccessGroupResponse]:
) -> Sequence[AccessGroupResponse]:
_require_admin_view(user_api_key_dict)
prisma_client: Final = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value)
table: Final = AccessGroupRepository(prisma_client).table
records: Final = await table.find_many(order={"created_at": "desc"})
attached: Final = await _attached_team_ids_for(TeamRepository(prisma_client).table, records)
return [_record_to_response(r, assigned_team_ids=attached[r.access_group_id]) for r in records]
return await _responses_for(prisma_client, records)
@router.get(
@ -462,8 +541,7 @@ async def get_access_group(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Access group '{access_group_id}' not found",
)
attached: Final = await _attached_team_ids_for(TeamRepository(prisma_client).table, (record,))
return _record_to_response(record, assigned_team_ids=attached[record.access_group_id])
return await _response_for(prisma_client, record)
@router.put(
@ -560,7 +638,7 @@ async def update_access_group(
await _patch_key_caches_add_access_group(keys_to_add, access_group_id, user_api_key_cache, proxy_logging_obj)
await _patch_key_caches_remove_access_group(keys_to_remove, access_group_id, user_api_key_cache, proxy_logging_obj)
return _record_to_response(record)
return await _response_for(prisma_client, record)
@router.delete(

View file

@ -3326,9 +3326,6 @@ async def team_member_delete(
data=data,
)
if not removed_team_members:
raise HTTPException(status_code=400, detail={"error": "User not found in team"})
existing_team_row.members_with_roles = new_team_members
_db_new_team_members: Final[list[dict]] = [m.model_dump() for m in new_team_members]
@ -3336,17 +3333,27 @@ async def team_member_delete(
## DELETE TEAM ID from USER ROW, IF EXISTS ##
# get user row
removed_user_ids: Final = frozenset(m.user_id for m in removed_team_members if m.user_id is not None)
addressed_user_ids: Final = (
removed_user_ids if removed_team_members else frozenset((data.user_id,) if data.user_id is not None else ())
)
key_val: Final[Mapping[str, object]] = (
{"user_id": {"in": sorted(removed_user_ids)}} if removed_user_ids else {"user_email": data.user_email}
{"user_id": {"in": sorted(addressed_user_ids)}} if addressed_user_ids else {"user_email": data.user_email}
)
member_tx: Final[_MemberDeleteTx] = tx
existing_user_rows: Final = await member_tx.litellm_usertable.find_many(where=key_val)
# Also clean up any existing team membership rows for this user and team
user_ids_to_delete: Final = removed_user_ids.union(
(data.user_id,) if data.user_id is not None else (),
(user.user_id for user in existing_user_rows if user.user_id),
)
# A user row can outlive its roster entry, and until the team is off user.teams the user
# still sees it and still fails key creation against it, so removal has to clear it too
stale_user_rows: Final = tuple(user for user in existing_user_rows if data.team_id in user.teams)
# Also clean up any existing team membership rows for this user and team. An email can
# match several user rows, so with no roster entry to name the member, only the rows
# actually carrying the team are the ones this request is allowed to touch
cleanup_user_rows: Final = existing_user_rows if removed_team_members else stale_user_rows
user_ids_to_delete: Final = addressed_user_ids.union(user.user_id for user in cleanup_user_rows if user.user_id)
if not removed_team_members and not stale_user_rows:
raise HTTPException(status_code=400, detail={"error": "User not found in team"})
## DELETE KEYS CREATED BY USER FOR THIS TEAM
# Fetch keys before deletion so their audit records can be persisted alongside the delete.
@ -3358,17 +3365,17 @@ async def team_member_delete(
}
)
await _team_tx_db(tx).update(
where={"team_id": data.team_id},
data={"members_with_roles": json.dumps(_db_new_team_members)},
)
if removed_team_members:
await _team_tx_db(tx).update(
where={"team_id": data.team_id},
data={"members_with_roles": json.dumps(_db_new_team_members)},
)
for existing_user in existing_user_rows:
if data.team_id in existing_user.teams:
await tx.litellm_usertable.update(
where={"user_id": existing_user.user_id},
data={"teams": {"set": [team for team in existing_user.teams if team != data.team_id]}},
)
for existing_user in stale_user_rows:
await tx.litellm_usertable.update(
where={"user_id": existing_user.user_id},
data={"teams": {"set": [team for team in existing_user.teams if team != data.team_id]}},
)
for _uid in sorted(user_ids_to_delete):
await tx.litellm_teammembership.delete_many(where={"team_id": data.team_id, "user_id": _uid})

View file

@ -0,0 +1,61 @@
"""Display names for ids stored on management objects. DB rows win; config-declared servers and agents fill the gaps."""
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import Final
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
from litellm.proxy.utils import PrismaClient
from litellm.repositories.table_repositories import AgentsRepository, MCPServerRepository
from litellm.repositories.verification_token_repository import VerificationTokenRepository
from litellm.types.mcp_server.mcp_server_manager import MCPServer
async def mcp_server_display_names(
prisma_client: PrismaClient,
server_ids: Sequence[str],
config_servers: Mapping[str, MCPServer],
) -> Mapping[str, str]:
"""server_id -> alias, falling back to server_name; config-only servers also fall back to their registry name."""
if not server_ids:
return MappingProxyType({})
wanted: Final = frozenset(server_ids)
where: Final = {"server_id": {"in": tuple(wanted)}} # mutable-ok: prisma where is a dict
rows: Final = await MCPServerRepository(prisma_client).table.find_many(where=where)
from_config: Final = {
server_id: server.alias or server.server_name or server.name
for server_id, server in config_servers.items()
if server_id in wanted
}
from_db: Final = {row.server_id: name for row in rows if (name := row.alias or row.server_name)}
return MappingProxyType({**from_config, **from_db})
async def agent_display_names(
prisma_client: PrismaClient,
agent_ids: Sequence[str],
registry: AgentRegistry,
) -> Mapping[str, str]:
"""agent_id -> agent_name. The registry covers config-declared agents and their legacy ids."""
if not agent_ids:
return MappingProxyType({})
wanted: Final = frozenset(agent_ids)
where: Final = {"agent_id": {"in": tuple(wanted)}} # mutable-ok: prisma where is a dict
rows: Final = await AgentsRepository(prisma_client).table.find_many(where=where)
from_registry: Final = {
alias_id: agent.agent_name
for agent in registry.get_agent_list()
for alias_id in registry.ids_for_agent(agent.agent_id)
if alias_id in wanted
}
from_db: Final = {row.agent_id: row.agent_name for row in rows}
return MappingProxyType({**from_registry, **from_db})
async def key_display_names(prisma_client: PrismaClient, tokens: Sequence[str]) -> Mapping[str, str]:
"""token hash -> key_alias for the keys that have one."""
if not tokens:
return MappingProxyType({})
where: Final = {"token": {"in": tuple(frozenset(tokens))}} # mutable-ok: prisma where is a dict
rows: Final = await VerificationTokenRepository(prisma_client).table.find_many(where=where)
return MappingProxyType({row.token: row.key_alias for row in rows if row.key_alias})

View file

@ -247,6 +247,58 @@ unless `modality_routing` is also on.
`session_affinity_ttl_seconds` is the idle window for both the model pin selected by session affinity and the deployment pin. Every request that reuses a pin refreshes its TTL, so a session actively sending requests stays pinned. After the window passes with no pin reuse, the next request classifies again and creates a fresh pin. Omit the setting to track the default of 3600 seconds.
### Mid-task stall escalation
A weak model working an agentic task can get stuck: it keeps calling the same tool with the
same arguments, or the same call keeps erroring, when a stronger model would have broken the
loop. `stall_escalation_enabled: true` catches this and bumps the request one tier higher, the
automatic counterpart to a user typing an escalation keyword:
```yaml
model_list:
- model_name: smart-router
litellm_params:
model: auto_router/complexity_router
complexity_router_config:
stall_escalation_enabled: true
stall_escalation_window: 6
stall_escalation_repeat_threshold: 3
tiers:
SIMPLE: gpt-4o-mini
MEDIUM: gpt-4o
COMPLEX: claude-sonnet-4
REASONING: o1-preview
```
Detection looks at the assistant's own tool calls, not the human's messages. The task counts as
stalled when the NEWEST tool call is still part of a stuck pattern: it repeats, or it errored, at
least `stall_escalation_repeat_threshold` times across the last `stall_escalation_window` calls.
The tier is then bumped one step by the same `_escalate_tier` ladder `escalation_keywords` uses,
capped at the highest configured tier. It reads both tool-call shapes: Anthropic Messages
`tool_use`/`tool_result` blocks (including `is_error`) and chat-completions `tool_calls`/`tool`
messages (which carry no standard error flag, so those calls are judged on repetition alone).
Anchoring on the newest call is what keeps a recovered task from being escalated on stale
evidence. A model that tried the same command three times and then moved on still has those
three calls sitting in the window for a few turns, and counting whichever pattern is most common
in the window would escalate a request that is already making progress again. Anchoring still
leaves room between the matches, so a retry loop broken up by an unrelated lookup counts.
There is no state to expire or leak: detection reruns on every classified turn from that
request's own message list, so the bump lasts only as long as the recent tool calls still look
stuck and lifts on its own the moment they don't. This also means it reads the whole
conversation rather than only the turns since the newest human ask, so a plain follow-up like
"try again" does not discard evidence from before it. Escalation records `stall_escalation` in
`routing_decision.signals`; unlike `escalation_keywords`, it does not set the
`escalated`/`escalation_keyword` pair, which is reserved for the keyword mechanism specifically.
`stall_escalation_enabled` cannot be combined with `session_affinity` or
`classification_mode: user_turn`: both replay a held routing decision on most turns instead of
classifying, so detection would never see the tool calls it needs to look at. It is also
rejected together with `tier_definitions`, for the same reason `escalation_keywords` is: both
rely on the built-in tier severity order, which a custom tier set does not define. Off by
default.
### Heuristic-first chaining
`classifier_type: heuristic_first` runs the local scorer on every request and only calls the LLM

View file

@ -40,7 +40,10 @@ from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
)
from litellm.litellm_core_utils.internal_call_metadata import forwarded_internal_call_metadata
from litellm.litellm_core_utils.prompt_templates.common_utils import request_contains_image_content
from litellm.litellm_core_utils.prompt_templates.common_utils import (
as_openai_image_part,
request_contains_image_content,
)
from litellm.litellm_core_utils.sensitive_data_masker import mask_credentials_in_payload
from litellm.llms.base_llm.base_utils import type_to_response_format_param
from litellm.router_strategy.adaptive_router.classifier import classify_prompt
@ -48,7 +51,11 @@ from litellm.router_strategy.complexity_router.tier_predictor import (
TierSuccessPredictor,
resolve_tier_artifact,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionImageObject,
ChatCompletionTextObject,
)
from litellm.types.utils import (
AUTOROUTER_CLASSIFIER_CALL_ORIGIN,
ModelResponse,
@ -76,6 +83,7 @@ from .config import (
ComplexityTier,
TierDefinition,
)
from .stall_detector import detect_stalled_task
if TYPE_CHECKING:
from semantic_router.routers import SemanticRouter
@ -434,6 +442,23 @@ def _strip_reminder_blocks(text: str, marker_pairs: tuple[tuple[str, str], ...]
return " ".join(kept for a, b in zip(keep_from, keep_to) if (kept := text[a:b].strip()))
def _inline_image_part(part: Mapping[str, object]) -> ChatCompletionImageObject | None:
"""One image content part safe to hand the classifier, or None.
Inline data URIs only. A remote URL is caller-controlled and provider adapters do not uniformly
delegate fetching to the provider: gigachat's file handler downloads any non-data URL with
`client.get` from the proxy host, so forwarding one would let a key scoped to this router aim a
proxy-side request at an internal address, on a call the caller never asked for. The routed
model still receives the original URL exactly as before.
"""
converted: Final = as_openai_image_part(part)
if converted is None:
return None
image_url: Final = converted["image_url"]
url: Final = image_url if isinstance(image_url, str) else image_url.get("url", "")
return converted if url.startswith("data:") else None
def _human_text(content: object, marker_pairs: tuple[tuple[str, str], ...] = _DEFAULT_REMINDER_MARKERS) -> str:
"""Message content as the text a human wrote, with complete reminder blocks removed.
@ -1591,6 +1616,10 @@ class ComplexityRouter(CustomLogger):
threshold check alone would hand that traffic to the cheapest model without ever consulting
the classifier. Scores also go negative when simple indicators fire, so a score threshold
would reject exactly the trivial prompts this path exists to serve.
A turn carrying images the classifier would see is never decided cheaply: the scorer reads
text alone, so its confidence describes a request it has only partly seen, and a trivial
caption beside a screenshot is exactly the misrouting vision classification exists to stop.
"""
tier, score, signals, cause = self._score_and_classify(prompt, system_prompt)
scored: Final = ClassificationOutcome(tier=tier, score=score, signals=signals, cause=cause)
@ -1598,6 +1627,7 @@ class ComplexityRouter(CustomLogger):
decided_cheaply: Final = (
threshold is not None
and bool(signals)
and not self._classifier_image_parts(messages)
and self._active_tier_severity(tier) <= self._active_tier_severity(threshold)
)
if decided_cheaply:
@ -1622,11 +1652,43 @@ class ComplexityRouter(CustomLogger):
tier, score, signals, cause = self._score_and_classify(prompt, system_prompt)
scored: Final = ClassificationOutcome(tier=tier, score=score, signals=signals, cause=cause)
margin: Final = self.config.hybrid_boundary_margin
decided: Final = margin is not None and bool(signals) and not self._is_near_tier_boundary(score, margin)
decided: Final = (
margin is not None
and bool(signals)
and not self._classifier_image_parts(messages)
and not self._is_near_tier_boundary(score, margin)
)
if decided:
return ClassificationOutcome(tier=tier, score=score, signals=signals, cause="hybrid_short_circuit")
return await self._llm_classifier_outcome(prompt, system_prompt, request_kwargs, messages, scored=scored)
def _classifier_image_parts(
self, messages: Sequence[Mapping[str, object]] | None
) -> tuple[ChatCompletionImageObject, ...]:
"""Images from the newest user turn to hand the classifier, capped by max_images.
Empty unless the operator opted in AND the classifier model is declared vision-capable, so
every other deployment keeps today's text-only payload byte for byte. Only the newest user
turn is read: earlier turns are context the classifier already gets as quoted text, and an
image nested in a tool_result is tool output rather than the ask being classified.
Remote-URL images are left out entirely; `_inline_image_part` carries why.
"""
llm_config: Final = self.config.classifier_llm_config
if llm_config is None or not llm_config.vision.enabled or not self.config.uses_llm_classifier or not messages:
return ()
if not self._model_declares_vision_support(llm_config.model):
return ()
newest_user_turn: Final = next((msg for msg in reversed(messages) if msg.get("role") == "user"), None)
content: Final = newest_user_turn.get("content") if newest_user_turn is not None else None
if not isinstance(content, list):
return ()
return tuple(
islice(
(part for raw in content if isinstance(raw, Mapping) and (part := _inline_image_part(raw)) is not None),
llm_config.vision.max_images,
)
)
async def _llm_classifier_outcome(
self,
prompt: str,
@ -1864,9 +1926,18 @@ class ComplexityRouter(CustomLogger):
}
turn_off_message_logging: Final = _effective_turn_off_message_logging(request_kwargs)
image_parts: Final = self._classifier_image_parts(messages)
user_content: Final[str | Sequence[ChatCompletionTextObject | ChatCompletionImageObject]] = (
[ # mutable-ok: SDK request payload content list is built once
{"type": "text", "text": user_payload},
*image_parts,
]
if image_parts
else user_payload
)
messages_for_call: Final[list[AllMessageValues]] = [ # mutable-ok: SDK request payload list is built once
{"role": "system", "content": classifier_system_prompt},
{"role": "user", "content": user_payload},
{"role": "user", "content": user_content},
]
response_format: Final = classifier_response_format
classifier_call_params: Mapping[str, str] = EMPTY_MAPPING
@ -2557,31 +2628,53 @@ class ComplexityRouter(CustomLogger):
return pinned_model
return self.get_model_for_tier(escalated_tier)
def _model_accepts_image_input(self, model_name: str) -> bool:
"""Whether a routed model or pool entry can serve an image request.
def _vision_verdicts(self, model_name: str) -> tuple[bool | None, ...]:
"""Declared vision support per deployment serving the name: True, False, or None when
nothing declares either way.
Resolved through the deployments that would actually serve the name; a name with no
deployment on the router is served by the SDK directly and is checked against the model
cost map itself. Only an explicit supports_vision false excludes, a deployment-level
model_info override first and the map otherwise, so unmapped custom names stay routable.
cost map itself. A deployment-level model_info override wins over the map.
One verdict set, two readings, because the two callers fail in opposite directions.
Routing a user's image asks whether anything RULES IT OUT, so an undeclared model stays
eligible and unmapped custom names keep routing. Handing an image to the classifier asks
whether something RULES IT IN: an undeclared model that turns out to be text-only rejects
every image request, and that rejection is swallowed by the classifier's own fallback, so
the router quietly serves all image traffic from the fallback tier and pays for the failed
call each time. An undeclared model instead keeps today's text-only payload, which is a
visible no-op the operator fixes by declaring supports_vision on the deployment.
"""
from litellm.utils import is_vision_explicitly_disabled, supports_vision
def model_verdict(model: str) -> bool | None:
if supports_vision(model):
return True
return False if is_vision_explicitly_disabled(model) else None
def deployment_verdict(deployment: Mapping[str, Any]) -> bool | None:
declared: Final = (deployment.get("model_info") or EMPTY_MAPPING).get("supports_vision")
if declared is not None:
return declared is True
return model_verdict((deployment.get("litellm_params") or EMPTY_MAPPING).get("model") or model_name)
deployments: Final = self.litellm_router_instance.get_model_list(model_name=model_name)
if not deployments:
return (model_verdict(model_name),)
return tuple(deployment_verdict(deployment) for deployment in deployments)
def _model_accepts_image_input(self, model_name: str) -> bool:
"""Whether a routed model or pool entry can serve an image request.
A multi-deployment group must accept on EVERY deployment: the router picks a deployment
inside the group after this gate runs, so a mixed group marked eligible could still hand
the image to its text-only member and fail with the exact 400 the gate exists to prevent.
"""
from litellm.utils import is_vision_explicitly_disabled
return all(verdict is not False for verdict in self._vision_verdicts(model_name))
def deployment_accepts(deployment: Mapping[str, Any]) -> bool:
declared: Final = (deployment.get("model_info") or EMPTY_MAPPING).get("supports_vision")
if declared is not None:
return declared is True
litellm_model: Final = (deployment.get("litellm_params") or EMPTY_MAPPING).get("model") or model_name
return not is_vision_explicitly_disabled(litellm_model)
deployments: Final = self.litellm_router_instance.get_model_list(model_name=model_name)
if not deployments:
return not is_vision_explicitly_disabled(model_name)
return all(deployment_accepts(deployment) for deployment in deployments)
def _model_declares_vision_support(self, model_name: str) -> bool:
"""Whether every deployment serving the name is declared vision-capable."""
return all(verdict is True for verdict in self._vision_verdicts(model_name))
def _modality_eligible_models(self) -> frozenset[str]:
"""Every configured pool entry, plus default_model, that can serve an image request."""
@ -3373,8 +3466,9 @@ class ComplexityRouter(CustomLogger):
has_original_messages: Final = messages is not None and len(messages) > 0
user_message, system_prompt = _extract_current_ask_and_system_prompt(resolved_messages, self._reminder_markers)
classifier_images: Final = self._classifier_image_parts(resolved_messages)
if user_message is None:
if user_message is None and not classifier_images:
verbose_router_logger.debug("ComplexityRouter: No user message found, routing to default model")
default_model_first: Final = not self.config.plugins and self.config.default_model
if default_model_first:
@ -3401,8 +3495,17 @@ class ComplexityRouter(CustomLogger):
),
)
ask: Final = user_message or ""
newest_ask: Final = _newest_turn_ask(resolved_messages, self._reminder_markers)
escalation_keyword: Final = self._matched_escalation_keyword(newest_ask) if newest_ask is not None else None
# Resolved here rather than beside the classifier because the keyword-override path below
# returns before any classification runs, and a forced tier gets stuck for the same reason
# a classified one does.
stalled: Final = self.config.stall_escalation_enabled and detect_stalled_task(
resolved_messages,
window=self.config.stall_escalation_window,
repeat_threshold=self.config.stall_escalation_repeat_threshold,
)
plan_mode_sentinel: Final = self._matched_plan_mode_signal(request_kwargs, resolved_messages)
plan_floor: Final = self._resolve_plan_mode_floor() if plan_mode_sentinel is not None else None
@ -3430,12 +3533,13 @@ class ComplexityRouter(CustomLogger):
),
)
override: Final = await self._resolve_keyword_tier_override(user_message, request_kwargs)
override: Final = await self._resolve_keyword_tier_override(ask, request_kwargs)
if override is not None:
escalated_tier: Final = (
keyword_bumped_tier: Final = (
self._escalate_tier(override.tier) if escalation_keyword is not None else override.tier
)
keyword_escalated: Final = escalated_tier != override.tier
escalated_tier: Final = self._escalate_tier(keyword_bumped_tier) if stalled else keyword_bumped_tier
keyword_escalated: Final = keyword_bumped_tier != override.tier
routed_tier: Final = (
self._apply_plan_mode_floor(escalated_tier) if plan_floor is not None else escalated_tier
)
@ -3463,6 +3567,7 @@ class ComplexityRouter(CustomLogger):
conversation_continuing=conversation_continuing,
cause=keyword_cause,
tier=routed_tier,
signals=("stall_escalation",) if stalled else None,
matched_keyword=plan_mode_sentinel if keyword_plan_floored else override.matched_keyword,
escalation_keyword=escalation_keyword,
escalated=keyword_escalated,
@ -3475,9 +3580,7 @@ class ComplexityRouter(CustomLogger):
outcome: Final = (
ClassificationOutcome(tier=housekeeping_tier, score=None, signals=("housekeeping",), cause="housekeeping")
if housekeeping_tier is not None
else await self.aclassify(
user_message, system_prompt, request_kwargs, resolved_messages, raw_messages=messages
)
else await self.aclassify(ask, system_prompt, request_kwargs, resolved_messages, raw_messages=messages)
)
tier, score, signals = outcome.tier, outcome.score, outcome.signals
classified_tier: Final = tier
@ -3486,6 +3589,9 @@ class ComplexityRouter(CustomLogger):
escalated: Final = tier != classified_tier
if escalated:
signals = (*signals, "escalation")
if stalled:
tier = self._escalate_tier(tier)
signals = (*signals, "stall_escalation")
pre_floor_tier: Final = tier
if plan_floor is not None:
tier = self._apply_plan_mode_floor(tier)
@ -3544,7 +3650,7 @@ class ComplexityRouter(CustomLogger):
# under is not a floor.
routed_model = self._soft_floor_pick(
tier,
user_message,
ask,
request_kwargs,
hard_floor=tier if context_original_tier is not None else plan_floor,
hard_ceiling=housekeeping_ceiling,

View file

@ -442,12 +442,47 @@ DEFAULT_TIER_MODELS: Final[dict[str, str]] = {
}
class ClassifierVisionConfig(BaseModel):
"""Whether the LLM classifier sees the images on the request it is classifying.
Off by default because images cost far more than the text ask they arrive with, and the
classifier runs on every request. A turn whose complexity lives in the image ("what is wrong in
this stack trace screenshot") is invisible to a text-only classifier, which is what this buys.
"""
enabled: bool = Field(
default=False,
description=(
"Forward image content to the classifier. Requires a classifier model declared "
"supports_vision, on the deployment's model_info or in the model cost map; images stay "
"stripped otherwise, so a classifier that cannot read them is never sent one. Declare "
"model_info.supports_vision on the deployment to enable a model the cost map does not "
"describe. Only inline data: URIs are forwarded. A request whose images are http(s) "
"URLs still classifies on its text alone, because some providers fetch such a URL from "
"the proxy rather than the provider, which would let a caller aim a proxy-side request "
"at an address of their choosing."
),
)
max_images: int = Field(
default=1,
ge=1,
description=(
"How many images from the newest user turn to forward, in wire order. Bounds the added "
"cost of a turn that attaches many images. Images on earlier turns are never forwarded."
),
)
class ClassifierLLMConfig(BaseModel):
"""Configuration for the LLM-based complexity classifier."""
model: str = Field(
description="Model name (from the router's model_list) to call for classification",
)
vision: ClassifierVisionConfig = Field(
default_factory=ClassifierVisionConfig,
description="Whether the classifier sees images on the request, and how many",
)
reasoning_effort: REASONING_EFFORT | None = Field(
default=None,
description=(
@ -852,6 +887,43 @@ class ComplexityRouterConfig(BaseModel):
description="Rules that force a specific tier when their keywords match the prompt",
)
stall_escalation_enabled: bool = Field(
default=False,
description=(
"Escalate mid-task to the next-higher configured tier when the assistant's own recent "
"tool calls look stuck: the newest tool call repeats, or errors, at least "
"stall_escalation_repeat_threshold times across the last stall_escalation_window "
"calls. Both tests are anchored on the newest call, so a task that tried the same "
"thing a few times and then moved on is not escalated on the strength of those older "
"calls alone, while a retry loop broken up by an unrelated lookup still counts. One "
"tier at most, on the same ladder escalation_keywords bumps along, and never above "
"the highest configured tier. Detection re-runs on every classified turn from the "
"tool calls visible in that request, so it needs no state and nothing survives past "
"the task. Mutually exclusive with session_affinity and classification_mode="
"'user_turn', which both replay a held routing decision instead of classifying most "
"turns, so this would never see the tool calls to look at. Off by default."
),
)
stall_escalation_window: int = Field(
default=6,
gt=0,
description=(
"How many of the assistant's most recent tool calls stall detection looks at, oldest "
"ones dropped as new calls happen. Counted across the whole visible conversation "
"rather than reset at the newest human ask, so evidence from before a plain follow-up "
"message like 'try again' is still visible on the turn after it."
),
)
stall_escalation_repeat_threshold: int = Field(
default=3,
ge=2,
description=(
"How many of the last stall_escalation_window tool calls must repeat the newest call, "
"or must have errored alongside it, before the task counts as stalled. Must not "
"exceed stall_escalation_window, or the condition could never be reached."
),
)
plan_mode_min_tier: str | None = Field(
default=None,
description=(
@ -1323,6 +1395,7 @@ class ComplexityRouterConfig(BaseModel):
("adaptive", self.adaptive),
("session_affinity", self.session_affinity),
("escalation_keywords", bool(self.escalation_keywords)),
("stall_escalation_enabled", self.stall_escalation_enabled),
("plugins", bool(self.plugins)),
)
if enabled
@ -1490,6 +1563,25 @@ class ComplexityRouterConfig(BaseModel):
)
return self
@model_validator(mode="after")
def _validate_stall_escalation(self) -> "ComplexityRouterConfig":
if not self.stall_escalation_enabled:
return self
if self.session_affinity or self.classification_mode == "user_turn":
raise ValueError(
"stall_escalation_enabled cannot be combined with session_affinity or "
"classification_mode='user_turn': both replay a held routing decision on most "
"turns instead of classifying, so stall detection would never see the tool calls "
"of the turns it needs to look at. Disable one or the other."
)
if self.stall_escalation_repeat_threshold > self.stall_escalation_window:
raise ValueError(
"stall_escalation_repeat_threshold "
f"({self.stall_escalation_repeat_threshold}) cannot exceed stall_escalation_window "
f"({self.stall_escalation_window}); the condition could never be reached."
)
return self
@model_validator(mode="after")
def _validate_tier_param_placement(self) -> "ComplexityRouterConfig":
"""Reject a router setting written into a tier entry's request params.

View file

@ -0,0 +1,118 @@
"""
Mid-task stall detection for the Complexity Router.
Reads the assistant's own recent tool calls, which every agentic client resends on each
turn, and reports whether the task currently looks stuck. No LLM call and no stored state:
the same window is rescanned per classified turn, so the verdict follows the conversation
rather than latching.
Tool calls arrive in two shapes and are read in place rather than translated:
- Anthropic Messages: assistant `tool_use` content blocks, answered by a user-turn
`tool_result` block carrying `is_error`
- Chat completions: assistant `tool_calls` entries, answered by a `role: "tool"` message,
which has no standard error flag, so those calls are judged on repetition alone
"""
from __future__ import annotations
import json
from collections.abc import Iterator, Mapping, Sequence
from itertools import islice
from typing import Final, NamedTuple
_ARGUMENTS_PARSE_FAILED: Final = object()
class _ToolCallEvent(NamedTuple):
signature: tuple[str, str]
is_error: bool | None
"""None where the surface reports no error status, and never counted as an error."""
def _json_arguments(raw: str) -> object:
try:
return json.loads(raw)
except (TypeError, ValueError):
return _ARGUMENTS_PARSE_FAILED
def _tool_call_signature(name: str, raw_arguments: object) -> tuple[str, str]:
"""Canonicalized so the same call compares equal across both surfaces, which carry
arguments as a dict and as a JSON string respectively."""
parsed: Final = _json_arguments(raw_arguments) if isinstance(raw_arguments, str) else raw_arguments
arguments: Final = raw_arguments if parsed is _ARGUMENTS_PARSE_FAILED else parsed
try:
return name, json.dumps(arguments, sort_keys=True, default=str)
except (TypeError, ValueError):
return name, str(arguments)
def _iter_tool_result_error_pairs(messages: Sequence[Mapping[str, object]]) -> Iterator[tuple[str, bool]]:
for msg in messages:
content = msg.get("content")
if msg.get("role") != "user" or not isinstance(content, list):
continue
for part in content:
if isinstance(part, Mapping) and part.get("type") == "tool_result":
call_id = part.get("tool_use_id")
if isinstance(call_id, str):
yield call_id, bool(part.get("is_error", False))
def _iter_tool_call_events_newest_first(messages: Sequence[Mapping[str, object]]) -> Iterator[_ToolCallEvent]:
error_by_call_id: Final = dict(_iter_tool_result_error_pairs(messages))
for msg in reversed(messages):
if msg.get("role") != "assistant":
continue
content = msg.get("content")
if isinstance(content, list):
for part in reversed(content):
if not (isinstance(part, Mapping) and part.get("type") == "tool_use"):
continue
name = part.get("name")
if isinstance(name, str):
call_id = part.get("id")
yield _ToolCallEvent(
signature=_tool_call_signature(name, part.get("input")),
is_error=error_by_call_id.get(call_id) if isinstance(call_id, str) else None,
)
tool_calls = msg.get("tool_calls")
if not isinstance(tool_calls, list):
continue
for call in reversed(tool_calls):
function = call.get("function") if isinstance(call, Mapping) else None
name = function.get("name") if isinstance(function, Mapping) else None
if isinstance(name, str):
yield _ToolCallEvent(
signature=_tool_call_signature(name, function.get("arguments") if function else None),
is_error=None,
)
def detect_stalled_task(
messages: Sequence[Mapping[str, object]] | None,
*,
window: int,
repeat_threshold: int,
) -> bool:
"""Whether the newest tool call is still part of a stuck pattern: it repeats, or it
errored, at least repeat_threshold times across the last `window` calls.
Both tests are anchored on the newest call rather than counting whichever pattern is
most common in the window. A task that tried the same thing three times and then moved
on has those three calls in the window for a while yet, and counting them alone would
escalate a request that already recovered. Anchoring also leaves room between the
matches, so a retry loop broken up by an unrelated lookup still reads as stuck.
"""
if not messages or repeat_threshold <= 0:
return False
recent: Final = tuple(islice(_iter_tool_call_events_newest_first(messages), window))
if len(recent) < repeat_threshold:
return False
newest: Final = recent[0]
repeats: Final = sum(1 for event in recent if event.signature == newest.signature)
if repeats >= repeat_threshold:
return True
if not newest.is_error:
return False
return sum(1 for event in recent if event.is_error) >= repeat_threshold

View file

@ -23,6 +23,13 @@ class AccessGroupUpdateRequest(BaseModel):
assigned_key_ids: list[str] | None = None
class AccessGroupResource(BaseModel):
"""A resource referenced by an access group. `name` is null when the id no longer resolves or has no alias."""
id: str
name: str | None
class AccessGroupResponse(BaseModel):
access_group_id: str
access_group_name: str
@ -32,6 +39,10 @@ class AccessGroupResponse(BaseModel):
access_agent_ids: list[str]
assigned_team_ids: list[str]
assigned_key_ids: list[str]
access_mcp_servers: tuple[AccessGroupResource, ...]
access_agents: tuple[AccessGroupResource, ...]
assigned_teams: tuple[AccessGroupResource, ...]
assigned_keys: tuple[AccessGroupResource, ...]
created_at: datetime
created_by: str | None = None
updated_at: datetime

View file

@ -3096,7 +3096,7 @@ PROMPT_CARRYING_GUARDRAIL_FIELDS: Final[frozenset[str]] = frozenset(
# The rest of the record: what the guardrail is, what it decided, how long it took and what it cost.
# None of these reproduce the prompt, so a redacted record keeps them and stays explainable.
# `test_every_guardrail_field_is_classified` fails if a field is added to the record without being
# `test_a_redacted_span_carries_every_declared_guardrail_field` fails if a field is added to the record without being
# placed in one set or the other, so a new field is dropped from redacted records rather than
# shipped unexamined.
AUDIT_GUARDRAIL_FIELDS: Final[frozenset[str]] = frozenset(
@ -3120,6 +3120,7 @@ AUDIT_GUARDRAIL_FIELDS: Final[frozenset[str]] = frozenset(
"guardrail_action",
"guardrail_usage",
"guardrail_cost",
"guardrail_cost_by_unit",
"guardrail_cost_in_spend",
}
)

View file

@ -7157,6 +7157,53 @@
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": false
},
"azure/gpt-6-astra": {
"cache_creation_input_token_cost": 1.25e-05,
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-05,
"cache_read_input_token_cost": 1e-06,
"cache_read_input_token_cost_above_272k_tokens": 2e-06,
"input_cost_per_token": 1e-05,
"input_cost_per_token_above_272k_tokens": 2e-05,
"litellm_provider": "azure",
"max_input_tokens": 922000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 5e-05,
"output_cost_per_token_above_272k_tokens": 7.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_computer_use": true,
"supports_function_calling": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": false,
"supports_native_streaming": true,
"supports_none_reasoning_effort": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true,
"supports_xhigh_reasoning_effort": true
},
"azure/us/gpt-5.6": {
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_272k_tokens": 1.375e-05,
@ -7376,6 +7423,53 @@
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": false
},
"azure/us/gpt-6-astra": {
"cache_creation_input_token_cost": 1.375e-05,
"cache_creation_input_token_cost_above_272k_tokens": 2.75e-05,
"cache_read_input_token_cost": 1.1e-06,
"cache_read_input_token_cost_above_272k_tokens": 2.2e-06,
"input_cost_per_token": 1.1e-05,
"input_cost_per_token_above_272k_tokens": 2.2e-05,
"litellm_provider": "azure",
"max_input_tokens": 922000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 5.5e-05,
"output_cost_per_token_above_272k_tokens": 8.25e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_computer_use": true,
"supports_function_calling": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": false,
"supports_native_streaming": true,
"supports_none_reasoning_effort": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true,
"supports_xhigh_reasoning_effort": true
},
"azure/eu/gpt-5.6": {
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_272k_tokens": 1.375e-05,

View file

@ -2008,6 +2008,49 @@ def test_generic_cost_per_token_azure_gpt56(_local_model_cost_map,
assert round(completion_cost, 10) == round(output_cost * completion_tokens, 10)
@pytest.mark.parametrize("model,zone_multiplier", [("azure/gpt-6-astra", 1.0), ("azure/us/gpt-6-astra", 1.1)])
@pytest.mark.parametrize(
"prompt_tokens,input_side_multiplier,output_multiplier",
[(100000, 1.0, 1.0), (300000, 2.0, 1.5)],
)
def test_generic_cost_per_token_azure_gpt_6_astra_foundry_price_sheet(
_local_model_cost_map,
model,
zone_multiplier,
prompt_tokens,
input_side_multiplier,
output_multiplier,
):
"""Microsoft Foundry sells gpt-6-astra at the OpenAI rates: $10 input, $1 cache read, $12.50 cache write,
$50 output per 1M tokens on Standard Global, with the input side doubling and output 1.5x above 272K
prompt tokens. Standard US Data Zone carries the usual 10% uplift on every rate.
"""
cached_tokens = 50000
cache_write_tokens = 40000
text_tokens = prompt_tokens - cached_tokens - cache_write_tokens
completion_tokens = 1000
usage = Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
prompt_tokens_details=PromptTokensDetailsWrapper(
cached_tokens=cached_tokens, cache_write_tokens=cache_write_tokens
),
)
prompt_cost, completion_cost = generic_cost_per_token(
model=model,
usage=usage,
custom_llm_provider="azure",
)
input_side = zone_multiplier * input_side_multiplier
assert prompt_cost == pytest.approx(
input_side * (text_tokens * 1e-5 + cached_tokens * 1e-6 + cache_write_tokens * 1.25e-5)
)
assert completion_cost == pytest.approx(zone_multiplier * output_multiplier * completion_tokens * 5e-5)
@pytest.mark.parametrize(
"model,expected_none,expected_xhigh,expected_minimal",
[

View file

@ -348,3 +348,31 @@ def test_azure_gpt_6_astra_takes_the_reasoning_series_request_shape():
assert params["max_completion_tokens"] == 100
assert "max_tokens" not in params
assert params["reasoning_effort"] == "max"
@pytest.mark.parametrize("model", ["azure/gpt-6-astra", "azure/us/gpt-6-astra"])
def test_azure_gpt6_astra_reasoning_effort_none_unlocks_temperature(config: AzureOpenAIGPT5Config, model: str):
"""Foundry's gpt-6-astra accepts reasoning_effort='none' and, only then, a non-default
temperature (verified live against a Foundry deployment), unlike OpenAI's gpt-6-astra."""
params = config.map_openai_params(
non_default_params={"temperature": 0.2, "reasoning_effort": "none"},
optional_params={},
model=model,
drop_params=False,
api_version="2025-04-01-preview",
)
assert params["temperature"] == 0.2
assert params["reasoning_effort"] == "none"
@pytest.mark.parametrize("model", ["azure/gpt-6-astra", "azure/us/gpt-6-astra"])
def test_azure_gpt6_astra_rejects_reasoning_effort_minimal(config: AzureOpenAIGPT5Config, model: str):
"""Foundry's gpt-6-astra lists none, low, medium, high, xhigh and max but not minimal."""
with pytest.raises(litellm.utils.UnsupportedParamsError):
config.map_openai_params(
non_default_params={"reasoning_effort": "minimal"},
optional_params={},
model=model,
drop_params=False,
api_version="2025-04-01-preview",
)

View file

@ -1,11 +1,10 @@
from copy import deepcopy
from unittest.mock import patch
from unittest.mock import MagicMock, patch
import pytest
from unittest.mock import MagicMock
import litellm
from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map
from litellm.llms.azure.responses.o_series_transformation import (
AzureOpenAIOSeriesResponsesAPIConfig,
)
@ -613,3 +612,39 @@ class TestAzureResponsesAPIConfig:
assert result["tools"][0] is tool
assert "anyOf" in result["tools"][0]["parameters"]
@pytest.fixture()
def local_model_cost_map(monkeypatch: pytest.MonkeyPatch) -> None:
"""Pin the bundled cost map: the published map lags a key added in this repo."""
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", get_model_cost_map(url=litellm.model_cost_map_url))
litellm.add_known_models(model_cost_map=litellm.model_cost)
def test_azure_responses_gpt6_astra_reasoning_effort_none_unlocks_temperature(local_model_cost_map: None):
"""Foundry's gpt-6-astra accepts reasoning.effort='none' with a non-default temperature
while OpenAI's gpt-6-astra does not, so the gate must read the azure/ cost-map entry
for the bare deployment name rather than OpenAI's."""
params = AzureOpenAIResponsesAPIConfig().map_openai_params(
response_api_optional_params=ResponsesAPIOptionalRequestParams(
temperature=0.2,
reasoning={"effort": "none"},
),
model="gpt-6-astra",
drop_params=False,
)
assert params["temperature"] == 0.2
assert params["reasoning"] == {"effort": "none"}
def test_azure_responses_gpt6_astra_rejects_temperature_while_reasoning(local_model_cost_map: None):
with pytest.raises(litellm.UnsupportedParamsError):
AzureOpenAIResponsesAPIConfig().map_openai_params(
response_api_optional_params=ResponsesAPIOptionalRequestParams(
temperature=0.2,
reasoning={"effort": "low"},
),
model="gpt-6-astra",
drop_params=False,
)

View file

@ -57,8 +57,20 @@ def _make_access_group_record(
return record
def _make_team_record(team_id: str, access_group_ids: list[str] | None = None):
return types.SimpleNamespace(team_id=team_id, access_group_ids=access_group_ids or [])
def _make_team_record(team_id: str, access_group_ids: list[str] | None = None, team_alias: str | None = None):
return types.SimpleNamespace(team_id=team_id, access_group_ids=access_group_ids or [], team_alias=team_alias)
def _make_mcp_server_record(server_id: str, alias: str | None = None, server_name: str | None = None):
return types.SimpleNamespace(server_id=server_id, alias=alias, server_name=server_name)
def _make_agent_record(agent_id: str, agent_name: str):
return types.SimpleNamespace(agent_id=agent_id, agent_name=agent_name)
def _make_key_record(token: str, key_alias: str | None = None):
return types.SimpleNamespace(token=token, key_alias=key_alias)
@pytest.fixture
@ -109,6 +121,12 @@ def client_and_mocks(monkeypatch):
mock_key_table.find_unique = AsyncMock(return_value=None)
mock_key_table.update = AsyncMock(return_value=None)
mock_mcp_server_table = MagicMock()
mock_mcp_server_table.find_many = AsyncMock(return_value=[])
mock_agents_table = MagicMock()
mock_agents_table.find_many = AsyncMock(return_value=[])
@asynccontextmanager
async def mock_tx():
tx = types.SimpleNamespace(
@ -122,6 +140,8 @@ def client_and_mocks(monkeypatch):
litellm_accessgrouptable=mock_access_group_table,
litellm_teamtable=mock_team_table,
litellm_verificationtoken=mock_key_table,
litellm_mcpservertable=mock_mcp_server_table,
litellm_agentstable=mock_agents_table,
tx=mock_tx,
)
mock_prisma.db = mock_db
@ -1447,3 +1467,169 @@ def test_update_access_group_null_assigned_ids_treated_as_empty(client_and_mocks
update_call_kwargs = mock_table.update.call_args.kwargs
assert update_call_kwargs["data"]["assigned_team_ids"] == []
assert update_call_kwargs["data"]["assigned_key_ids"] == []
# ---------------------------------------------------------------------------
# Resolved resource names (LIT-6594)
# ---------------------------------------------------------------------------
def _mock_resource_tables(mock_prisma, *, mcp_servers=(), agents=(), teams=(), keys=()):
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=list(mcp_servers))
mock_prisma.db.litellm_agentstable.find_many = AsyncMock(return_value=list(agents))
mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=list(teams))
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=list(keys))
@pytest.mark.parametrize("base_path", ACCESS_GROUP_PATHS)
def test_get_access_group_resolves_resource_names(client_and_mocks, base_path):
"""Every id list gets a sibling list of {id, name}; name is null when the id has no alias or no longer resolves."""
client, mock_prisma, mock_table, *_ = client_and_mocks
mock_table.find_unique = AsyncMock(
return_value=_make_access_group_record(
access_group_id="ag-123",
access_mcp_server_ids=["mcp-a", "mcp-b", "mcp-ghost"],
access_agent_ids=["agent-a", "agent-ghost"],
assigned_team_ids=["team-a", "team-b"],
assigned_key_ids=["key-a", "key-b"],
)
)
_mock_resource_tables(
mock_prisma,
mcp_servers=[
_make_mcp_server_record("mcp-a", alias="GitHub"),
_make_mcp_server_record("mcp-b", server_name="jira_tools"),
],
agents=[_make_agent_record("agent-a", "support-bot")],
teams=[
_make_team_record("team-a", ["ag-123"], team_alias="Platform"),
_make_team_record("team-b", ["ag-123"]),
],
keys=[_make_key_record("key-a", key_alias="ci-key"), _make_key_record("key-b")],
)
resp = client.get(f"{base_path}/ag-123")
assert resp.status_code == 200
body = resp.json()
assert body["access_mcp_servers"] == [
{"id": "mcp-a", "name": "GitHub"},
{"id": "mcp-b", "name": "jira_tools"},
{"id": "mcp-ghost", "name": None},
]
assert body["access_agents"] == [{"id": "agent-a", "name": "support-bot"}, {"id": "agent-ghost", "name": None}]
assert body["assigned_teams"] == [{"id": "team-a", "name": "Platform"}, {"id": "team-b", "name": None}]
assert body["assigned_keys"] == [{"id": "key-a", "name": "ci-key"}, {"id": "key-b", "name": None}]
assert body["access_mcp_server_ids"] == ["mcp-a", "mcp-b", "mcp-ghost"]
assert body["assigned_team_ids"] == ["team-a", "team-b"]
mcp_where = mock_prisma.db.litellm_mcpservertable.find_many.call_args.kwargs["where"]
assert sorted(mcp_where["server_id"]["in"]) == ["mcp-a", "mcp-b", "mcp-ghost"]
agent_where = mock_prisma.db.litellm_agentstable.find_many.call_args.kwargs["where"]
assert sorted(agent_where["agent_id"]["in"]) == ["agent-a", "agent-ghost"]
key_where = mock_prisma.db.litellm_verificationtoken.find_many.call_args.kwargs["where"]
assert sorted(key_where["token"]["in"]) == ["key-a", "key-b"]
def test_list_access_groups_resolves_names_with_one_query_per_table(client_and_mocks):
"""List batches every group's ids into one lookup per table and attributes names back to the right group."""
client, mock_prisma, mock_table, *_ = client_and_mocks
mock_table.find_many = AsyncMock(
return_value=[
_make_access_group_record(
access_group_id="ag-1", access_mcp_server_ids=["mcp-a"], access_agent_ids=["agent-a"], assigned_key_ids=["key-a"]
),
_make_access_group_record(
access_group_id="ag-2", access_mcp_server_ids=["mcp-b"], access_agent_ids=["agent-b"], assigned_key_ids=["key-b"]
),
]
)
_mock_resource_tables(
mock_prisma,
mcp_servers=[_make_mcp_server_record("mcp-a", alias="A"), _make_mcp_server_record("mcp-b", alias="B")],
agents=[_make_agent_record("agent-a", "Agent A"), _make_agent_record("agent-b", "Agent B")],
keys=[_make_key_record("key-a", key_alias="Key A"), _make_key_record("key-b", key_alias="Key B")],
)
resp = client.get("/v1/access_group")
assert resp.status_code == 200
first, second = resp.json()
assert first["access_mcp_servers"] == [{"id": "mcp-a", "name": "A"}]
assert first["access_agents"] == [{"id": "agent-a", "name": "Agent A"}]
assert first["assigned_keys"] == [{"id": "key-a", "name": "Key A"}]
assert second["access_mcp_servers"] == [{"id": "mcp-b", "name": "B"}]
assert second["access_agents"] == [{"id": "agent-b", "name": "Agent B"}]
assert second["assigned_keys"] == [{"id": "key-b", "name": "Key B"}]
for table, column in (
(mock_prisma.db.litellm_mcpservertable, "server_id"),
(mock_prisma.db.litellm_agentstable, "agent_id"),
(mock_prisma.db.litellm_verificationtoken, "token"),
):
table.find_many.assert_awaited_once()
assert len(table.find_many.call_args.kwargs["where"][column]["in"]) == 2
def test_list_access_groups_skips_lookups_when_nothing_to_resolve(client_and_mocks):
"""Groups with no MCP servers, agents or keys must not trigger an empty IN () query per table."""
client, mock_prisma, mock_table, *_ = client_and_mocks
mock_table.find_many = AsyncMock(
return_value=[_make_access_group_record(access_group_id="ag-1"), _make_access_group_record(access_group_id="ag-2")]
)
resp = client.get("/v1/access_group")
assert resp.status_code == 200
assert all(group["access_mcp_servers"] == [] and group["assigned_keys"] == [] for group in resp.json())
mock_prisma.db.litellm_mcpservertable.find_many.assert_not_awaited()
mock_prisma.db.litellm_agentstable.find_many.assert_not_awaited()
mock_prisma.db.litellm_verificationtoken.find_many.assert_not_awaited()
def test_create_access_group_response_carries_resolved_names(client_and_mocks):
"""The create response already shows names so the UI never has to refetch to label what it just saved."""
client, mock_prisma, *_ = client_and_mocks
team_record = _make_team_record("team-1", team_alias="Platform")
mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_record)
_mock_resource_tables(
mock_prisma,
mcp_servers=[_make_mcp_server_record("mcp-a", alias="GitHub")],
agents=[_make_agent_record("agent-a", "support-bot")],
teams=[team_record],
)
resp = client.post(
"/v1/access_group",
json={
"access_group_name": "new-group",
"access_mcp_server_ids": ["mcp-a"],
"access_agent_ids": ["agent-a"],
"assigned_team_ids": ["team-1"],
},
)
assert resp.status_code == 201
body = resp.json()
assert body["access_mcp_servers"] == [{"id": "mcp-a", "name": "GitHub"}]
assert body["access_agents"] == [{"id": "agent-a", "name": "support-bot"}]
assert body["assigned_teams"] == [{"id": "team-1", "name": "Platform"}]
def test_update_access_group_response_carries_resolved_names(client_and_mocks):
"""The update response reflects the new ids with their names, not the pre-update state."""
client, mock_prisma, mock_table, *_ = client_and_mocks
mock_table.find_unique = AsyncMock(
return_value=_make_access_group_record(access_group_id="ag-update", access_mcp_server_ids=["mcp-old"])
)
_mock_resource_tables(
mock_prisma,
mcp_servers=[_make_mcp_server_record("mcp-new", alias="Linear")],
agents=[_make_agent_record("agent-a", "support-bot")],
)
resp = client.put(
"/v1/access_group/ag-update", json={"access_mcp_server_ids": ["mcp-new"], "access_agent_ids": ["agent-a"]}
)
assert resp.status_code == 200
body = resp.json()
assert body["access_mcp_servers"] == [{"id": "mcp-new", "name": "Linear"}]
assert body["access_agents"] == [{"id": "agent-a", "name": "support-bot"}]
assert body["access_mcp_server_ids"] == ["mcp-new"]

View file

@ -4863,6 +4863,302 @@ async def test_team_member_delete_by_email_the_user_row_does_not_carry(
)
@pytest.mark.asyncio
async def test_team_member_delete_clears_team_left_on_the_user_row_without_a_roster_entry(
mock_db_client, mock_admin_auth
):
"""
A user row can keep a team (several times over, from older duplicate-prone adds) after the
roster entry is gone, which leaves the team listed on the user, offered in the key creation
dropdown, and rejected by key creation itself. Reporting "User not found in team" left that
residue unremovable, so the delete now cleans every copy of the team off the user row.
"""
from litellm.proxy._types import TeamMemberDeleteRequest
from litellm.proxy.management_endpoints.team_endpoints import team_member_delete
test_team_id = "team-del-orphan-123"
test_user_id = "user-del-orphan-123"
mock_team_row = MagicMock()
mock_team_row.model_dump.return_value = {
"team_id": test_team_id,
"members_with_roles": [],
"team_member_permissions": [],
"metadata": {},
"models": [],
"spend": 0.0,
}
mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(
return_value=mock_team_row
)
mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_team_row)
mock_user_row = MagicMock()
mock_user_row.user_id = test_user_id
mock_user_row.user_email = None
mock_user_row.teams = [test_team_id, "other-team", test_team_id]
mock_db_client.db.litellm_usertable.find_many = AsyncMock(
return_value=[mock_user_row]
)
mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock())
mock_db_client.db.litellm_teammembership = MagicMock()
mock_db_client.db.litellm_teammembership.delete_many = AsyncMock(
return_value=MagicMock()
)
mock_db_client.db.litellm_verificationtoken = MagicMock()
mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_db_client.db.litellm_verificationtoken.delete_many = AsyncMock(
return_value=MagicMock()
)
_wire_member_delete_tx(mock_db_client)
await team_member_delete(
data=TeamMemberDeleteRequest(team_id=test_team_id, user_id=test_user_id),
user_api_key_dict=mock_admin_auth,
)
mock_db_client.db.litellm_usertable.update.assert_awaited_once_with(
where={"user_id": test_user_id},
data={"teams": {"set": ["other-team"]}},
)
mock_db_client.db.litellm_teammembership.delete_many.assert_awaited_once_with(
where={"team_id": test_team_id, "user_id": test_user_id}
)
@pytest.mark.asyncio
async def test_team_member_delete_still_rejects_a_user_the_team_has_no_trace_of(
mock_db_client, mock_admin_auth
):
from litellm.proxy._types import TeamMemberDeleteRequest
from litellm.proxy.management_endpoints.team_endpoints import team_member_delete
test_team_id = "team-del-absent-123"
test_user_id = "user-del-absent-123"
mock_team_row = MagicMock()
mock_team_row.model_dump.return_value = {
"team_id": test_team_id,
"members_with_roles": [],
"team_member_permissions": [],
"metadata": {},
"models": [],
"spend": 0.0,
}
mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(
return_value=mock_team_row
)
mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_team_row)
mock_user_row = MagicMock()
mock_user_row.user_id = test_user_id
mock_user_row.user_email = None
mock_user_row.teams = ["other-team"]
mock_db_client.db.litellm_usertable.find_many = AsyncMock(
return_value=[mock_user_row]
)
mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock())
mock_db_client.db.litellm_teammembership = MagicMock()
mock_db_client.db.litellm_teammembership.delete_many = AsyncMock(
return_value=MagicMock()
)
_wire_member_delete_tx(mock_db_client)
with pytest.raises(HTTPException) as exc_info:
await team_member_delete(
data=TeamMemberDeleteRequest(team_id=test_team_id, user_id=test_user_id),
user_api_key_dict=mock_admin_auth,
)
assert exc_info.value.status_code == 400
assert exc_info.value.detail == {"error": "User not found in team"}
mock_db_client.db.litellm_usertable.update.assert_not_awaited()
mock_db_client.db.litellm_teammembership.delete_many.assert_not_awaited()
@pytest.mark.asyncio
async def test_team_member_delete_leaves_a_bystander_named_by_a_conflicting_user_id_alone(
mock_db_client, mock_admin_auth
):
"""
A request can carry a user_id and a user_email that point at two different people, and only the
email matches a roster entry. Cleaning up both ids would strip the team, the membership row and
the keys off the bystander the roster never listed, so the user_id only widens the cleanup when
the roster came back empty.
"""
from litellm.proxy._types import TeamMemberDeleteRequest
from litellm.proxy.management_endpoints.team_endpoints import team_member_delete
test_team_id = "team-del-conflict-123"
roster_user_id = "user-del-conflict-roster"
bystander_user_id = "user-del-conflict-bystander"
roster_email = "roster@example.com"
mock_team_row = MagicMock()
mock_team_row.model_dump.return_value = {
"team_id": test_team_id,
"members_with_roles": [
{"user_id": roster_user_id, "user_email": roster_email, "role": "user"}
],
"team_member_permissions": [],
"metadata": {},
"models": [],
"spend": 0.0,
}
mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(
return_value=mock_team_row
)
mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_team_row)
roster_user_row = MagicMock()
roster_user_row.user_id = roster_user_id
roster_user_row.user_email = roster_email
roster_user_row.teams = [test_team_id]
bystander_user_row = MagicMock()
bystander_user_row.user_id = bystander_user_id
bystander_user_row.user_email = "bystander@example.com"
bystander_user_row.teams = [test_team_id]
rows_by_user_id = {
roster_user_id: roster_user_row,
bystander_user_id: bystander_user_row,
}
async def find_user_rows(where):
user_id_filter = where.get("user_id")
if isinstance(user_id_filter, dict):
return [
rows_by_user_id[uid]
for uid in user_id_filter.get("in", [])
if uid in rows_by_user_id
]
return [
row
for row in rows_by_user_id.values()
if row.user_email == where.get("user_email")
]
mock_db_client.db.litellm_usertable.find_many = AsyncMock(
side_effect=find_user_rows
)
mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock())
mock_db_client.db.litellm_teammembership = MagicMock()
mock_db_client.db.litellm_teammembership.delete_many = AsyncMock(
return_value=MagicMock()
)
mock_db_client.db.litellm_verificationtoken = MagicMock()
mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_db_client.db.litellm_verificationtoken.delete_many = AsyncMock(
return_value=MagicMock()
)
_wire_member_delete_tx(mock_db_client)
await team_member_delete(
data=TeamMemberDeleteRequest(
team_id=test_team_id,
user_id=bystander_user_id,
user_email=roster_email,
),
user_api_key_dict=mock_admin_auth,
)
mock_db_client.db.litellm_usertable.update.assert_awaited_once_with(
where={"user_id": roster_user_id},
data={"teams": {"set": []}},
)
mock_db_client.db.litellm_teammembership.delete_many.assert_awaited_once_with(
where={"team_id": test_team_id, "user_id": roster_user_id}
)
mock_db_client.db.litellm_verificationtoken.delete_many.assert_awaited_once_with(
where={"user_id": {"in": [roster_user_id]}, "team_id": test_team_id}
)
@pytest.mark.asyncio
async def test_team_member_delete_by_email_only_touches_the_row_carrying_the_stale_team(
mock_db_client, mock_admin_auth
):
"""
user_email is not unique, so an email delete against an empty roster can match several user
rows. Only the row that actually carries the team is stale; the namesake keeps its team, its
membership row and its keys.
"""
from litellm.proxy._types import TeamMemberDeleteRequest
from litellm.proxy.management_endpoints.team_endpoints import team_member_delete
test_team_id = "team-del-shared-email-123"
stale_user_id = "user-del-shared-email-stale"
namesake_user_id = "user-del-shared-email-namesake"
shared_email = "shared@example.com"
mock_team_row = MagicMock()
mock_team_row.model_dump.return_value = {
"team_id": test_team_id,
"members_with_roles": [],
"team_member_permissions": [],
"metadata": {},
"models": [],
"spend": 0.0,
}
mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(
return_value=mock_team_row
)
mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_team_row)
stale_user_row = MagicMock()
stale_user_row.user_id = stale_user_id
stale_user_row.user_email = shared_email
stale_user_row.teams = [test_team_id]
namesake_user_row = MagicMock()
namesake_user_row.user_id = namesake_user_id
namesake_user_row.user_email = shared_email
namesake_user_row.teams = ["other-team"]
mock_db_client.db.litellm_usertable.find_many = AsyncMock(
return_value=[stale_user_row, namesake_user_row]
)
mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock())
mock_db_client.db.litellm_teammembership = MagicMock()
mock_db_client.db.litellm_teammembership.delete_many = AsyncMock(
return_value=MagicMock()
)
mock_db_client.db.litellm_verificationtoken = MagicMock()
mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_db_client.db.litellm_verificationtoken.delete_many = AsyncMock(
return_value=MagicMock()
)
_wire_member_delete_tx(mock_db_client)
await team_member_delete(
data=TeamMemberDeleteRequest(team_id=test_team_id, user_email=shared_email),
user_api_key_dict=mock_admin_auth,
)
mock_db_client.db.litellm_usertable.update.assert_awaited_once_with(
where={"user_id": stale_user_id},
data={"teams": {"set": []}},
)
mock_db_client.db.litellm_teammembership.delete_many.assert_awaited_once_with(
where={"team_id": test_team_id, "user_id": stale_user_id}
)
mock_db_client.db.litellm_verificationtoken.delete_many.assert_awaited_once_with(
where={"user_id": {"in": [stale_user_id]}, "team_id": test_team_id}
)
class _InjectedMemberDeleteFailure(Exception):
pass

View file

@ -0,0 +1,130 @@
import types
from types import MappingProxyType
from unittest.mock import AsyncMock
import pytest
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
from litellm.proxy.management_helpers.resource_display_names import (
agent_display_names,
key_display_names,
mcp_server_display_names,
)
from litellm.types.agents import AgentResponse
from litellm.types.mcp_server.mcp_server_manager import MCPServer
def _table(rows=()):
return types.SimpleNamespace(find_many=AsyncMock(return_value=list(rows)))
def _prisma(**tables):
return types.SimpleNamespace(db=types.SimpleNamespace(**tables))
def _config_server(server_id: str, name: str, alias: str | None = None, server_name: str | None = None) -> MCPServer:
return MCPServer(server_id=server_id, name=name, alias=alias, server_name=server_name, transport="http")
def _registry_with(*agents: AgentResponse, legacy_ids: dict[str, str] | None = None) -> AgentRegistry:
registry = AgentRegistry()
for agent in agents:
registry.register_agent(agent)
registry.config_agent_legacy_ids = MappingProxyType(legacy_ids or {})
return registry
def _agent(agent_id: str, agent_name: str) -> AgentResponse:
return AgentResponse(agent_id=agent_id, agent_name=agent_name, agent_card_params={})
@pytest.mark.asyncio
async def test_mcp_db_row_beats_config_entry_for_the_same_server():
"""The DB is authoritative when both sources know a server; the registry may lag behind a rename on another pod."""
prisma = _prisma(
litellm_mcpservertable=_table([types.SimpleNamespace(server_id="s1", alias="db-alias", server_name=None)])
)
names = await mcp_server_display_names(prisma, ("s1",), {"s1": _config_server("s1", "config-name")})
assert dict(names) == {"s1": "db-alias"}
@pytest.mark.asyncio
@pytest.mark.parametrize(
("alias", "server_name", "expected"),
[("Alias", "server_name", "Alias"), (None, "server_name", "server_name"), (None, None, "config-name")],
)
async def test_mcp_config_only_server_falls_back_alias_then_server_name_then_name(alias, server_name, expected):
"""Config-declared servers have no DB row, so their registry entry supplies the label."""
prisma = _prisma(litellm_mcpservertable=_table())
config = {"s1": _config_server("s1", "config-name", alias=alias, server_name=server_name)}
names = await mcp_server_display_names(prisma, ("s1",), config)
assert dict(names) == {"s1": expected}
@pytest.mark.asyncio
async def test_mcp_db_row_without_alias_or_server_name_yields_no_label():
"""A bare DB row must not produce an empty string label; the caller falls back to the id."""
prisma = _prisma(
litellm_mcpservertable=_table([types.SimpleNamespace(server_id="s1", alias=None, server_name=None)])
)
assert dict(await mcp_server_display_names(prisma, ("s1",), {})) == {}
@pytest.mark.asyncio
async def test_mcp_only_requested_ids_are_returned_and_the_query_is_deduped():
"""Unrequested config servers stay out of the result and repeated ids collapse to one IN filter entry."""
table = _table([types.SimpleNamespace(server_id="s1", alias="A", server_name=None)])
prisma = _prisma(litellm_mcpservertable=table)
config = {"other": _config_server("other", "not-requested")}
names = await mcp_server_display_names(prisma, ("s1", "s1", "missing"), config)
assert dict(names) == {"s1": "A"}
assert sorted(table.find_many.call_args.kwargs["where"]["server_id"]["in"]) == ["missing", "s1"]
@pytest.mark.asyncio
async def test_mcp_empty_ids_skip_the_db():
table = _table()
names = await mcp_server_display_names(_prisma(litellm_mcpservertable=table), (), {})
assert dict(names) == {}
table.find_many.assert_not_awaited()
@pytest.mark.asyncio
async def test_agent_db_name_beats_registry_name():
prisma = _prisma(litellm_agentstable=_table([types.SimpleNamespace(agent_id="a1", agent_name="from-db")]))
registry = _registry_with(_agent("a1", "from-registry"))
assert dict(await agent_display_names(prisma, ("a1",), registry)) == {"a1": "from-db"}
@pytest.mark.asyncio
async def test_agent_legacy_config_id_resolves_to_the_stable_agent_name():
"""Access groups saved before agent ids were stabilised still carry the legacy hash; it must still get a name."""
prisma = _prisma(litellm_agentstable=_table())
registry = _registry_with(_agent("stable-id", "config-agent"), legacy_ids={"legacy-id": "stable-id"})
names = await agent_display_names(prisma, ("legacy-id", "stable-id", "unknown"), registry)
assert dict(names) == {"legacy-id": "config-agent", "stable-id": "config-agent"}
@pytest.mark.asyncio
async def test_agent_empty_ids_skip_the_db():
table = _table()
names = await agent_display_names(_prisma(litellm_agentstable=table), (), _registry_with())
assert dict(names) == {}
table.find_many.assert_not_awaited()
@pytest.mark.asyncio
async def test_key_alias_only_for_keys_that_have_one():
table = _table(
[types.SimpleNamespace(token="k1", key_alias="ci-key"), types.SimpleNamespace(token="k2", key_alias=None)]
)
names = await key_display_names(_prisma(litellm_verificationtoken=table), ("k1", "k2", "k1"))
assert dict(names) == {"k1": "ci-key"}
assert sorted(table.find_many.call_args.kwargs["where"]["token"]["in"]) == ["k1", "k2"]
@pytest.mark.asyncio
async def test_key_empty_ids_skip_the_db():
table = _table()
assert dict(await key_display_names(_prisma(litellm_verificationtoken=table), ())) == {}
table.find_many.assert_not_awaited()

View file

@ -2897,9 +2897,7 @@ class TestRouterPreRoutingAliasOverrides:
import time
monkeypatch.setenv("GITHUB_COPILOT_TOKEN_DIR", str(tmp_path))
(tmp_path / "api-key.json").write_text(
json.dumps({"token": "tid=test", "expires_at": int(time.time()) + 3600})
)
(tmp_path / "api-key.json").write_text(json.dumps({"token": "tid=test", "expires_at": int(time.time()) + 3600}))
router = Router(
model_list=[
{
@ -2924,7 +2922,9 @@ class TestRouterPreRoutingAliasOverrides:
copilot_resolutions: List = []
def _guarded(*args, **kwargs):
target = str(kwargs.get("model") or (args[0] if args else "")) + str(kwargs.get("custom_llm_provider") or "")
target = str(kwargs.get("model") or (args[0] if args else "")) + str(
kwargs.get("custom_llm_provider") or ""
)
if "github_copilot" in target:
copilot_resolutions.append(target)
raise RuntimeError("routing must not resolve an authenticating provider")
@ -6133,6 +6133,150 @@ class TestEscalationKeywords:
assert result.model == "o1-b" # unchanged: no random hop to o1-a / o1-c
def _stalled_tool_history(repeats: int = 3) -> List[Dict]:
"""`repeats` identical bash tool calls in a row, the automatic counterpart to a user
typing an escalation keyword: the assistant, not the human, is the one stuck."""
return [
turn
for i in range(repeats)
for turn in (
{
"role": "assistant",
"content": [{"type": "tool_use", "id": f"call-{i}", "name": "bash", "input": {"cmd": "pytest"}}],
},
{
"role": "user",
"content": [{"type": "tool_result", "tool_use_id": f"call-{i}", "is_error": True, "content": "fail"}],
},
)
]
class TestStallEscalation:
"""Mid-task auto-escalation when the assistant's own recent tool calls look stuck: the
automatic counterpart to escalation_keywords, gated by stall_escalation_enabled and off
by default."""
@pytest.mark.asyncio
async def test_repeated_tool_calls_escalate_the_classified_tier(self, mock_router_instance, basic_config):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={**basic_config, "stall_escalation_enabled": True},
)
messages = [*_stalled_tool_history(), {"role": "user", "content": "Hello there!"}]
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
assert result.model == "gpt-4o" # SIMPLE bumped to MEDIUM
@pytest.mark.asyncio
async def test_varied_tool_calls_do_not_escalate(self, mock_router_instance, basic_config):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={**basic_config, "stall_escalation_enabled": True},
)
messages = [
{
"role": "assistant",
"content": [{"type": "tool_use", "id": "c1", "name": "bash", "input": {"cmd": "ls"}}],
},
{
"role": "user",
"content": [{"type": "tool_result", "tool_use_id": "c1", "is_error": False, "content": "ok"}],
},
{"role": "user", "content": "Hello there!"},
]
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
assert result.model == "gpt-4o-mini" # not escalated
@pytest.mark.asyncio
async def test_disabled_by_default_ignores_repeated_tool_calls(self, mock_router_instance, basic_config):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config=basic_config,
)
messages = [*_stalled_tool_history(), {"role": "user", "content": "Hello there!"}]
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
assert result.model == "gpt-4o-mini" # stall_escalation_enabled defaults False
@pytest.mark.asyncio
async def test_signals_record_stall_escalation(self, mock_router_instance, basic_config):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={**basic_config, "stall_escalation_enabled": True},
)
messages = [*_stalled_tool_history(), {"role": "user", "content": "Hello there!"}]
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
assert "stall_escalation" in result.routing_decision["signals"]
@pytest.mark.asyncio
async def test_stall_escalation_caps_at_highest_tier(self, mock_router_instance, basic_config):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={**basic_config, "stall_escalation_enabled": True},
)
messages = [
*_stalled_tool_history(),
{"role": "user", "content": "Let's think step by step and reason through this carefully."},
]
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
assert result.model == "o1-preview" # already REASONING, stays there
@pytest.mark.asyncio
async def test_stall_escalation_stacks_with_keyword_escalation(self, mock_router_instance, basic_config):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={**basic_config, "stall_escalation_enabled": True},
)
messages = [*_stalled_tool_history(), {"role": "user", "content": "LITELLM ESCALATE Hello there!"}]
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
assert result.model == "claude-sonnet-4-20250514" # SIMPLE -> MEDIUM (keyword) -> COMPLEX (stall)
@pytest.mark.asyncio
async def test_a_keyword_forced_tier_still_escalates_when_stalled(self, mock_router_instance, basic_config):
"""A keyword rule forces its tier and returns before any classification runs, so
without its own bump the one path that can pin a weak model to a whole conversation
would be the one path a stall could never lift."""
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
**basic_config,
"stall_escalation_enabled": True,
"keyword_tier_rules": [{"keywords": ["billing"], "tier": "SIMPLE"}],
},
)
healthy = await router.async_pre_routing_hook(
model="test-model", request_kwargs={}, messages=[{"role": "user", "content": "a billing question"}]
)
assert healthy.model == "gpt-4o-mini" # forced SIMPLE, nothing stuck
stalled = await router.async_pre_routing_hook(
model="test-model",
request_kwargs={},
messages=[*_stalled_tool_history(), {"role": "user", "content": "a billing question"}],
)
assert stalled.model == "gpt-4o" # forced SIMPLE bumped to MEDIUM
assert "stall_escalation" in stalled.routing_decision["signals"]
@pytest.mark.asyncio
async def test_evidence_survives_a_new_human_ask(self, mock_router_instance, basic_config):
"""A plain follow-up like 'try again' must not erase the stall evidence that came
before it: escalation still fires on the turn carrying that follow-up."""
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={**basic_config, "stall_escalation_enabled": True},
)
messages = [*_stalled_tool_history(), {"role": "user", "content": "try again"}]
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
assert result.model == "gpt-4o" # SIMPLE ("try again" carries no signal) bumped to MEDIUM
class TestRoutingDecisionContents:
"""Every routing path must return a PreRoutingHookResponse carrying a routing_decision
that names the mechanism that actually decided, with the facts of that path only."""
@ -8027,7 +8171,6 @@ class TestClientHousekeepingCalls:
assert result is not None
assert result.model == "claude-sonnet-4-20250514"
@pytest.mark.asyncio
async def test_a_classifier_plugin_still_decides_its_own_routers(self, mock_router_instance):
"""A plugin is where an operator encodes policy the tier ladder cannot express.
@ -8062,9 +8205,7 @@ class TestClientHousekeepingCalls:
assert result.model == "o1-preview"
assert result.routing_decision["cause"] == "classifier_plugin"
def _adaptive_router(
self, tier_distance_penalty: float, plan_mode_min_tier: str | None = None
) -> ComplexityRouter:
def _adaptive_router(self, tier_distance_penalty: float, plan_mode_min_tier: str | None = None) -> ComplexityRouter:
adaptive_instance = MagicMock()
adaptive_instance.model_list = [
{
@ -8101,9 +8242,7 @@ class TestClientHousekeepingCalls:
return router
@pytest.mark.asyncio
async def test_the_bandit_cannot_route_a_housekeeping_call_above_the_cheapest_tier(
self, mock_router_instance
):
async def test_the_bandit_cannot_route_a_housekeeping_call_above_the_cheapest_tier(self, mock_router_instance):
"""The tier here is what the request IS, not how hard it is, so the bandit has nothing to win.
Without a ceiling the tier distance penalty is the only thing holding the tier, so a
@ -8136,7 +8275,6 @@ class TestClientHousekeepingCalls:
assert result is not None
assert result.model == "premium"
@pytest.mark.asyncio
async def test_a_housekeeping_call_never_becomes_the_session_pin(self, mock_router_instance):
"""Pinning this is the most expensive mistake of the transient causes.
@ -8178,9 +8316,7 @@ class TestClientHousekeepingCalls:
assert work_turn.routing_decision["cause"] == "llm_classifier"
@pytest.mark.asyncio
async def test_the_decision_records_which_sentinel_matched(
self, mock_router_instance, llm_classifier_config
):
async def test_the_decision_records_which_sentinel_matched(self, mock_router_instance, llm_classifier_config):
"""The cause's contract says the sentinel rides in matched_keyword, so it has to be there.
Without it an operator reading the logs can see that a call was treated as housekeeping but
@ -8201,7 +8337,6 @@ class TestClientHousekeepingCalls:
"Write the title in the predominant language of the session"
)
@pytest.mark.asyncio
async def test_the_plan_mode_floor_raises_a_housekeeping_call_under_adaptive(self, mock_router_instance):
"""Floor and ceiling must not contradict each other on the same request.
@ -9428,6 +9563,7 @@ class TestTierDefinitions:
({"adaptive": True}, "severity order"),
({"session_affinity": True}, "severity order"),
({"escalation_keywords": ["GO UP"]}, "severity order"),
({"stall_escalation_enabled": True}, "severity order"),
(
{"classifier_llm_config": {"model": "haiku-classifier", "system_prompt": "grade it"}},
"system_prompt",
@ -10735,9 +10871,7 @@ class TestHeuristicFirst:
# Scores 0.175 with one signal, so it sits 0.025 from simple_medium: the pair of tiers either side of
# that boundary are different model pools, and a hair's difference in score picks the other one.
NEAR_BOUNDARY_PROMPT = (
"design a distributed cache with consistent hashing, then explain the failure modes step by step"
)
NEAR_BOUNDARY_PROMPT = "design a distributed cache with consistent hashing, then explain the failure modes step by step"
# Scores 0.075 with signals, the far side of any margin under 0.075: the scorer is decided here.
CLEAR_OF_BOUNDARY_PROMPT = "explain step by step how consistent hashing rebalances keys"
@ -11179,6 +11313,7 @@ class TestContextWindowEscalation:
litellm_router_instance=_windowed_router(_SMALL, _BIG),
complexity_router_config=_tier_config(session_affinity=True),
)
def session_kwargs() -> dict[str, object]:
return {"metadata": {"session_id": "s-1", "user_api_key_hash": "k-1"}}
@ -11203,6 +11338,7 @@ class TestContextWindowEscalation:
litellm_router_instance=_windowed_router(_SMALL, _BIG),
complexity_router_config=_tier_config(session_affinity=True),
)
def session_kwargs() -> dict[str, object]:
return {"metadata": {"session_id": "s-2", "user_api_key_hash": "k-2"}}
@ -11280,7 +11416,9 @@ class TestContextWindowEscalation:
copilot_resolutions: List = []
def _guarded(*args, **kwargs):
target = str(kwargs.get("model") or (args[0] if args else "")) + str(kwargs.get("custom_llm_provider") or "")
target = str(kwargs.get("model") or (args[0] if args else "")) + str(
kwargs.get("custom_llm_provider") or ""
)
if "github_copilot" in target:
copilot_resolutions.append(target)
raise RuntimeError("the gate must not resolve an authenticating provider")
@ -12291,3 +12429,277 @@ class TestTierHealthFailover:
for _ in range(20)
]
assert {r.model for r in results} == {"live-c"}
ANTHROPIC_IMG_PART = {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "aGk="}}
RESPONSES_IMG_PART = {"type": "input_image", "image_url": "data:image/png;base64,aGk="}
class TestClassifierVision:
"""classifier_llm_config.vision: what the LLM classifier is shown for an image-bearing turn."""
TIERS = {"SIMPLE": "t-simple", "MEDIUM": "t-medium", "COMPLEX": "t-complex", "REASONING": "t-reasoning"}
@staticmethod
def _router(mock_router_instance, *, vision, classifier_declares_vision=True, classifier_type="llm", **extra):
def get_model_list(model_name=None):
if model_name != "clf":
return [{"model_name": model_name, "litellm_params": {"model": "openai/gpt-4o"}}]
declared = classifier_declares_vision
return [
{
"model_name": "clf",
"litellm_params": {"model": "openai/unmapped-classifier"},
"model_info": {} if declared is None else {"supports_vision": declared},
}
]
mock_router_instance.get_model_list = get_model_list
classifier_llm_config = {"model": "clf", "circuit_breaker_enabled": False}
return ComplexityRouter(
model_name="vision-classifier-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"classifier_type": classifier_type,
"classifier_llm_config": (
classifier_llm_config if vision is None else {**classifier_llm_config, "vision": vision}
),
"tiers": dict(TestClassifierVision.TIERS),
**extra,
},
)
@staticmethod
def _classifier_user_content(mock_router_instance):
return mock_router_instance.acompletion.call_args.kwargs["messages"][-1]["content"]
@staticmethod
def _turn(*parts):
return [{"role": "user", "content": list(parts)}]
@pytest.fixture(autouse=True)
def _classifier_answers_complex(self, mock_router_instance):
mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "COMPLEX"}'))
@pytest.mark.asyncio
@pytest.mark.parametrize(
"vision, classifier_declares_vision",
[
(None, True),
({"enabled": False}, True),
({"enabled": True}, False),
({"enabled": True}, None),
],
ids=["vision_unset", "vision_disabled", "classifier_declared_text_only", "classifier_undeclared"],
)
async def test_payload_stays_text_only(self, mock_router_instance, vision, classifier_declares_vision):
"""Off, or a classifier not declared vision-capable, keeps the plain-string payload.
The undeclared case is the polarity. A text-only classifier handed an image rejects the
call, the rejection is swallowed by the classifier's own fallback, and every image request
then serves from the fallback tier while still paying for the failed call. Staying text-only
is instead a visible no-op the operator fixes by declaring supports_vision.
"""
router = self._router(
mock_router_instance, vision=vision, classifier_declares_vision=classifier_declares_vision
)
await router.async_pre_routing_hook(
model="m", request_kwargs={}, messages=self._turn({"type": "text", "text": "what is this"}, IMG_PART)
)
content = self._classifier_user_content(mock_router_instance)
assert isinstance(content, str)
assert "what is this" in content
@pytest.mark.asyncio
async def test_deployment_model_info_enables_a_classifier_the_cost_map_does_not_describe(
self, mock_router_instance
):
"""The escape hatch for an unmapped classifier name, and the reason undeclared can stay off.
`_router` gives every deployment an `openai/unmapped-*` litellm_params model, so nothing in
the cost map declares it and the verdict comes only from model_info.
"""
router = self._router(mock_router_instance, vision={"enabled": True}, classifier_declares_vision=True)
await router.async_pre_routing_hook(
model="m", request_kwargs={}, messages=self._turn({"type": "text", "text": "what is this"}, IMG_PART)
)
assert [b["type"] for b in self._classifier_user_content(mock_router_instance)] == ["text", "image_url"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"part",
[IMG_PART, ANTHROPIC_IMG_PART, RESPONSES_IMG_PART],
ids=["chat_completions", "anthropic_messages", "responses"],
)
async def test_image_reaches_the_classifier_in_chat_completions_dialect(self, mock_router_instance, part):
"""Every surface's dialect arrives as a chat-completions image_url on the classifier call.
/v1/messages hands the hook an Anthropic image block untranslated, so forwarding verbatim
would send the classifier a content part its own request dialect has no meaning for.
"""
router = self._router(mock_router_instance, vision={"enabled": True})
await router.async_pre_routing_hook(
model="m", request_kwargs={}, messages=self._turn({"type": "text", "text": "what is this"}, part)
)
content = self._classifier_user_content(mock_router_instance)
assert [block["type"] for block in content] == ["text", "image_url"]
assert content[1]["image_url"] == {"url": "data:image/png;base64,aGk="}
assert "what is this" in content[0]["text"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"part",
[
{"type": "image_url", "image_url": {"url": "http://169.254.169.254/latest/meta-data/"}},
{"type": "image_url", "image_url": {"url": "https://example.internal/secret.png"}},
{"type": "input_image", "image_url": "https://example.internal/secret.png"},
{"type": "image", "source": {"type": "url", "url": "https://example.internal/secret.png"}},
],
ids=["metadata_service", "chat_completions", "responses", "anthropic"],
)
async def test_remote_url_images_are_never_forwarded(self, mock_router_instance, part):
"""A caller-supplied URL must not reach an internal call the caller did not ask for.
Provider adapters do not uniformly delegate fetching: gigachat downloads any non-data URL
from the proxy host, so forwarding one would turn a router-scoped key into a proxy-side GET
at an address of the caller's choosing.
"""
router = self._router(mock_router_instance, vision={"enabled": True})
await router.async_pre_routing_hook(
model="m", request_kwargs={}, messages=self._turn({"type": "text", "text": "what is this"}, part)
)
assert isinstance(self._classifier_user_content(mock_router_instance), str)
@pytest.mark.asyncio
async def test_remote_url_image_only_turn_does_not_reach_the_classifier(self, mock_router_instance):
"""With nothing forwardable left, the turn stays unclassifiable rather than sending the URL."""
router = self._router(mock_router_instance, vision={"enabled": True})
response = await router.async_pre_routing_hook(
model="m",
request_kwargs={},
messages=self._turn({"type": "image_url", "image_url": {"url": "https://example.internal/x.png"}}),
)
assert response.routing_decision["cause"] == "default_fallback"
mock_router_instance.acompletion.assert_not_awaited()
@pytest.mark.asyncio
async def test_image_only_turn_is_classified_instead_of_falling_back(self, mock_router_instance):
"""A turn carrying only an image reaches the classifier rather than the default model.
It flattens to empty text, so before this it never reached the classifier at all and was
routed as default_fallback on text the request never contained.
"""
router = self._router(mock_router_instance, vision={"enabled": True})
response = await router.async_pre_routing_hook(
model="m", request_kwargs={}, messages=self._turn(IMG_PART)
)
assert response.routing_decision["cause"] == "llm_classifier"
assert response.model == "t-complex"
assert [block["type"] for block in self._classifier_user_content(mock_router_instance)] == [
"text",
"image_url",
]
@pytest.mark.asyncio
async def test_image_only_turn_still_falls_back_when_vision_is_off(self, mock_router_instance):
router = self._router(mock_router_instance, vision={"enabled": False})
response = await router.async_pre_routing_hook(
model="m", request_kwargs={}, messages=self._turn(IMG_PART)
)
assert response.routing_decision["cause"] == "default_fallback"
mock_router_instance.acompletion.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("max_images, expected", [(1, 1), (2, 2), (5, 3)])
async def test_max_images_caps_what_is_forwarded(self, mock_router_instance, max_images, expected):
router = self._router(mock_router_instance, vision={"enabled": True, "max_images": max_images})
images = [dict(IMG_PART, image_url={"url": f"data:image/png;base64,{n}"}) for n in ("a", "b", "c")]
await router.async_pre_routing_hook(
model="m", request_kwargs={}, messages=self._turn({"type": "text", "text": "look"}, *images)
)
content = self._classifier_user_content(mock_router_instance)
forwarded = [block for block in content if block["type"] == "image_url"]
assert len(forwarded) == expected
assert [block["image_url"]["url"] for block in forwarded] == [
f"data:image/png;base64,{n}" for n in ("a", "b", "c")[:expected]
]
@pytest.mark.asyncio
async def test_earlier_turn_images_are_not_forwarded(self, mock_router_instance):
"""Only the newest user turn's images ride along, so history cannot inflate every call.
The two turns carry different images on purpose: identical ones would pass this assertion
whichever turn the helper read.
"""
older = dict(IMG_PART, image_url={"url": "data:image/png;base64,OLDER"})
newer = dict(IMG_PART, image_url={"url": "data:image/png;base64,NEWER"})
router = self._router(mock_router_instance, vision={"enabled": True, "max_images": 5})
await router.async_pre_routing_hook(
model="m",
request_kwargs={},
messages=[
{"role": "user", "content": [{"type": "text", "text": "first"}, older]},
{"role": "assistant", "content": "ok"},
{"role": "user", "content": [{"type": "text", "text": "second"}, newer]},
],
)
content = self._classifier_user_content(mock_router_instance)
forwarded = [block for block in content if block["type"] == "image_url"]
assert [block["image_url"]["url"] for block in forwarded] == ["data:image/png;base64,NEWER"]
@pytest.mark.asyncio
async def test_logged_request_body_matches_what_was_sent(self, mock_router_instance):
"""proxy_server_request is the logged copy of the classifier call and must not drift."""
router = self._router(mock_router_instance, vision={"enabled": True})
await router.async_pre_routing_hook(
model="m", request_kwargs={}, messages=self._turn({"type": "text", "text": "what is this"}, IMG_PART)
)
call_kwargs = mock_router_instance.acompletion.call_args.kwargs
assert call_kwargs["proxy_server_request"]["body"]["messages"] == call_kwargs["messages"]
SHORT_CIRCUIT_ARMS = [
("heuristic_first", {"heuristic_first_max_tier": "SIMPLE"}, "heuristic_first_short_circuit"),
("hybrid", {"hybrid_boundary_margin": 0.05}, "hybrid_short_circuit"),
]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"classifier_type, extra, short_circuit_cause", SHORT_CIRCUIT_ARMS, ids=["heuristic_first", "hybrid"]
)
async def test_local_scorer_cannot_short_circuit_a_turn_it_cannot_see(
self, mock_router_instance, classifier_type, extra, short_circuit_cause
):
"""The scorer reads text alone, so its confidence is not a verdict on an image turn.
Both arms are tuned so the scorer WOULD short-circuit on this exact text, which is what
makes the image the only variable; a margin loose enough to leave the score undecided
would pass whether or not the guard exists.
"""
router = self._router(
mock_router_instance, vision={"enabled": True}, classifier_type=classifier_type, **extra
)
response = await router.async_pre_routing_hook(
model="m", request_kwargs={}, messages=self._turn({"type": "text", "text": "what is this"}, IMG_PART)
)
assert response.routing_decision["cause"] == "llm_classifier"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"classifier_type, extra, short_circuit_cause", SHORT_CIRCUIT_ARMS, ids=["heuristic_first", "hybrid"]
)
async def test_local_scorer_still_short_circuits_without_images(
self, mock_router_instance, classifier_type, extra, short_circuit_cause
):
"""The negative class: same router, same text, no image, and the scorer still decides."""
router = self._router(
mock_router_instance, vision={"enabled": True}, classifier_type=classifier_type, **extra
)
response = await router.async_pre_routing_hook(
model="m", request_kwargs={}, messages=[{"role": "user", "content": "what is this"}]
)
assert response.routing_decision["cause"] == short_circuit_cause
mock_router_instance.acompletion.assert_not_awaited()
def test_max_images_must_be_positive(self):
with pytest.raises(ValidationError):
ClassifierLLMConfig(model="clf", vision={"enabled": True, "max_images": 0})

View file

@ -0,0 +1,154 @@
"""
Tests for mid-task stall detection: repeated identical tool calls or repeated tool
errors, read from both Anthropic Messages and chat-completions tool-call shapes.
"""
from litellm.router_strategy.complexity_router.stall_detector import detect_stalled_task
def _anthropic_call(call_id: str, name: str, arguments: dict, *, is_error: bool) -> list[dict]:
return [
{"role": "assistant", "content": [{"type": "tool_use", "id": call_id, "name": name, "input": arguments}]},
{
"role": "user",
"content": [{"type": "tool_result", "tool_use_id": call_id, "is_error": is_error, "content": "result"}],
},
]
def _chat_completions_call(call_id: str, name: str, arguments_json: str) -> list[dict]:
return [
{
"role": "assistant",
"tool_calls": [
{"id": call_id, "type": "function", "function": {"name": name, "arguments": arguments_json}}
],
},
{"role": "tool", "tool_call_id": call_id, "content": "result"},
]
class TestDetectStalledTask:
def test_repeated_identical_anthropic_calls_are_stalled(self):
messages = [
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=False),
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=False),
*_anthropic_call("t3", "bash", {"cmd": "pytest"}, is_error=False),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is True
def test_repeated_errors_are_stalled_even_with_varied_arguments(self):
messages = [
*_anthropic_call("t1", "bash", {"cmd": "pytest tests/a.py"}, is_error=True),
*_anthropic_call("t2", "bash", {"cmd": "pytest tests/b.py"}, is_error=True),
*_anthropic_call("t3", "bash", {"cmd": "pytest tests/c.py"}, is_error=True),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is True
def test_varied_successful_calls_are_not_stalled(self):
messages = [
*_anthropic_call("t1", "bash", {"cmd": "ls"}, is_error=False),
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=False),
*_anthropic_call("t3", "grep", {"pattern": "x"}, is_error=False),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is False
def test_chat_completions_repeats_are_stalled(self):
messages = [
*_chat_completions_call("c1", "bash", '{"cmd": "pytest"}'),
*_chat_completions_call("c2", "bash", '{"cmd": "pytest"}'),
*_chat_completions_call("c3", "bash", '{"cmd": "pytest"}'),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is True
def test_chat_completions_has_no_structured_error_signal(self):
"""A chat-completions tool message carries no standard error flag, so varied calls
whose content happens to read like failures still aren't flagged on error alone."""
messages = [
*_chat_completions_call("c1", "bash", '{"cmd": "a"}'),
*_chat_completions_call("c2", "bash", '{"cmd": "b"}'),
*_chat_completions_call("c3", "bash", '{"cmd": "c"}'),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is False
def test_dict_and_json_string_arguments_compare_equal_across_surfaces(self):
messages = [
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=False),
*_chat_completions_call("c2", "bash", '{"cmd": "pytest"}'),
*_anthropic_call("t3", "bash", {"cmd": "pytest"}, is_error=False),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is True
def test_below_repeat_threshold_is_not_stalled(self):
messages = [
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=False),
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=False),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is False
def test_evidence_older_than_the_window_does_not_count(self):
"""Only the most recent `window` tool calls are considered, so a stall the model
already recovered from does not keep re-triggering forever."""
messages = [
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=False),
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=False),
*_anthropic_call("t3", "bash", {"cmd": "pytest"}, is_error=False),
*_anthropic_call("t4", "grep", {"pattern": "a"}, is_error=False),
*_anthropic_call("t5", "grep", {"pattern": "b"}, is_error=False),
]
assert detect_stalled_task(messages, window=2, repeat_threshold=2) is False
def test_evidence_survives_a_new_human_ask(self):
"""A follow-up like 'try again' must not erase evidence from before it: detection
reads the whole message list, not just the turns since the newest human ask."""
messages = [
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=False),
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=False),
*_anthropic_call("t3", "bash", {"cmd": "pytest"}, is_error=False),
{"role": "user", "content": [{"type": "text", "text": "try again"}]},
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is True
def test_a_recovered_task_is_not_stalled_while_its_old_failures_sit_in_the_window(self):
"""The three identical failures stay in the window for a few turns after the model
breaks out of them, and counting them on their own would escalate a request that is
already making progress again."""
messages = [
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=True),
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=True),
*_anthropic_call("t3", "bash", {"cmd": "pytest"}, is_error=True),
*_anthropic_call("t4", "read_file", {"path": "conftest.py"}, is_error=False),
*_anthropic_call("t5", "edit_file", {"path": "conftest.py"}, is_error=False),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is False
def test_a_retry_loop_broken_up_by_an_unrelated_call_still_counts(self):
"""Anchoring on the newest call must not require the repeats to be adjacent: a model
re-running the same failing command around a lookup in between is still stuck."""
messages = [
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=True),
*_anthropic_call("t2", "read_file", {"path": "conftest.py"}, is_error=False),
*_anthropic_call("t3", "bash", {"cmd": "pytest"}, is_error=True),
*_anthropic_call("t4", "bash", {"cmd": "pytest"}, is_error=True),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is True
def test_errors_only_count_while_the_newest_call_is_still_failing(self):
messages = [
*_anthropic_call("t1", "bash", {"cmd": "pytest a"}, is_error=True),
*_anthropic_call("t2", "bash", {"cmd": "pytest b"}, is_error=True),
*_anthropic_call("t3", "bash", {"cmd": "pytest c"}, is_error=True),
*_anthropic_call("t4", "bash", {"cmd": "pytest d"}, is_error=False),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is False
def test_no_messages_is_not_stalled(self):
assert detect_stalled_task(None, window=6, repeat_threshold=3) is False
assert detect_stalled_task([], window=6, repeat_threshold=3) is False
def test_zero_threshold_never_flags_stalled(self):
messages = [
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=True),
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=True),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=0) is False

View file

@ -388,3 +388,21 @@ class TestGpt6AstraAdvertisesItsDocumentedLevels:
"xhigh",
"max",
)
@pytest.mark.parametrize("model", ["azure/gpt-6-astra", "azure/us/gpt-6-astra"])
def test_a_foundry_deployment_also_advertises_none(self, local_model_cost_map, model):
"""Microsoft Foundry serves the same model but its API accepts reasoning_effort none
(verified live: 200 with zero reasoning tokens, and it unlocks temperature), which
OpenAI's rejects, so an Azure deployment offers none on top of low through max."""
from litellm.utils import _get_model_info_helper
model_info = dict(_get_model_info_helper(model=model, custom_llm_provider="azure"))
assert resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=True) == (
"none",
"low",
"medium",
"high",
"xhigh",
"max",
)

View file

@ -7,6 +7,7 @@ import { renderWithProviders } from "../../../../../tests/test-utils";
import { AccessGroupDetail } from "./AccessGroupsDetailsPage";
vi.mock("@/app/(dashboard)/hooks/accessGroups/useAccessGroupDetails");
vi.mock("next/navigation", () => ({ useRouter: () => ({ push: vi.fn() }) }));
vi.mock("./AccessGroupsModal/AccessGroupEditModal", () => ({
AccessGroupEditModal: ({ visible, onCancel }: { visible: boolean; onCancel: () => void }) =>
visible ? (
@ -44,6 +45,8 @@ const baseMockReturnValue = {
refetch: vi.fn(),
} as unknown as ReturnType<typeof useAccessGroupDetails>;
const unnamed = (ids: readonly string[]) => ids.map((id) => ({ id, name: null }));
const createMockAccessGroup = (overrides: Partial<AccessGroupResponse> = {}): AccessGroupResponse => ({
access_group_id: "ag-1",
access_group_name: "Test Group",
@ -53,6 +56,13 @@ const createMockAccessGroup = (overrides: Partial<AccessGroupResponse> = {}): Ac
access_agent_ids: ["agent-1"],
assigned_team_ids: ["team-1"],
assigned_key_ids: ["key-1", "key-2"],
access_mcp_servers: [{ id: "mcp-1", name: "GitHub MCP" }],
access_agents: [{ id: "agent-1", name: "Support Agent" }],
assigned_teams: [{ id: "team-1", name: "Platform Team" }],
assigned_keys: [
{ id: "key-1", name: "ci-key" },
{ id: "key-2", name: null },
],
created_at: "2025-01-01T00:00:00Z",
created_by: null,
updated_at: "2025-01-02T00:00:00Z",
@ -60,6 +70,14 @@ const createMockAccessGroup = (overrides: Partial<AccessGroupResponse> = {}): Ac
...overrides,
});
const renderWith = (overrides: Partial<AccessGroupResponse> = {}) => {
mockUseAccessGroupDetails.mockReturnValue({
...baseMockReturnValue,
data: createMockAccessGroup(overrides),
} as ReturnType<typeof useAccessGroupDetails>);
return renderWithProviders(<AccessGroupDetail accessGroupId="ag-1" onBack={vi.fn()} />);
};
describe("AccessGroupDetail", () => {
const mockOnBack = vi.fn();
const accessGroupId = "ag-1";
@ -106,9 +124,7 @@ describe("AccessGroupDetail", () => {
const user = userEvent.setup();
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
const buttons = screen.getAllByRole("button");
const backButton = buttons.find((btn) => !btn.textContent?.includes("Edit"));
await user.click(backButton!);
await user.click(screen.getByRole("button", { name: "Back" }));
expect(mockOnBack).toHaveBeenCalledTimes(1);
});
@ -128,12 +144,7 @@ describe("AccessGroupDetail", () => {
});
it("should display em dash when description is empty", () => {
mockUseAccessGroupDetails.mockReturnValue({
...baseMockReturnValue,
data: createMockAccessGroup({ description: null }),
} as ReturnType<typeof useAccessGroupDetails>);
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
renderWith({ description: null });
expect(screen.getByText("—")).toBeInTheDocument();
});
@ -144,8 +155,7 @@ describe("AccessGroupDetail", () => {
expect(screen.queryByRole("dialog", { name: "Edit Access Group" })).not.toBeInTheDocument();
const editButton = screen.getByRole("button", { name: /Edit Access Group/i });
await user.click(editButton);
await user.click(screen.getByRole("button", { name: /Edit Access Group/i }));
expect(screen.getByRole("dialog", { name: "Edit Access Group" })).toBeInTheDocument();
});
@ -161,88 +171,126 @@ describe("AccessGroupDetail", () => {
expect(screen.queryByRole("dialog", { name: "Edit Access Group" })).not.toBeInTheDocument();
});
it("should display attached keys", () => {
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
describe("attached keys", () => {
it("should show the key alias and hide the token when the key has an alias", () => {
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
expect(screen.getByText("Attached Keys")).toBeInTheDocument();
expect(screen.getByText("key-1")).toBeInTheDocument();
expect(screen.getByText("key-2")).toBeInTheDocument();
expect(screen.getByText("Attached Keys")).toBeInTheDocument();
expect(screen.getByText("ci-key")).toBeInTheDocument();
expect(screen.queryByText("key-1")).not.toBeInTheDocument();
});
it("should fall back to the token when the key has no alias", () => {
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
expect(screen.getByText("key-2")).toBeInTheDocument();
});
it("should link each key to its detail page", () => {
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
expect(screen.getByRole("link", { name: "ci-key" })).toHaveAttribute(
"href",
expect.stringContaining("key=key-1"),
);
expect(screen.getByRole("link", { name: "key-2" })).toHaveAttribute("href", expect.stringContaining("key=key-2"));
});
it("should reveal the token in a tooltip when hovering an aliased key", async () => {
const user = userEvent.setup();
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
await user.hover(screen.getByText("ci-key"));
expect(await screen.findByText("key-1")).toBeInTheDocument();
});
it("should show View All button for keys when more than 5", () => {
renderWith({ assigned_keys: unnamed(["k1", "k2", "k3", "k4", "k5", "k6"]) });
expect(screen.getByRole("button", { name: "View All (6)" })).toBeInTheDocument();
expect(screen.queryByText("k6")).not.toBeInTheDocument();
});
it("should toggle between View All and Show Less for keys", async () => {
const user = userEvent.setup();
renderWith({ assigned_keys: unnamed(["k1", "k2", "k3", "k4", "k5", "k6"]) });
await user.click(screen.getByRole("button", { name: "View All (6)" }));
expect(screen.getByRole("button", { name: "Show Less" })).toBeInTheDocument();
expect(screen.getByText("k6")).toBeInTheDocument();
await user.click(screen.getByRole("button", { name: "Show Less" }));
expect(screen.getByRole("button", { name: "View All (6)" })).toBeInTheDocument();
});
it("should show empty state when no keys attached", () => {
renderWith({ assigned_keys: [] });
expect(screen.getByText("No keys attached")).toBeInTheDocument();
});
it("should truncate long unaliased tokens with ellipsis", () => {
renderWith({ assigned_keys: unnamed(["a".repeat(25)]) });
expect(screen.getByText(/^a{10}\.\.\.a{6}$/)).toBeInTheDocument();
});
it("should not truncate a long alias", () => {
const alias = "b".repeat(25);
renderWith({ assigned_keys: [{ id: "a".repeat(25), name: alias }] });
expect(screen.getByText(alias)).toBeInTheDocument();
});
});
it("should display attached teams", () => {
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
describe("attached teams", () => {
it("should show the team alias and hide the id when the team has an alias", () => {
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
expect(screen.getByText("Attached Teams")).toBeInTheDocument();
expect(screen.getByText("team-1")).toBeInTheDocument();
expect(screen.getByText("Attached Teams")).toBeInTheDocument();
expect(screen.getByText("Platform Team")).toBeInTheDocument();
expect(screen.queryByText("team-1")).not.toBeInTheDocument();
});
it("should link each team to its detail page", () => {
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
expect(screen.getByRole("link", { name: "Platform Team" })).toHaveAttribute(
"href",
expect.stringContaining("team=team-1"),
);
});
it("should reveal the team id in a tooltip when hovering an aliased team", async () => {
const user = userEvent.setup();
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
await user.hover(screen.getByText("Platform Team"));
expect(await screen.findByText("team-1")).toBeInTheDocument();
});
it("should fall back to the team id when the team has no alias", () => {
renderWith({ assigned_teams: unnamed(["team-ghost"]) });
expect(screen.getByText("team-ghost")).toBeInTheDocument();
});
it("should show View All button for teams when more than 5", () => {
renderWith({ assigned_teams: unnamed(["t1", "t2", "t3", "t4", "t5", "t6"]) });
expect(screen.getByRole("button", { name: "View All (6)" })).toBeInTheDocument();
});
it("should show empty state when no teams attached", () => {
renderWith({ assigned_teams: [] });
expect(screen.getByText("No teams attached")).toBeInTheDocument();
});
});
it("should show View All button for keys when more than 5", () => {
mockUseAccessGroupDetails.mockReturnValue({
...baseMockReturnValue,
data: createMockAccessGroup({
assigned_key_ids: ["k1", "k2", "k3", "k4", "k5", "k6"],
}),
} as ReturnType<typeof useAccessGroupDetails>);
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
expect(screen.getByRole("button", { name: "View All (6)" })).toBeInTheDocument();
});
it("should toggle between View All and Show Less for keys", async () => {
const user = userEvent.setup();
mockUseAccessGroupDetails.mockReturnValue({
...baseMockReturnValue,
data: createMockAccessGroup({
assigned_key_ids: ["k1", "k2", "k3", "k4", "k5", "k6"],
}),
} as ReturnType<typeof useAccessGroupDetails>);
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
await user.click(screen.getByRole("button", { name: "View All (6)" }));
expect(screen.getByRole("button", { name: "Show Less" })).toBeInTheDocument();
await user.click(screen.getByRole("button", { name: "Show Less" }));
expect(screen.getByRole("button", { name: "View All (6)" })).toBeInTheDocument();
});
it("should show View All button for teams when more than 5", () => {
mockUseAccessGroupDetails.mockReturnValue({
...baseMockReturnValue,
data: createMockAccessGroup({
assigned_team_ids: ["t1", "t2", "t3", "t4", "t5", "t6"],
}),
} as ReturnType<typeof useAccessGroupDetails>);
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
expect(screen.getByRole("button", { name: "View All (6)" })).toBeInTheDocument();
});
it("should show empty state when no keys attached", () => {
mockUseAccessGroupDetails.mockReturnValue({
...baseMockReturnValue,
data: createMockAccessGroup({ assigned_key_ids: [] }),
} as ReturnType<typeof useAccessGroupDetails>);
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
expect(screen.getByText("No keys attached")).toBeInTheDocument();
});
it("should show empty state when no teams attached", () => {
mockUseAccessGroupDetails.mockReturnValue({
...baseMockReturnValue,
data: createMockAccessGroup({ assigned_team_ids: [] }),
} as ReturnType<typeof useAccessGroupDetails>);
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
expect(screen.getByText("No teams attached")).toBeInTheDocument();
});
it("should display Models tab with model IDs", () => {
it("should display Models tab with model names", () => {
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
expect(screen.getByRole("tab", { name: /Models/i })).toBeInTheDocument();
@ -250,73 +298,90 @@ describe("AccessGroupDetail", () => {
expect(screen.getByText("model-2")).toBeInTheDocument();
});
it("should display MCP Servers tab with server IDs", async () => {
const user = userEvent.setup();
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
describe("MCP Servers tab", () => {
it("should show server names instead of ids", async () => {
const user = userEvent.setup();
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
const mcpTab = screen.getByRole("tab", { name: /MCP Servers/i });
expect(mcpTab).toBeInTheDocument();
await user.click(mcpTab);
expect(screen.getByText("mcp-1")).toBeInTheDocument();
await user.click(screen.getByRole("tab", { name: /MCP Servers/i }));
expect(screen.getByText("GitHub MCP")).toBeInTheDocument();
expect(screen.queryByText("mcp-1")).not.toBeInTheDocument();
});
it("should reveal the server id in a tooltip when hovering the name", async () => {
const user = userEvent.setup();
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
await user.click(screen.getByRole("tab", { name: /MCP Servers/i }));
await user.hover(screen.getByText("GitHub MCP"));
expect(await screen.findByText("mcp-1")).toBeInTheDocument();
});
it("should fall back to the id when the server has no name", async () => {
const user = userEvent.setup();
renderWith({ access_mcp_servers: unnamed(["mcp-deleted"]) });
await user.click(screen.getByRole("tab", { name: /MCP Servers/i }));
expect(screen.getByText("mcp-deleted")).toBeInTheDocument();
});
it("should show empty state when none assigned", async () => {
const user = userEvent.setup();
renderWith({ access_mcp_servers: [] });
await user.click(screen.getByRole("tab", { name: /MCP Servers/i }));
expect(screen.getByText("No MCP servers assigned to this group")).toBeInTheDocument();
});
});
it("should display Agents tab with agent IDs", async () => {
const user = userEvent.setup();
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
describe("Agents tab", () => {
it("should show agent names instead of ids", async () => {
const user = userEvent.setup();
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
const agentsTab = screen.getByRole("tab", { name: /Agents/i });
expect(agentsTab).toBeInTheDocument();
await user.click(agentsTab);
expect(screen.getByText("agent-1")).toBeInTheDocument();
await user.click(screen.getByRole("tab", { name: /Agents/i }));
expect(screen.getByText("Support Agent")).toBeInTheDocument();
expect(screen.queryByText("agent-1")).not.toBeInTheDocument();
});
it("should fall back to the id when the agent has no name", async () => {
const user = userEvent.setup();
renderWith({ access_agents: unnamed(["agent-deleted"]) });
await user.click(screen.getByRole("tab", { name: /Agents/i }));
expect(screen.getByText("agent-deleted")).toBeInTheDocument();
});
it("should show empty state when none assigned", async () => {
const user = userEvent.setup();
renderWith({ access_agents: [] });
await user.click(screen.getByRole("tab", { name: /Agents/i }));
expect(screen.getByText("No agents assigned to this group")).toBeInTheDocument();
});
});
it("should show empty state in Models tab when no models assigned", () => {
mockUseAccessGroupDetails.mockReturnValue({
...baseMockReturnValue,
data: createMockAccessGroup({ access_model_names: [] }),
} as ReturnType<typeof useAccessGroupDetails>);
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
renderWith({ access_model_names: [] });
expect(screen.getByText("No models assigned to this group")).toBeInTheDocument();
});
it("should show empty state in MCP Servers tab when none assigned", async () => {
const user = userEvent.setup();
mockUseAccessGroupDetails.mockReturnValue({
...baseMockReturnValue,
data: createMockAccessGroup({ access_mcp_server_ids: [] }),
} as ReturnType<typeof useAccessGroupDetails>);
it("should count resources from the resolved lists in the tab badges", () => {
renderWith({
access_mcp_servers: unnamed(["m1", "m2", "m3"]),
access_agents: unnamed(["a1", "a2"]),
});
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
await user.click(screen.getByRole("tab", { name: /MCP Servers/i }));
expect(screen.getByText("No MCP servers assigned to this group")).toBeInTheDocument();
});
it("should show empty state in Agents tab when none assigned", async () => {
const user = userEvent.setup();
mockUseAccessGroupDetails.mockReturnValue({
...baseMockReturnValue,
data: createMockAccessGroup({ access_agent_ids: [] }),
} as ReturnType<typeof useAccessGroupDetails>);
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
await user.click(screen.getByRole("tab", { name: /Agents/i }));
expect(screen.getByText("No agents assigned to this group")).toBeInTheDocument();
});
it("should truncate long key IDs with ellipsis", () => {
const longKeyId = "a".repeat(25);
mockUseAccessGroupDetails.mockReturnValue({
...baseMockReturnValue,
data: createMockAccessGroup({ assigned_key_ids: [longKeyId] }),
} as ReturnType<typeof useAccessGroupDetails>);
renderWithProviders(<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />);
expect(screen.getByText(/a{10}\.\.\.a{6}/)).toBeInTheDocument();
expect(screen.getByRole("tab", { name: /MCP Servers/i })).toHaveTextContent("3");
expect(screen.getByRole("tab", { name: /Agents/i })).toHaveTextContent("2");
});
it("should display created and last updated timestamps", () => {

View file

@ -2,14 +2,20 @@ import { useAccessGroupDetails } from "@/app/(dashboard)/hooks/accessGroups/useA
import { ArrowLeftIcon, BotIcon, EditIcon, KeyIcon, LayersIcon, ServerIcon, UsersIcon } from "lucide-react";
import { useState } from "react";
import DefaultProxyAdminTag from "@/components/common_components/DefaultProxyAdminTag";
import { BadgeLink } from "@/components/shared/BadgeLink";
import CopyButton from "@/components/shared/CopyButton";
import { Badge } from "@/components/ui/badge";
import { Button } from "@/components/ui/button";
import { Card, CardAction, CardContent, CardHeader, CardTitle } from "@/components/ui/card";
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
import { SimpleTooltip } from "@/components/ui/tooltip";
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
import type { components } from "@/lib/http/schema";
import { keyDetailHref, teamDetailHref } from "@/utils/entityLinks";
import { AccessGroupEditModal } from "./AccessGroupsModal/AccessGroupEditModal";
type AccessGroupResource = components["schemas"]["AccessGroupResource"];
interface AccessGroupDetailProps {
accessGroupId: string;
onBack: () => void;
@ -17,16 +23,24 @@ interface AccessGroupDetailProps {
const MAX_PREVIEW = 5;
function ResourceList({ ids, emptyMessage }: { ids: string[]; emptyMessage: string }) {
if (ids.length === 0) {
const shortId = (id: string) => (id.length > 20 ? `${id.slice(0, 10)}...${id.slice(-6)}` : id);
function ResourceList({ items, emptyMessage }: { items: readonly AccessGroupResource[]; emptyMessage: string }) {
if (items.length === 0) {
return <p className="py-8 text-center text-sm text-muted-foreground">{emptyMessage}</p>;
}
return (
<div className="grid grid-cols-1 gap-4 sm:grid-cols-2 md:grid-cols-3 lg:grid-cols-4">
{ids.map((id) => (
{items.map(({ id, name }) => (
<Card key={id} size="sm">
<CardContent>
<code className="font-mono text-xs break-all text-foreground">{id}</code>
{name ? (
<SimpleTooltip content={id}>
<span className="text-sm font-medium break-all text-foreground">{name}</span>
</SimpleTooltip>
) : (
<code className="font-mono text-xs break-all text-foreground">{id}</code>
)}
</CardContent>
</Card>
))}
@ -34,6 +48,23 @@ function ResourceList({ ids, emptyMessage }: { ids: string[]; emptyMessage: stri
);
}
function ResourceBadge({
resource: { id, name },
href,
fallback,
}: {
resource: AccessGroupResource;
href: string;
fallback: (id: string) => string;
}) {
const badge = (
<BadgeLink href={href} className={name ? undefined : "font-mono"}>
{name ?? fallback(id)}
</BadgeLink>
);
return name ? <SimpleTooltip content={id}>{badge}</SimpleTooltip> : badge;
}
export function AccessGroupDetail({ accessGroupId, onBack }: AccessGroupDetailProps) {
const { data: accessGroup, isLoading } = useAccessGroupDetails(accessGroupId);
const [isEditModalVisible, setIsEditModalVisible] = useState(false);
@ -61,14 +92,14 @@ export function AccessGroupDetail({ accessGroupId, onBack }: AccessGroupDetailPr
);
}
const modelIds = accessGroup.access_model_names ?? [];
const mcpServerIds = accessGroup.access_mcp_server_ids ?? [];
const agentIds = accessGroup.access_agent_ids ?? [];
const keyIds = accessGroup.assigned_key_ids ?? [];
const teamIds = accessGroup.assigned_team_ids ?? [];
const models = accessGroup.access_model_names.map((id) => ({ id, name: null }));
const mcpServers = accessGroup.access_mcp_servers;
const agents = accessGroup.access_agents;
const keys = accessGroup.assigned_keys;
const teams = accessGroup.assigned_teams;
const displayedKeys = showAllKeys ? keyIds : keyIds.slice(0, MAX_PREVIEW);
const displayedTeams = showAllTeams ? teamIds : teamIds.slice(0, MAX_PREVIEW);
const displayedKeys = showAllKeys ? keys : keys.slice(0, MAX_PREVIEW);
const displayedTeams = showAllTeams ? teams : teams.slice(0, MAX_PREVIEW);
return (
<div className="p-6 px-12">
@ -129,23 +160,21 @@ export function AccessGroupDetail({ accessGroupId, onBack }: AccessGroupDetailPr
<CardTitle className="flex items-center gap-2">
<KeyIcon className="size-4" />
Attached Keys
<Badge variant="secondary">{keyIds.length}</Badge>
<Badge variant="secondary">{keys.length}</Badge>
</CardTitle>
{keyIds.length > MAX_PREVIEW && (
{keys.length > MAX_PREVIEW && (
<CardAction>
<Button variant="link" size="sm" onClick={() => setShowAllKeys(!showAllKeys)}>
{showAllKeys ? "Show Less" : `View All (${keyIds.length})`}
{showAllKeys ? "Show Less" : `View All (${keys.length})`}
</Button>
</CardAction>
)}
</CardHeader>
<CardContent>
{keyIds.length > 0 ? (
{keys.length > 0 ? (
<div className="flex flex-wrap gap-2">
{displayedKeys.map((id) => (
<Badge key={id} variant="secondary" className="font-mono">
{id.length > 20 ? `${id.slice(0, 10)}...${id.slice(-6)}` : id}
</Badge>
{displayedKeys.map((key) => (
<ResourceBadge key={key.id} resource={key} href={keyDetailHref(key.id)} fallback={shortId} />
))}
</div>
) : (
@ -159,23 +188,21 @@ export function AccessGroupDetail({ accessGroupId, onBack }: AccessGroupDetailPr
<CardTitle className="flex items-center gap-2">
<UsersIcon className="size-4" />
Attached Teams
<Badge variant="secondary">{teamIds.length}</Badge>
<Badge variant="secondary">{teams.length}</Badge>
</CardTitle>
{teamIds.length > MAX_PREVIEW && (
{teams.length > MAX_PREVIEW && (
<CardAction>
<Button variant="link" size="sm" onClick={() => setShowAllTeams(!showAllTeams)}>
{showAllTeams ? "Show Less" : `View All (${teamIds.length})`}
{showAllTeams ? "Show Less" : `View All (${teams.length})`}
</Button>
</CardAction>
)}
</CardHeader>
<CardContent>
{teamIds.length > 0 ? (
{teams.length > 0 ? (
<div className="flex flex-wrap gap-2">
{displayedTeams.map((id) => (
<Badge key={id} variant="secondary" className="font-mono">
{id}
</Badge>
{displayedTeams.map((team) => (
<ResourceBadge key={team.id} resource={team} href={teamDetailHref(team.id)} fallback={(id) => id} />
))}
</div>
) : (
@ -192,27 +219,27 @@ export function AccessGroupDetail({ accessGroupId, onBack }: AccessGroupDetailPr
<TabsTrigger value="models" className="flex-none gap-2 rounded-none px-4 py-2">
<LayersIcon className="size-4" />
Models
<Badge variant="secondary">{modelIds.length}</Badge>
<Badge variant="secondary">{models.length}</Badge>
</TabsTrigger>
<TabsTrigger value="mcp" className="flex-none gap-2 rounded-none px-4 py-2">
<ServerIcon className="size-4" />
MCP Servers
<Badge variant="secondary">{mcpServerIds.length}</Badge>
<Badge variant="secondary">{mcpServers.length}</Badge>
</TabsTrigger>
<TabsTrigger value="agents" className="flex-none gap-2 rounded-none px-4 py-2">
<BotIcon className="size-4" />
Agents
<Badge variant="secondary">{agentIds.length}</Badge>
<Badge variant="secondary">{agents.length}</Badge>
</TabsTrigger>
</TabsList>
<TabsContent value="models" className="pt-4">
<ResourceList ids={modelIds} emptyMessage="No models assigned to this group" />
<ResourceList items={models} emptyMessage="No models assigned to this group" />
</TabsContent>
<TabsContent value="mcp" className="pt-4">
<ResourceList ids={mcpServerIds} emptyMessage="No MCP servers assigned to this group" />
<ResourceList items={mcpServers} emptyMessage="No MCP servers assigned to this group" />
</TabsContent>
<TabsContent value="agents" className="pt-4">
<ResourceList ids={agentIds} emptyMessage="No agents assigned to this group" />
<ResourceList items={agents} emptyMessage="No agents assigned to this group" />
</TabsContent>
</Tabs>
</CardContent>

View file

@ -42,6 +42,10 @@ const accessGroup: AccessGroupResponse = {
access_agent_ids: ["agent-1"],
assigned_team_ids: [],
assigned_key_ids: [],
access_mcp_servers: [{ id: "srv-1", name: "Server One" }],
access_agents: [{ id: "agent-1", name: "Agent One" }],
assigned_teams: [],
assigned_keys: [],
created_at: "2024-01-01T00:00:00Z",
created_by: "user-1",
updated_at: "2024-01-02T00:00:00Z",

View file

@ -15,6 +15,10 @@ const mockAccessGroups: AccessGroupResponse[] = [
access_agent_ids: ["a1"],
assigned_team_ids: [],
assigned_key_ids: [],
access_mcp_servers: [{ id: "s1", name: "Server One" }],
access_agents: [{ id: "a1", name: "Agent One" }],
assigned_teams: [],
assigned_keys: [],
created_at: "2024-01-15T10:00:00Z",
created_by: "user-1",
updated_at: "2024-01-20T12:00:00Z",
@ -29,6 +33,10 @@ const mockAccessGroups: AccessGroupResponse[] = [
access_agent_ids: [],
assigned_team_ids: [],
assigned_key_ids: [],
access_mcp_servers: [],
access_agents: [],
assigned_teams: [],
assigned_keys: [],
created_at: "2024-01-10T09:00:00Z",
created_by: null,
updated_at: "2024-01-12T11:00:00Z",

View file

@ -46,6 +46,10 @@ const mockAccessGroups: AccessGroupResponse[] = [
access_agent_ids: [],
assigned_team_ids: [],
assigned_key_ids: [],
access_mcp_servers: [],
access_agents: [],
assigned_teams: [],
assigned_keys: [],
created_at: "2025-01-01T00:00:00Z",
created_by: "user-1",
updated_at: "2025-01-01T00:00:00Z",

View file

@ -3,23 +3,11 @@ import { createQueryKeys } from "../common/queryKeysFactory";
import { getProxyBaseUrl, getGlobalLitellmHeaderName, deriveErrorMessage, handleError } from "@/components/networking";
import { all_admin_roles } from "@/utils/roles";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import type { components } from "@/lib/http/schema";
// ── Types ────────────────────────────────────────────────────────────────────
export interface AccessGroupResponse {
access_group_id: string;
access_group_name: string;
description: string | null;
access_model_names: string[];
access_mcp_server_ids: string[];
access_agent_ids: string[];
assigned_team_ids: string[];
assigned_key_ids: string[];
created_at: string;
created_by: string | null;
updated_at: string;
updated_by: string | null;
}
export type AccessGroupResponse = components["schemas"]["AccessGroupResponse"];
// ── Query keys (shared across access-group hooks) ────────────────────────────

View file

@ -33,6 +33,8 @@ import { ModelGroup } from "@/components/llm_calls/fetch_models";
import AdaptiveRoutingConfig from "./AdaptiveRoutingConfig";
import ClassificationMethodConfig from "./ClassificationMethodConfig";
import ContextWindowEscalationConfig from "./ContextWindowEscalationConfig";
import ResponseFormatControls from "./ResponseFormatControls";
import StallEscalationConfig from "./StallEscalationConfig";
import { Restricted, restrictedBy } from "./TierRestrictions";
import { type TierSetAction, applyTierSetAction, setFallbackTier } from "./tier_set_actions";
import {
@ -422,6 +424,14 @@ export interface ComplexityRouterConfigValue {
deployment_affinity?: boolean;
/** Plan-mode floor as a tier ROW ID, unset meaning off. The wire carries the row's name. */
plan_mode_min_tier?: string;
/**
* Mid-task stall escalation. Undefined means off, which keeps all three keys out of the payload:
* the backend rejects them alongside session pinning, user-turn classification and a custom tier
* set, so an off router must stay silent about them rather than send an explicit false.
*/
stall_escalation_enabled?: boolean;
stall_escalation_window?: number;
stall_escalation_repeat_threshold?: number;
adaptive?: boolean;
adaptive_weights?: AdaptiveRouterWeights;
tier_distance_penalty?: number;
@ -575,25 +585,6 @@ const PlanModeOverrideControls: React.FC<{
</>
);
const ResponseFormatControls: React.FC<{
value: ComplexityRouterConfigValue;
onChange: (value: ComplexityRouterConfigValue) => void;
}> = ({ value, onChange }) => (
<>
<div className="flex items-center gap-2 mb-2">
<Switch
checked={value.return_raw_model_name ?? false}
onCheckedChange={(returnRawModelName) => onChange({ ...value, return_raw_model_name: returnRawModelName })}
aria-label="Return raw model name"
/>
<strong className="font-semibold">Return raw model name</strong>
</div>
<span className="block text-xs text-muted-foreground">
Return the resolved underlying model name in responses instead of the autorouter alias.
</span>
</>
);
const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
modelInfo,
value,
@ -859,6 +850,15 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
label: <strong className="text-foreground font-semibold">Advanced: Context Window Escalation</strong>,
children: <ContextWindowEscalationConfig value={value} onChange={onChange} />,
},
{
key: "stall-escalation",
label: <strong className="text-foreground font-semibold">Advanced: Stalled Task Escalation</strong>,
children: (
<Restricted by={restrictedBy(value, "stallEscalation")}>
<StallEscalationConfig value={value} onChange={onChange} />
</Restricted>
),
},
{
key: "response",
label: <strong className="text-foreground font-semibold">Advanced: Response Format</strong>,

View file

@ -0,0 +1,24 @@
import { Switch } from "@/components/ui/switch";
import React from "react";
import type { ComplexityRouterConfigValue } from "./ComplexityRouterConfig";
const ResponseFormatControls: React.FC<{
value: ComplexityRouterConfigValue;
onChange: (value: ComplexityRouterConfigValue) => void;
}> = ({ value, onChange }) => (
<>
<div className="flex items-center gap-2 mb-2">
<Switch
checked={value.return_raw_model_name ?? false}
onCheckedChange={(returnRawModelName) => onChange({ ...value, return_raw_model_name: returnRawModelName })}
aria-label="Return raw model name"
/>
<strong className="font-semibold">Return raw model name</strong>
</div>
<span className="block text-xs text-muted-foreground">
Return the resolved underlying model name in responses instead of the autorouter alias.
</span>
</>
);
export default ResponseFormatControls;

View file

@ -0,0 +1,115 @@
import { fireEvent, renderWithProviders, screen } from "../../../tests/test-utils";
import { vi } from "vitest";
import type { ComplexityRouterConfigValue } from "./ComplexityRouterConfig";
import StallEscalationConfig, { stallEscalationBlockedReason } from "./StallEscalationConfig";
const tiers = { SIMPLE: "gpt-4o-mini", MEDIUM: "gpt-4o", COMPLEX: "claude-sonnet-4", REASONING: "o1-preview" };
const baseValue: ComplexityRouterConfigValue = {
tiers,
classifier_type: "heuristic",
};
const renderConfig = (value: Partial<ComplexityRouterConfigValue> = {}) => {
const onChange = vi.fn();
renderWithProviders(<StallEscalationConfig value={{ ...baseValue, ...value }} onChange={onChange} />);
return onChange;
};
const toggle = () => screen.getByRole("switch", { name: "Escalate a stalled task to a stronger model" });
describe("stallEscalationBlockedReason", () => {
it("blocks on session pinning, which replays a model instead of classifying", () => {
expect(stallEscalationBlockedReason({ ...baseValue, session_affinity: true })).toContain("Classification Method");
});
it("blocks on user-turn classification, which skips the agent-loop turns a stall shows up in", () => {
expect(stallEscalationBlockedReason({ ...baseValue, classification_mode: "user_turn" })).toContain("every request");
});
it("allows the default every-request router", () => {
expect(stallEscalationBlockedReason(baseValue)).toBeNull();
});
});
describe("StallEscalationConfig", () => {
it("hides the knobs until the feature is turned on", () => {
renderConfig();
expect(toggle()).not.toBeChecked();
expect(screen.queryByLabelText("Repeats before escalating")).not.toBeInTheDocument();
});
it("turning it on seeds both knobs so the saved config is explicit rather than half-set", () => {
const onChange = renderConfig();
fireEvent.click(toggle());
expect(onChange).toHaveBeenCalledWith(
expect.objectContaining({
stall_escalation_enabled: true,
stall_escalation_window: 6,
stall_escalation_repeat_threshold: 3,
}),
);
});
it("turning it off clears all three keys, since the backend rejects them next to session pinning", () => {
const onChange = renderConfig({
stall_escalation_enabled: true,
stall_escalation_window: 6,
stall_escalation_repeat_threshold: 3,
});
fireEvent.click(toggle());
expect(onChange).toHaveBeenCalledWith(
expect.objectContaining({
stall_escalation_enabled: undefined,
stall_escalation_window: undefined,
stall_escalation_repeat_threshold: undefined,
}),
);
});
it("raises the window to match a larger threshold, which could otherwise never be reached", () => {
const onChange = renderConfig({
stall_escalation_enabled: true,
stall_escalation_window: 4,
stall_escalation_repeat_threshold: 3,
});
fireEvent.change(screen.getByLabelText("Repeats before escalating"), { target: { value: "9" } });
expect(onChange).toHaveBeenCalledWith(
expect.objectContaining({ stall_escalation_repeat_threshold: 9, stall_escalation_window: 9 }),
);
});
it("holds the window at the threshold when someone types a smaller one", () => {
const onChange = renderConfig({
stall_escalation_enabled: true,
stall_escalation_window: 6,
stall_escalation_repeat_threshold: 3,
});
fireEvent.change(screen.getByLabelText("Recent calls examined"), { target: { value: "1" } });
expect(onChange).toHaveBeenCalledWith(expect.objectContaining({ stall_escalation_window: 3 }));
});
it("floors the threshold at 2, below which a single ordinary retry would escalate", () => {
const onChange = renderConfig({ stall_escalation_enabled: true });
fireEvent.change(screen.getByLabelText("Repeats before escalating"), { target: { value: "1" } });
expect(onChange).toHaveBeenCalledWith(expect.objectContaining({ stall_escalation_repeat_threshold: 2 }));
});
it("disables the toggle and says why when session pinning is on", () => {
renderConfig({ session_affinity: true });
expect(toggle()).toHaveAttribute("aria-disabled", "true");
expect(screen.getByText(/How often to classify/)).toBeInTheDocument();
});
it("hides the knobs when a blocker is switched on under an already-enabled router", () => {
renderConfig({ stall_escalation_enabled: true, session_affinity: true });
expect(screen.queryByLabelText("Repeats before escalating")).not.toBeInTheDocument();
});
it("still lets an already-on router turn it off once a blocker appears, which the save needs", () => {
const onChange = renderConfig({ stall_escalation_enabled: true, session_affinity: true });
expect(toggle()).not.toHaveAttribute("aria-disabled", "true");
fireEvent.click(toggle());
expect(onChange).toHaveBeenCalledWith(expect.objectContaining({ stall_escalation_enabled: undefined }));
});
});

View file

@ -0,0 +1,119 @@
import { Input } from "@/components/ui/input";
import { Switch } from "@/components/ui/switch";
import React from "react";
import { type ComplexityRouterConfigValue, classificationFrequency } from "./ComplexityRouterConfig";
export const DEFAULT_STALL_ESCALATION_WINDOW = 6;
export const DEFAULT_STALL_ESCALATION_REPEAT_THRESHOLD = 3;
/**
* Why the toggle is unavailable, or null when it can be turned on. Both blockers replay a held
* routing decision instead of classifying most turns, so detection would never see the tool
* calls it reads.
*/
export const stallEscalationBlockedReason = (value: ComplexityRouterConfigValue): string | null => {
const frequency = classificationFrequency(value);
if (frequency === "session")
return 'Set "How often to classify" to every request under Advanced: Classification Method to use this. Scoring once per session replays that model instead of classifying, so a stall never reaches the classifier.';
if (frequency === "user_turn")
return 'Set "How often to classify" to every request under Advanced: Classification Method to use this. Scoring only new user messages skips the tool-call turns a stall shows up in.';
return null;
};
const clampedInt = (raw: string, min: number, fallback: number): number => {
const parsed = Number(raw);
if (!Number.isFinite(parsed)) return fallback;
return Math.max(min, Math.trunc(parsed));
};
const StallEscalationConfig: React.FC<{
value: ComplexityRouterConfigValue;
onChange: (value: ComplexityRouterConfigValue) => void;
}> = ({ value, onChange }) => {
const enabled = value.stall_escalation_enabled ?? false;
const blockedReason = stallEscalationBlockedReason(value);
const window = value.stall_escalation_window ?? DEFAULT_STALL_ESCALATION_WINDOW;
const threshold = value.stall_escalation_repeat_threshold ?? DEFAULT_STALL_ESCALATION_REPEAT_THRESHOLD;
// A threshold above the window can never be reached, and the backend rejects the pair, so the
// window rises with the threshold rather than letting the form save something inert.
const commitThreshold = (raw: string) => {
const nextThreshold = clampedInt(raw, 2, DEFAULT_STALL_ESCALATION_REPEAT_THRESHOLD);
onChange({
...value,
stall_escalation_repeat_threshold: nextThreshold,
stall_escalation_window: Math.max(window, nextThreshold),
});
};
const commitWindow = (raw: string) => {
const nextWindow = clampedInt(raw, 1, DEFAULT_STALL_ESCALATION_WINDOW);
onChange({
...value,
stall_escalation_window: Math.max(nextWindow, threshold),
});
};
const toggle = (next: boolean) => {
const enabledValue: ComplexityRouterConfigValue = {
...value,
stall_escalation_enabled: next || undefined,
stall_escalation_window: next ? window : undefined,
stall_escalation_repeat_threshold: next ? threshold : undefined,
};
onChange(enabledValue);
};
return (
<>
<div className="flex items-center gap-2 mb-2">
<Switch
checked={enabled}
// Blocked only prevents turning it on: an already-on router that just became
// blocked (e.g. session pinning turned on afterward) still needs a way to turn
// this back off, since the backend rejects saving both together.
disabled={blockedReason !== null && !enabled}
onCheckedChange={toggle}
aria-label="Escalate a stalled task to a stronger model"
/>
<strong className="font-semibold">Escalate a stalled task to a stronger model</strong>
</div>
<span className="block text-xs mb-3 text-muted-foreground">
When the model keeps repeating the same tool call, or the same call keeps erroring, bump the request one tier
higher for as long as it looks stuck. The automatic counterpart to an escalation keyword: nobody has to notice
the loop and ask. Off means a stuck task keeps the model it was classified onto.
{blockedReason !== null && ` ${blockedReason}`}
</span>
{enabled && blockedReason === null && (
<div className="flex flex-wrap gap-4">
<div style={{ maxWidth: 240 }}>
<label className="block text-sm font-medium mb-1" htmlFor="stall-escalation-repeat-threshold">
Repeats before escalating
</label>
<Input
id="stall-escalation-repeat-threshold"
inputMode="numeric"
value={threshold}
onChange={(event) => commitThreshold(event.target.value)}
/>
<span className="block text-xs mt-1 text-muted-foreground">
How many identical or failing calls count as stuck. At least 2; lower reacts sooner and misfires more.
</span>
</div>
<div style={{ maxWidth: 240 }}>
<label className="block text-sm font-medium mb-1" htmlFor="stall-escalation-window">
Recent calls examined
</label>
<Input
id="stall-escalation-window"
inputMode="numeric"
value={window}
onChange={(event) => commitWindow(event.target.value)}
/>
<span className="block text-xs mt-1 text-muted-foreground">
How far back to look, in tool calls. Never below the repeat count, since that could never be reached.
</span>
</div>
</div>
)}
</>
);
};
export default StallEscalationConfig;

View file

@ -395,6 +395,9 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({
embeddingModel,
matchThreshold,
escalationKeywords,
stallEscalationEnabled: complexityRouterConfig.stall_escalation_enabled,
stallEscalationWindow: complexityRouterConfig.stall_escalation_window,
stallEscalationRepeatThreshold: complexityRouterConfig.stall_escalation_repeat_threshold,
adaptive: complexityRouterConfig.adaptive ?? false,
adaptiveWeights: complexityRouterConfig.adaptive_weights ?? DEFAULT_ADAPTIVE_WEIGHTS,
tierDistancePenalty: complexityRouterConfig.tier_distance_penalty ?? DEFAULT_TIER_DISTANCE_PENALTY,

View file

@ -592,6 +592,12 @@ describe("classifier prompt and fallback", () => {
timeout_ms: 1,
});
});
it.each([{}, { system_prompt: "x" }])("normalizeClassifierLlmConfig carries vision through %o", (extra) => {
const base = { model: "m", timeout_ms: 1, ...extra };
const vision = { enabled: true, max_images: 2 };
expect(normalizeClassifierLlmConfig({ ...base, vision })).toEqual({ ...base, vision });
});
});
describe("tier labels", () => {
@ -1058,6 +1064,9 @@ describe("buildComplexityRouterConfig with an edited tier set", () => {
heuristicFirstMaxTier: "SIMPLE",
hybridBoundaryMargin: 0.03,
customTechnicalKeywords: ["kubernetes"],
stallEscalationEnabled: true,
stallEscalationWindow: 6,
stallEscalationRepeatThreshold: 3,
};
const emittingType = key === "heuristic_first_max_tier" ? "heuristic_first" : "llm";
const typeForKey = key === "hybrid_boundary_margin" ? "hybrid" : emittingType;
@ -1152,6 +1161,34 @@ describe("hydrateCustomTierSet", () => {
});
});
describe("buildComplexityRouterConfig stall escalation", () => {
it("omits all three keys when the toggle is off, since the backend rejects them next to session pinning", () => {
const config = buildComplexityRouterConfig({ ...baseParams, stallEscalationEnabled: false });
expect(config).not.toHaveProperty("stall_escalation_enabled");
expect(config).not.toHaveProperty("stall_escalation_window");
expect(config).not.toHaveProperty("stall_escalation_repeat_threshold");
});
it("emits the toggle and both knobs when it is on", () => {
const config = buildComplexityRouterConfig({
...baseParams,
stallEscalationEnabled: true,
stallEscalationWindow: 8,
stallEscalationRepeatThreshold: 4,
});
expect(config.stall_escalation_enabled).toBe(true);
expect(config.stall_escalation_window).toBe(8);
expect(config.stall_escalation_repeat_threshold).toBe(4);
});
it("emits the toggle alone when neither knob was touched, so both track the backend defaults", () => {
const config = buildComplexityRouterConfig({ ...baseParams, stallEscalationEnabled: true });
expect(config.stall_escalation_enabled).toBe(true);
expect(config).not.toHaveProperty("stall_escalation_window");
expect(config).not.toHaveProperty("stall_escalation_repeat_threshold");
});
});
describe("dryRunRejection", () => {
it("blocks the save on a rejection whose message is missing, which the write would return as a raw 400", () => {
expect(dryRunRejection({ valid: false })).toBe("The proxy rejected this auto-router configuration");

View file

@ -1,4 +1,7 @@
import { KeywordTierRule } from "./KeywordTierRules";
type ClassifierLLMConfigWire = ClassifierLLMConfig & { vision?: { enabled?: boolean; max_images?: number } };
import type { ModelGroup } from "../llm_calls/fetch_models";
import {
type CustomTierSet,
@ -61,7 +64,8 @@ export const normalizeClassifierLlmConfig = ({
reasoning_effort,
classification_rubric,
system_prompt,
}: ClassifierLLMConfig): ClassifierLLMConfig =>
vision,
}: ClassifierLLMConfigWire): ClassifierLLMConfigWire =>
system_prompt?.trim()
? {
model,
@ -69,6 +73,7 @@ export const normalizeClassifierLlmConfig = ({
...(circuit_breaker_enabled !== undefined && { circuit_breaker_enabled }),
...(circuit_breaker_cooldown_seconds !== undefined && { circuit_breaker_cooldown_seconds }),
...(reasoning_effort && { reasoning_effort }),
...(vision && { vision }),
system_prompt,
}
: {
@ -78,6 +83,7 @@ export const normalizeClassifierLlmConfig = ({
...(circuit_breaker_cooldown_seconds !== undefined && { circuit_breaker_cooldown_seconds }),
...(reasoning_effort && { reasoning_effort }),
...(classification_rubric && { classification_rubric }),
...(vision && { vision }),
};
interface ScorerKnobInputs {
@ -138,6 +144,9 @@ export interface BuildComplexityRouterConfigParams {
embeddingModel: string | undefined;
matchThreshold: number;
escalationKeywords: string[];
stallEscalationEnabled?: boolean;
stallEscalationWindow?: number;
stallEscalationRepeatThreshold?: number;
adaptive: boolean;
adaptiveWeights: AdaptiveRouterWeights;
tierDistancePenalty: number;
@ -199,6 +208,9 @@ export interface ComplexityRouterConfigPayload {
embedding_model?: string;
match_threshold?: number;
escalation_keywords?: string[];
stall_escalation_enabled?: boolean;
stall_escalation_window?: number;
stall_escalation_repeat_threshold?: number;
adaptive?: boolean;
adaptive_weights?: AdaptiveRouterWeights;
tier_distance_penalty?: number;
@ -318,7 +330,7 @@ export const getSemanticConfigError = ({
};
interface CustomTierWireFieldInputs {
classifierLlmConfig: ClassifierLLMConfig | undefined;
classifierLlmConfig: ClassifierLLMConfigWire | undefined;
planModeMinTierId: string | undefined;
classificationPrompt: string | undefined;
classificationExamples: string | undefined;
@ -350,6 +362,7 @@ export const customTierWireFields = (
circuit_breaker_cooldown_seconds: classifierLlmConfig.circuit_breaker_cooldown_seconds,
}),
...(classifierLlmConfig.reasoning_effort && { reasoning_effort: classifierLlmConfig.reasoning_effort }),
...(classifierLlmConfig.vision && { vision: classifierLlmConfig.vision }),
},
}),
session_affinity: false,
@ -472,6 +485,9 @@ export const buildComplexityRouterConfig = ({
embeddingModel,
matchThreshold,
escalationKeywords,
stallEscalationEnabled,
stallEscalationWindow,
stallEscalationRepeatThreshold,
adaptive,
adaptiveWeights,
tierDistancePenalty,
@ -541,6 +557,15 @@ export const buildComplexityRouterConfig = ({
...(customTechnicalKeywords.length > 0 && { custom_technical_keywords: customTechnicalKeywords }),
...(cleanedKeywordTierRules.length > 0 && { keyword_tier_rules: cleanedKeywordTierRules }),
escalation_keywords: cleanedEscalationKeywords,
// Only written when on: the backend rejects it alongside session_affinity, user_turn mode and
// a custom tier set, so an off router must not carry the key into any of those saves.
...(stallEscalationEnabled && {
stall_escalation_enabled: true,
...(stallEscalationWindow !== undefined && { stall_escalation_window: stallEscalationWindow }),
...(stallEscalationRepeatThreshold !== undefined && {
stall_escalation_repeat_threshold: stallEscalationRepeatThreshold,
}),
}),
...(semanticMatchingEnabled && {
semantic_keyword_matching: true,
embedding_model: embeddingModel,

View file

@ -113,6 +113,10 @@ export const CUSTOM_TIER_RESTRICTIONS = {
omit: ["escalation_keywords"],
reason: "Escalation bumps a request along the built-in tier ladder, which your tier set replaces",
},
stallEscalation: {
omit: ["stall_escalation_enabled", "stall_escalation_window", "stall_escalation_repeat_threshold"],
reason: "Stall escalation bumps a request along the built-in tier ladder, which your tier set replaces",
},
adaptive: {
omit: ["adaptive", "adaptive_weights", "tier_distance_penalty", "adaptive_eligible"],
reason: "Adaptive routing scores models along the built-in tier ladder, which your tier set replaces",

View file

@ -590,16 +590,54 @@ describe("managed keys survive an untouched open-and-save", () => {
// hold every managed key. Each gets its own round trip below.
const KEYS_ANOTHER_CLASSIFIER_TYPE_OWNS = new Set(["tier_definitions", "fallback_tier", "hybrid_boundary_margin"]);
// The stall keys are rejected beside the session pinning and user-turn classification this
// fixture sets, so they get their own round trip below rather than widening this one.
const KEYS_ANOTHER_CLASSIFICATION_FREQUENCY_OWNS = new Set([
"stall_escalation_enabled",
"stall_escalation_window",
"stall_escalation_repeat_threshold",
]);
it("carries every managed key a built-in router can hold through hydrate then save", () => {
const hydrated = hydrateComplexityRouterConfig(STORED_ALL_MANAGED, undefined);
const saved = buildUpdatedComplexityRouterConfig(STORED_ALL_MANAGED, hydrated);
const dropped = [...MANAGED_COMPLEXITY_ROUTER_KEYS]
.filter((key) => !KEYS_ANOTHER_CLASSIFIER_TYPE_OWNS.has(key))
.filter((key) => !KEYS_ANOTHER_CLASSIFICATION_FREQUENCY_OWNS.has(key))
.filter((key) => saved[key] === undefined);
expect(dropped).toEqual([]);
});
it("carries the stall-escalation keys through their own round trip", () => {
const stored: Record<string, unknown> = {
...STORED_ALL_MANAGED,
session_affinity: false,
classification_mode: "every_request",
stall_escalation_enabled: true,
stall_escalation_window: 8,
stall_escalation_repeat_threshold: 4,
};
const hydrated = hydrateComplexityRouterConfig(stored, undefined);
const saved = buildUpdatedComplexityRouterConfig(stored, hydrated);
expect(saved.stall_escalation_enabled).toBe(true);
expect(saved.stall_escalation_window).toBe(8);
expect(saved.stall_escalation_repeat_threshold).toBe(4);
});
it("leaves the stall keys out of a saved config that never had them on", () => {
const stored: Record<string, unknown> = {
...STORED_ALL_MANAGED,
session_affinity: false,
classification_mode: "every_request",
};
const hydrated = hydrateComplexityRouterConfig(stored, undefined);
const saved = buildUpdatedComplexityRouterConfig(stored, hydrated);
expect(saved).not.toHaveProperty("stall_escalation_enabled");
});
it("drops a stored local-scorer threshold when the operator converts the router to custom tiers", () => {
const hydrated = hydrateComplexityRouterConfig(STORED_ALL_MANAGED, undefined);
const converted = {

View file

@ -116,6 +116,9 @@ export interface StoredComplexityRouterConfig {
return_raw_model_name?: boolean;
enable_context_window_escalation?: unknown;
context_window_escalation_buffer?: unknown;
stall_escalation_enabled?: unknown;
stall_escalation_window?: unknown;
stall_escalation_repeat_threshold?: unknown;
}
/**
@ -213,6 +216,13 @@ export const hydrateComplexityRouterConfig = (
typeof parsedConfig.context_window_escalation_buffer === "number"
? parsedConfig.context_window_escalation_buffer
: undefined,
stall_escalation_enabled: parsedConfig.stall_escalation_enabled === true || undefined,
stall_escalation_window:
typeof parsedConfig.stall_escalation_window === "number" ? parsedConfig.stall_escalation_window : undefined,
stall_escalation_repeat_threshold:
typeof parsedConfig.stall_escalation_repeat_threshold === "number"
? parsedConfig.stall_escalation_repeat_threshold
: undefined,
};
};
@ -251,6 +261,9 @@ export const MANAGED_COMPLEXITY_ROUTER_KEYS = new Set([
"reasoning_override_min_score",
"enable_context_window_escalation",
"context_window_escalation_buffer",
"stall_escalation_enabled",
"stall_escalation_window",
"stall_escalation_repeat_threshold",
]);
// Managed only when the caller passes the corresponding state. A caller that does not render
@ -358,6 +371,9 @@ export const buildUpdatedComplexityRouterConfig = (
tierModelParams: value.tier_model_params,
enableContextWindowEscalation: value.enable_context_window_escalation,
contextWindowEscalationBuffer: value.context_window_escalation_buffer,
stallEscalationEnabled: value.stall_escalation_enabled,
stallEscalationWindow: value.stall_escalation_window,
stallEscalationRepeatThreshold: value.stall_escalation_repeat_threshold,
};
const built = buildComplexityRouterConfig(builderParams);

View file

@ -22764,22 +22764,40 @@ export interface components {
/** Spend */
spend?: number | null;
};
/**
* AccessGroupResource
* @description A resource referenced by an access group. `name` is null when the id no longer resolves or has no alias.
*/
AccessGroupResource: {
/** Id */
id: string;
/** Name */
name: string | null;
};
/** AccessGroupResponse */
AccessGroupResponse: {
/** Access Agent Ids */
access_agent_ids: string[];
/** Access Agents */
access_agents: components["schemas"]["AccessGroupResource"][];
/** Access Group Id */
access_group_id: string;
/** Access Group Name */
access_group_name: string;
/** Access Mcp Server Ids */
access_mcp_server_ids: string[];
/** Access Mcp Servers */
access_mcp_servers: components["schemas"]["AccessGroupResource"][];
/** Access Model Names */
access_model_names: string[];
/** Assigned Key Ids */
assigned_key_ids: string[];
/** Assigned Keys */
assigned_keys: components["schemas"]["AccessGroupResource"][];
/** Assigned Team Ids */
assigned_team_ids: string[];
/** Assigned Teams */
assigned_teams: components["schemas"]["AccessGroupResource"][];
/**
* Created At
* Format: date-time
@ -25305,6 +25323,30 @@ export interface components {
* @default 3000
*/
timeout_ms: number;
/** @description Whether the classifier sees images on the request, and how many */
vision?: components["schemas"]["ClassifierVisionConfig"];
};
/**
* ClassifierVisionConfig
* @description Whether the LLM classifier sees the images on the request it is classifying.
*
* Off by default because images cost far more than the text ask they arrive with, and the
* classifier runs on every request. A turn whose complexity lives in the image ("what is wrong in
* this stack trace screenshot") is invisible to a text-only classifier, which is what this buys.
*/
ClassifierVisionConfig: {
/**
* Enabled
* @description Forward image content to the classifier. Requires a classifier model declared supports_vision, on the deployment's model_info or in the model cost map; images stay stripped otherwise, so a classifier that cannot read them is never sent one. Declare model_info.supports_vision on the deployment to enable a model the cost map does not describe. Only inline data: URIs are forwarded. A request whose images are http(s) URLs still classifies on its text alone, because some providers fetch such a URL from the proxy rather than the provider, which would let a caller aim a proxy-side request at an address of their choosing.
* @default false
*/
enabled: boolean;
/**
* Max Images
* @description How many images from the newest user turn to forward, in wire order. Bounds the added cost of a turn that attaches many images. Images on earlier turns are never forwarded.
* @default 1
*/
max_images: number;
};
/**
* CloudZeroExportRequest
@ -34923,6 +34965,24 @@ export interface components {
* @description Keywords indicating simple/basic queries
*/
simple_keywords?: string[] | null;
/**
* Stall Escalation Enabled
* @description Escalate mid-task to the next-higher configured tier when the assistant's own recent tool calls look stuck: the newest tool call repeats, or errors, at least stall_escalation_repeat_threshold times across the last stall_escalation_window calls. Both tests are anchored on the newest call, so a task that tried the same thing a few times and then moved on is not escalated on the strength of those older calls alone, while a retry loop broken up by an unrelated lookup still counts. One tier at most, on the same ladder escalation_keywords bumps along, and never above the highest configured tier. Detection re-runs on every classified turn from the tool calls visible in that request, so it needs no state and nothing survives past the task. Mutually exclusive with session_affinity and classification_mode='user_turn', which both replay a held routing decision instead of classifying most turns, so this would never see the tool calls to look at. Off by default.
* @default false
*/
stall_escalation_enabled: boolean;
/**
* Stall Escalation Repeat Threshold
* @description How many of the last stall_escalation_window tool calls must repeat the newest call, or must have errored alongside it, before the task counts as stalled. Must not exceed stall_escalation_window, or the condition could never be reached.
* @default 3
*/
stall_escalation_repeat_threshold: number;
/**
* Stall Escalation Window
* @description How many of the assistant's most recent tool calls stall detection looks at, oldest ones dropped as new calls happen. Counted across the whole visible conversation rather than reset at the newest human ask, so evidence from before a plain follow-up message like 'try again' is still visible on the turn after it.
* @default 6
*/
stall_escalation_window: number;
/**
* Technical Keywords
* @description Keywords indicating technical content