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

This commit is contained in:
moe-berri 2026-09-04 20:54:26 -07:00
commit 955baf8a5c
54 changed files with 2637 additions and 332 deletions

View file

@ -0,0 +1 @@
ALTER TABLE "LiteLLM_ShadowEvalJob" ADD COLUMN IF NOT EXISTS "models" TEXT[] NOT NULL DEFAULT ARRAY[]::TEXT[];

View file

@ -1536,6 +1536,7 @@ model LiteLLM_ShadowEvalJob {
target_id String // hashed virtual key, team_id, or user_id whose traffic this leg shadows
router_name String // first (often only) auto-router under evaluation; router_names is the full set
router_names String[] @default([]) // all routers this job runs as shadow arms; empty on legacy rows, whose set is (router_name)
models String[] @default([]) // model groups the sampled traffic is narrowed to; empty samples every model
direction String @default("forward") // forward | reverse
baseline_model String? // reverse only: the fixed model the router is judged against
judge_model String

View file

@ -35,6 +35,15 @@ span orphaned into its own trace). The anchor — a contextvar inherited by thos
child tasks — gives a stable parent in both cases. DB/service spans keep ambient
parenting so an auth DB lookup still nests under `auth`.
The anchor is also what `litellm.request.route` is read from: `request_root_http_route`
returns the server span's own `http.route`, so the LLM call span cannot disagree with
its parent about which endpoint served the request. That means the route template on a
normal route and the literal path on a passthrough prefix, because the passthrough hook
rewrote the attribute; an MCP call anchors the same server span, so it reports the
`/mcp` mount point. Attributes stay readable after a span ends, so the async close
callback reads the same value. Where no server span was anchored at all, the route the
proxy recorded at auth (`metadata.user_api_key_request_route`) is the backstop.
**Which service calls become spans (`spans.span_role_for_service`).** LiteLLM's
service-logging layer instruments many internal functions, but only some are
traceable units of work:

View file

@ -49,6 +49,7 @@ from litellm.integrations.otel.model.utils import to_ns
from litellm.integrations.otel.plumbing.context import (
is_recordable_span,
mcp_message_transport_span,
request_root_http_route,
request_root_span,
resolve_mcp_span_context,
resolve_parent_context,
@ -541,6 +542,7 @@ class OpenTelemetryV2(CustomLogger):
payload,
capture_content=self.config.capture_span_content,
time_to_first_chunk_seconds=call.time_to_first_chunk_seconds,
request_route=request_root_http_route(),
)
end_time_ns: Final = to_ns(end_time)
if carrier is not None and carrier.span is not None:

View file

@ -89,6 +89,7 @@ class GenAIMapper:
f"{LiteLLM.COST_PREFIX}margin_percent": lambda d: d.cost.margin_percent,
f"{LiteLLM.COST_PREFIX}margin_total_amount": lambda d: d.cost.margin_total_amount,
LiteLLM.REQUEST_STREAMING: lambda d: d.is_streaming,
LiteLLM.REQUEST_ROUTE: lambda d: d.request_route,
}
_TOOL_ATTRS: dict[str, Callable[[ToolDefinition], AttrValue | None]] = {

View file

@ -64,6 +64,7 @@ class RequestIdentity:
# completes (routing has picked a deployment), so it's absent from the
# auth-time seed and filled only from the payload.
provider_model: str | None = None
request_route: str | None = None
metadata: Mapping[str, str] = field(default_factory=dict)
@classmethod
@ -87,6 +88,7 @@ class RequestIdentity:
key_hash=as_str(raw_meta.get("user_api_key_hash")),
end_user=as_str(payload.get("end_user")) or as_str(raw_meta.get("user_api_key_end_user_id")),
provider_model=resolve_provider_model(payload),
request_route=as_str(raw_meta.get("user_api_key_request_route")),
metadata=metadata,
)

View file

@ -386,6 +386,7 @@ class LLMCallSpanData:
# keeps routes the convention folds into one operation distinguishable.
output_type: GenAIOutputType | None = None
call_type: str | None = None
request_route: str | None = None
@classmethod
def from_standard_logging_payload(
@ -393,6 +394,7 @@ class LLMCallSpanData:
payload: StandardLoggingPayload,
capture_content: bool = False,
time_to_first_chunk_seconds: float | None = None,
request_route: str | None = None,
) -> LLMCallSpanData:
params: Final = cast(Mapping[str, object], payload.get("model_parameters") or {})
# The single parse of the request's metadata — the request-vs-provider
@ -433,6 +435,7 @@ class LLMCallSpanData:
time_to_first_chunk_seconds=time_to_first_chunk_seconds,
output_type=resolve_output_type(call_type),
call_type=call_type or None,
request_route=request_route or context.identity.request_route,
)

View file

@ -295,6 +295,7 @@ class LiteLLM:
# ``litellm_params.model``), distinct from the user-facing ``gen_ai.request.model``.
PROVIDER_MODEL: Final = "litellm.provider.model"
REQUEST_STREAMING: Final = "litellm.request.streaming"
REQUEST_ROUTE: Final = "litellm.request.route"
TOOLS_DECLARED: Final = "litellm.request.tools.declared"
GUARDRAIL_NAME: Final = "litellm.guardrail.name"
GUARDRAIL_MODE: Final = "litellm.guardrail.mode"

View file

@ -6,6 +6,7 @@ from typing import Final
from opentelemetry import baggage
from opentelemetry.context import Context, get_current
from opentelemetry.sdk.trace import ReadableSpan
from opentelemetry.trace import (
Link,
NonRecordingSpan,
@ -18,6 +19,8 @@ from opentelemetry.trace.propagation.tracecontext import (
TraceContextTextMapPropagator,
)
from litellm.integrations.otel.model.semconv import HTTP
_PROPAGATOR: Final = TraceContextTextMapPropagator()
# The request's root span — the FastAPI-owned SERVER span — captured ONCE when the
@ -55,6 +58,25 @@ def request_root_span() -> "Span | None":
return span if is_recordable_span(span) else None
def request_root_http_route() -> str | None:
"""``http.route`` exactly as the request's root SERVER span reports it.
Read off the span rather than re-derived, so the LLM call span cannot disagree
with its own parent about which endpoint served the request: the template the
instrumentation matched, or the literal path where
``mount._passthrough_span_name_hook`` rewrote it, are already in the attribute.
An MCP call anchors that same server span, so it reports the ``/mcp`` mount
point the instrumentation matched. Attributes stay readable after a span ends,
so this answers just as well from the async logging callback.
None when no server span is anchored, which is the SDK path and any deployment
where the FastAPI instrumentation did not mount.
"""
span: Final = request_root_span()
route: Final = span.attributes.get(HTTP.ROUTE) if isinstance(span, ReadableSpan) and span.attributes else None
return route if isinstance(route, str) and route else None
# The W3C trace-context carrier (``traceparent``/``tracestate``/``baggage``) the
# MCP client propagated in the current request's ``params._meta``. The MCP gateway
# sets it per message so the MCP span can record the client's span as a span

View file

@ -37,6 +37,7 @@ from litellm.litellm_core_utils.llm_judge import (
)
from litellm.litellm_core_utils.redact_messages import should_redact_message_logging
from litellm.llms.base_llm.base_utils import type_to_response_format_param
from litellm.router_utils.common_utils import resolve_model_group_alias
from litellm.types.management_endpoints.auto_router_endpoints import ShadowEvalDirection
from litellm.types.utils import SHADOW_EVAL_JUDGE_CALL_ORIGIN, SHADOW_EVAL_ROUTER_CALL_ORIGIN
@ -664,6 +665,7 @@ class ActiveShadowEvalJob(BaseModel):
id: str
router_name: str
router_names: tuple[str, ...] = ()
models: frozenset[str] = frozenset()
direction: ShadowEvalDirection = "forward"
baseline_model: str | None = None
shadow_percentage: float
@ -706,6 +708,21 @@ class ActiveShadowEvalJob(BaseModel):
return self.baseline_model or arm_router
def _canonical_group(router: "Router | None", model_group: str) -> str:
"""A model group in the one spelling both a job's scope and a request's model compare
under: an alias resolves to its target so the two never fail to match on spelling."""
return (
resolve_model_group_alias(router.model_group_alias, model_group) if router is not None else None
) or model_group
def _scope_admits(router: "Router | None", job: "ActiveShadowEvalJob", model_group: str) -> bool:
"""Whether the request's group is in the job's model scope. Both sides resolve through
the router's alias map at match time, so a re-pointed alias applies to the next request
rather than after the jobs cache rolls."""
return not job.models or any(_canonical_group(router, name) == model_group for name in job.models)
def _as_active_job(record: object, attempts: int, spend: float) -> ActiveShadowEvalJob | None:
"""The sampling path's view of one job row, or None for a row it cannot sample: an
unknown direction, or a reverse job with no baseline model to duplicate against.
@ -728,7 +745,8 @@ class ShadowEvalLogger(CustomLogger):
A job targets a virtual key, a team, or a user; a request qualifies for a job when
any of its resolved identities (key hash, team id, user id) matches the job's
target, so team and user jobs cover JWT-authenticated traffic, which carries no
key hash at all."""
key hash at all. A job scoped to model groups further requires the request's
requested group to be one of them."""
def __init__(
self,
@ -815,19 +833,24 @@ class ShadowEvalLogger(CustomLogger):
active_jobs: Sequence[ActiveShadowEvalJob],
request_metadata: Mapping[str, object],
request_id: str,
model_group: str,
) -> tuple[ActiveShadowEvalJob, ...]:
"""The jobs that sample this request. A key can hold one job per direction, and a
request routed by one job's router while bypassing the other's qualifies for both;
each is separately budgeted, so both fire. An admitting job that loses the sampling
dice is counted, so results can weigh judged rows against the traffic they stand for."""
dice is counted, so results can weigh judged rows against the traffic they stand for.
A request outside a job's direction or model scope is not that job's traffic and
goes uncounted, so the funnel stays a fraction of the traffic the job admits."""
eligible: list[ActiveShadowEvalJob] = [] # mutable-ok: bucketed per-job admission
now: Final = datetime.now(timezone.utc)
router: Final = self._router_provider()
for job in active_jobs:
if (
now >= job.ends_at
or job.attempts + self._job_starts.get(job.id, 0) >= job.max_turns
or (job.max_budget is not None and job.spend >= job.max_budget)
or not _direction_admits(request_metadata, job)
or not _scope_admits(router, job, model_group)
):
continue
if not _sample_hits(request_id, job.id, job.shadow_percentage):
@ -882,6 +905,7 @@ class ShadowEvalLogger(CustomLogger):
tuple(job for target in targets for job in active_jobs.get(target, ())),
request_metadata,
request_id,
_canonical_group(self._router_provider(), str(payload.get("model_group") or "")),
)
if not eligible:
return

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

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

@ -0,0 +1,372 @@
"""`lite debug claude`: one-shot debug report for a Claude Code session routed through the proxy.
Claude Code puts its session id in `metadata.user_id`, which the proxy lifts into
`LiteLLM_SpendLogs.session_id`. This command pulls every turn of that session, plus
the request / response bodies for failures and the most recent turns, and renders a
single markdown report that can be pasted into a bug report or handed to another agent.
"""
import json
import os
import re
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from types import MappingProxyType
from typing import Final
import click
import requests
from pydantic import BaseModel, ConfigDict, Field, JsonValue, TypeAdapter, ValidationError, field_validator
from ...http_client import HTTPClient
from ._cli_context import cli_context_values
CLAUDE_DIR: Final = Path.home() / ".claude"
REPORT_DIR: Final = Path.home() / ".litellm" / "debug"
SESSION_ID_ENV: Final = "CLAUDE_CODE_SESSION_ID"
SLASH_COMMAND_NAME: Final = "debug-lite"
SLASH_COMMAND_BODY: Final = """---
description: Pull the LiteLLM debug report (spend, request, response, error) for this Claude Code session
allowed-tools: Bash(lite debug claude:*)
---
Below is the LiteLLM debug report for this Claude Code session. Summarize the failing
request(s) in a few sentences (model, error, request id) and tell me the path the full
report was saved to so I can hand it off. If nothing failed, say so.
!`lite debug claude $ARGUMENTS`
"""
@dataclass(frozen=True, slots=True)
class DebugFailure:
message: str
class ErrorInformation(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
error_code: str | None = None
error_class: str | None = None
error_message: str | None = None
llm_provider: str | None = None
class SpendLogMetadata(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
status: str | None = None
error_information: ErrorInformation | None = None
class SpendLogRow(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore", populate_by_name=True)
request_id: str
start_time: str | None = Field(default=None, alias="startTime")
end_time: str | None = Field(default=None, alias="endTime")
model: str | None = None
model_group: str | None = None
custom_llm_provider: str | None = None
api_base: str | None = None
call_type: str | None = None
status: str | None = None
spend: float = 0.0
prompt_tokens: int = 0
completion_tokens: int = 0
total_tokens: int = 0
metadata: SpendLogMetadata = SpendLogMetadata()
@field_validator("metadata", mode="before")
@classmethod
def _parse_metadata(cls, value: object) -> object:
if value is None:
return SpendLogMetadata()
if isinstance(value, str):
return json.loads(value) if value else SpendLogMetadata()
return value
@property
def failed(self) -> bool:
return (self.status or self.metadata.status) == "failure"
@property
def error(self) -> ErrorInformation | None:
return self.metadata.error_information
class SessionLogsPage(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
data: tuple[SpendLogRow, ...]
total: int
total_pages: int
class RequestResponsePayload(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
proxy_server_request: JsonValue = None
response: JsonValue = None
messages: JsonValue = None
_SESSION_PAGE: Final = TypeAdapter(SessionLogsPage)
_PAYLOAD: Final[TypeAdapter[RequestResponsePayload | None]] = TypeAdapter(RequestResponsePayload | None)
_JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
_SESSION_PAGE_SIZE: Final = 100
_TRANSPORT_BODY_CHARS: Final = 500
_SESSION_TRANSCRIPT_STEM: Final = re.compile(r"[0-9a-f]{8}(?:-[0-9a-f]{4}){3}-[0-9a-f]{12}")
def detect_claude_session_id(env: Mapping[str, str], claude_dir: Path) -> str | None:
explicit: Final = env.get(SESSION_ID_ENV)
if explicit:
return explicit
transcripts: Final = tuple(
path for path in claude_dir.glob("projects/*/*.jsonl") if _SESSION_TRANSCRIPT_STEM.fullmatch(path.stem)
)
if not transcripts:
return None
newest: Final = max(transcripts, key=lambda p: p.stat().st_mtime)
return newest.stem
def _transport_failure(uri: str, error: requests.exceptions.RequestException) -> DebugFailure:
body: Final = error.response.text[:_TRANSPORT_BODY_CHARS] if error.response is not None else ""
detail: Final = f"\n{body}" if body else ""
return DebugFailure(f"GET {uri} failed: {error}{detail}")
class SpendLogsFetcher:
def __init__(self, http: HTTPClient) -> None:
self._http = http
def session_rows(self, session_id: str) -> tuple[SpendLogRow, ...] | DebugFailure:
first: Final = self._page(session_id, 1)
if isinstance(first, DebugFailure):
return first
rest: Final = tuple(self._page(session_id, page) for page in range(2, first.total_pages + 1))
failed_page: Final = next((page for page in rest if isinstance(page, DebugFailure)), None)
if failed_page is not None:
return failed_page
rows: Final = first.data + tuple(row for page in rest if isinstance(page, SessionLogsPage) for row in page.data)
return tuple(sorted(rows, key=lambda r: r.start_time or ""))
def _get(self, uri: str, params: Mapping[str, str | int] | None = None) -> JsonValue | DebugFailure:
try:
return _JSON.validate_python(self._http.request("GET", uri, params=params)) # pyright: ignore[reportUnknownMemberType] # HTTPClient.request is untyped
except requests.exceptions.RequestException as e:
return _transport_failure(uri, e)
def _page(self, session_id: str, page: int) -> SessionLogsPage | DebugFailure:
uri: Final = "/spend/logs/session/ui"
raw: Final = self._get(
uri, MappingProxyType({"session_id": session_id, "page": page, "page_size": _SESSION_PAGE_SIZE})
)
if isinstance(raw, DebugFailure):
return raw
try:
return _SESSION_PAGE.validate_python(raw)
except ValidationError as e:
return DebugFailure(f"Unexpected {uri} response: {e}")
def payload(self, request_id: str) -> RequestResponsePayload | None | DebugFailure:
uri: Final = f"/spend/logs/ui/{request_id}"
raw: Final = self._get(uri)
if isinstance(raw, DebugFailure):
return raw
try:
return _PAYLOAD.validate_python(raw)
except ValidationError as e:
return DebugFailure(f"Unexpected {uri} response: {e}")
def _fmt_json(value: JsonValue, max_chars: int) -> str:
text: Final = value if isinstance(value, str) else json.dumps(value, indent=2, default=str)
if len(text) <= max_chars:
return text
return f"{text[:max_chars]}\n... (truncated, {len(text) - max_chars} more chars)"
def _fenced(text: str, info: str = "") -> tuple[str, str, str]:
longest_run: Final = max((len(run) for run in re.findall(r"`+", text)), default=0)
fence: Final = "`" * max(3, longest_run + 1)
return (f"{fence}{info}", text, fence)
def _row_section(row: SpendLogRow, index: int, payload: RequestResponsePayload | None, max_chars: int) -> str:
err: Final = row.error
error_lines: Final = (
(
f"- error: `{err.error_code or '?'}` {err.error_class or ''}".rstrip(),
"",
*_fenced(err.error_message or ""),
)
if err is not None and row.failed
else ()
)
body_lines: Final = (
(
"",
"<details><summary>request body</summary>",
"",
*_fenced(_fmt_json(payload.proxy_server_request, max_chars), "json"),
"</details>",
"",
"<details><summary>response</summary>",
"",
*_fenced(_fmt_json(payload.response, max_chars), "json"),
"</details>",
)
if payload is not None
else ()
)
header: Final = f"### {index}. {'FAILED' if row.failed else 'ok'} {row.model or row.model_group or '?'}"
facts: Final = (
f"- request_id: `{row.request_id}`",
f"- time: {row.start_time} -> {row.end_time}",
f"- provider: {row.custom_llm_provider or '?'} ({row.api_base or 'n/a'}), call_type: {row.call_type or '?'}",
f"- spend: ${row.spend:.6f}, tokens: {row.prompt_tokens} in / {row.completion_tokens} out",
)
return "\n".join((header, *facts, *error_lines, *body_lines))
def render_report(
*,
session_id: str,
base_url: str,
rows: Sequence[SpendLogRow],
payloads: Mapping[str, RequestResponsePayload | None],
max_chars: int,
) -> str:
failures: Final = tuple(r for r in rows if r.failed)
summary: Final = (
f"# LiteLLM debug report: Claude Code session `{session_id}`",
"",
f"- proxy: {base_url}",
f"- generated: {datetime.now(timezone.utc).isoformat(timespec='seconds')}",
f"- turns: {len(rows)}, failed: {len(failures)}",
f"- total spend: ${sum(r.spend for r in rows):.6f}",
f"- models: {', '.join(sorted(frozenset(r.model or r.model_group or '?' for r in rows))) or 'n/a'}",
"",
"Bodies are included for failed turns and the most recent turns. "
"Bodies are empty unless the proxy runs with `general_settings.store_prompts_in_spend_logs: true`.",
"",
"## Turns",
"",
)
sections: Final = tuple(
_row_section(row, i, payloads.get(row.request_id), max_chars) for i, row in enumerate(rows, start=1)
)
return "\n".join(summary) + "\n\n".join(sections) + "\n"
def build_report(
*,
fetcher: SpendLogsFetcher,
session_id: str,
base_url: str,
recent_bodies: int,
max_chars: int,
) -> str | DebugFailure:
rows: Final = fetcher.session_rows(session_id)
if isinstance(rows, DebugFailure):
return rows
if not rows:
return DebugFailure(
f"No spend logs found for session {session_id!r} on {base_url}. "
"Is Claude Code routed through this proxy (`lite up`), and does your key have log access?"
)
wanted: Final = frozenset(r.request_id for r in rows if r.failed) | frozenset(
r.request_id for r in rows[-recent_bodies:] if recent_bodies > 0
)
fetched: Final = MappingProxyType({rid: fetcher.payload(rid) for rid in sorted(wanted)})
failed_payload: Final = next((p for p in fetched.values() if isinstance(p, DebugFailure)), None)
if failed_payload is not None:
return failed_payload
payloads: Final = MappingProxyType({rid: p for rid, p in fetched.items() if not isinstance(p, DebugFailure)})
return render_report(session_id=session_id, base_url=base_url, rows=rows, payloads=payloads, max_chars=max_chars)
def write_report(report: str, session_id: str, report_dir: Path) -> Path:
report_dir.mkdir(parents=True, exist_ok=True)
path: Final = report_dir / f"claude-{session_id}.md"
path.write_text(report, encoding="utf-8")
path.chmod(0o600)
return path
def install_slash_command(claude_dir: Path) -> Path:
commands_dir: Final = claude_dir / "commands"
commands_dir.mkdir(parents=True, exist_ok=True)
path: Final = commands_dir / f"{SLASH_COMMAND_NAME}.md"
path.write_text(SLASH_COMMAND_BODY, encoding="utf-8")
return path
@click.group()
def debug() -> None:
"""Pull debug reports (spend, request, response, error) for coding-agent sessions"""
@debug.command("claude")
@click.option(
"--session-id",
default=None,
help=f"Claude Code session id. Defaults to ${SESSION_ID_ENV}, else the most recently used transcript in ~/.claude",
)
@click.option(
"--recent-bodies",
default=3,
show_default=True,
type=click.IntRange(min=0),
help="Also include request/response bodies for the N most recent turns (failed turns always get bodies)",
)
@click.option(
"--max-body-chars",
default=20_000,
show_default=True,
type=click.IntRange(min=100),
help="Truncate each request/response body to this many characters",
)
@click.option("--no-save", is_flag=True, help="Print only, do not write the report under ~/.litellm/debug")
@click.pass_context
def debug_claude(
ctx: click.Context, session_id: str | None, recent_bodies: int, max_body_chars: int, no_save: bool
) -> None:
"""Render a markdown debug report for one Claude Code session routed through the proxy
Examples:
lite debug claude
lite debug claude --session-id e96634a3-fa28-4083-b354-55542e2dca01
"""
resolved: Final = session_id or detect_claude_session_id(os.environ, CLAUDE_DIR)
if resolved is None:
raise click.ClickException(f"Could not find a Claude Code session. Pass --session-id or set ${SESSION_ID_ENV}.")
values: Final = cli_context_values(ctx)
base_url: Final = values["base_url"]
fetcher: Final = SpendLogsFetcher(HTTPClient(base_url, values["api_key"]))
outcome: Final = build_report(
fetcher=fetcher,
session_id=resolved,
base_url=base_url,
recent_bodies=recent_bodies,
max_chars=max_body_chars,
)
if isinstance(outcome, DebugFailure):
raise click.ClickException(outcome.message)
click.echo(outcome)
if not no_save:
path: Final = write_report(outcome, resolved, REPORT_DIR)
click.echo(f"Saved to {path}", err=True)
@debug.command("install-claude-command")
def debug_install_claude_command() -> None:
"""Install the /debug-lite slash command into ~/.claude/commands so Claude Code can run `lite debug claude`"""
path: Final = install_slash_command(CLAUDE_DIR)
click.echo(f"Installed /{SLASH_COMMAND_NAME}: {path}")
click.echo("Restart Claude Code (or start a new session), then type /debug-lite.")

View file

@ -14,6 +14,7 @@ from .commands.autoroute.commands import autoroute_group
from .commands.chat import chat
from .commands.config import config_commands, get_config_value, hidden_command_names
from .commands.credentials import credentials
from .commands.debug import debug
from .commands.encryption import encryption
from .commands.http import http
from .commands.keys import keys
@ -143,6 +144,7 @@ cli.add_command(encryption)
cli.add_command(chat)
# Add the http command group
cli.add_command(http)
cli.add_command(debug)
# Add the keys command group
cli.add_command(keys)
# Add the teams command group

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

@ -789,6 +789,26 @@ def _for_teams(team_ids: Sequence[str | None]) -> str:
return f" for team {', '.join(named)}" if named else ""
def _validate_model_scope(llm_router: "Router | None", models: Sequence[str]) -> None:
"""Reject a scope naming a model no request on this proxy could carry, at start rather
than as a job that silently samples nothing. The question is "could any caller ask for
this name", not "does it resolve for the job's teams": a user target's traffic can arrive
on any team's key, so a team-public name is a legitimate scope for it, and an auto-router
is one too (a forward job on router A scoped to router B samples what B serves today).
Nothing here is ever dispatched to."""
unreachable: Final = tuple(
model
for model in models
if judge_target(llm_router, model).via == "nothing"
and (llm_router is None or model not in llm_router.team_public_model_names)
)
if unreachable:
raise HTTPException(
status_code=400,
detail="models not served by this proxy: " + ", ".join(f"'{model}'" for model in unreachable),
)
_JUDGED_ROLES: Final[frozenset[StrategyRouterDependencyRole]] = frozenset({"tier", "default"})
@ -1080,6 +1100,7 @@ class _LegRow(BaseModel):
target_id: str
router_name: str
router_names: tuple[str, ...] = ()
models: tuple[str, ...] = ()
direction: ShadowEvalDirection
baseline_model: str | None = None
judge_model: str
@ -1150,6 +1171,7 @@ def _group_response(
for leg in sorted(legs, key=lambda leg: (leg.target_type, leg.target_id))
),
router_names=first.arm_router_names,
models=first.models,
direction=first.direction,
baseline_model=first.baseline_model,
judge_model=first.judge_model,
@ -1322,7 +1344,10 @@ async def start_shadow_eval(
A target is a virtual key, a team, or a user. Team and user targets match on the
identity every request resolves to at auth time, so they cover JWT-authenticated
traffic, which presents no virtual key; a user target samples that user's traffic
across all their teams, whether it arrives on a JWT or a key they own.
across all their teams, whether it arrives on a JWT or a key they own. models narrows
every target to requests for those model groups, so a user plus one model samples that
user's traffic on that model across every key they own; it is forward-only, since a
reverse job already samples exactly the traffic its own router served.
A forward job answers whether the targets should adopt router_name: it samples the
requests the router did not serve and duplicates them through it. A reverse job
@ -1411,6 +1436,7 @@ async def start_shadow_eval(
if data.baseline_model is not None:
_validate_plain_model(llm_router, data.baseline_model, "baseline_model", team_ids)
_validate_judge_is_not_a_candidate(llm_router, data, team_ids)
_validate_model_scope(llm_router, data.models)
requested_targets: Final[tuple[tuple[ShadowEvalTargetType, str], ...]] = (
*(("key", key) for key in data.api_key_ids),
@ -1456,6 +1482,7 @@ async def start_shadow_eval(
# a pre-router_names pod samples router_name alone, so it must be a real arm
"router_name": data.router_names[0],
"router_names": list(data.router_names), # mutable-ok: Prisma payload
"models": list(data.models), # mutable-ok: Prisma payload
"direction": data.direction,
"baseline_model": data.baseline_model,
"judge_model": data.judge_model,
@ -1517,6 +1544,7 @@ async def start_shadow_eval(
for target_type, target_id in sorted(requested_targets)
),
router_names=data.router_names,
models=data.models,
direction=data.direction,
baseline_model=data.baseline_model,
judge_model=data.judge_model,

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

@ -1536,6 +1536,7 @@ model LiteLLM_ShadowEvalJob {
target_id String // hashed virtual key, team_id, or user_id whose traffic this leg shadows
router_name String // first (often only) auto-router under evaluation; router_names is the full set
router_names String[] @default([]) // all routers this job runs as shadow arms; empty on legacy rows, whose set is (router_name)
models String[] @default([]) // model groups the sampled traffic is narrowed to; empty samples every model
direction String @default("forward") // forward | reverse
baseline_model String? // reverse only: the fixed model the router is judged against
judge_model String

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,
@ -435,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.
@ -1592,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)
@ -1599,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:
@ -1623,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,
@ -1865,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
@ -2558,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."""
@ -3374,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:
@ -3402,6 +3495,7 @@ 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
@ -3439,7 +3533,7 @@ 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:
keyword_bumped_tier: Final = (
self._escalate_tier(override.tier) if escalation_keyword is not None else override.tier
@ -3486,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
@ -3558,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=(

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

@ -292,6 +292,18 @@ class StartShadowEvalRequest(BaseModel):
"to across all their teams: JWT requests carrying their subject claim and virtual keys they own"
),
)
models: tuple[str, ...] = Field(
default=(),
max_length=100,
description=(
"Model groups to narrow the sampled traffic to, matched on the group the caller "
"requested and resolved through model_group_alias, so an alias and its target are one "
"name. Empty samples every model the targets use. This ANDs with the targets: a job "
"over a user and one model samples that user's requests on that model across every key "
"they own, and none of their other traffic. Forward jobs only: a reverse job samples "
"exactly the traffic its own router served, which no other model group can name"
),
)
router_name: str | None = Field(
default=None,
description=(
@ -372,12 +384,20 @@ class StartShadowEvalRequest(BaseModel):
def _round_percentage(cls, value: float) -> float:
return round(value, 2)
@field_validator("api_key_ids", "team_ids", "user_ids")
@field_validator("api_key_ids", "team_ids", "user_ids", "models")
@classmethod
def _dedupe_targets(cls, value: tuple[str, ...]) -> tuple[str, ...]:
"""A target named twice would collide with itself on the one-active-per-(target, direction) index."""
"""A target named twice would collide with itself on the one-active-per-(target, direction)
index; a model named twice is one scope entry."""
return tuple(dict.fromkeys(value))
@field_validator("models")
@classmethod
def _models_are_names(cls, value: tuple[str, ...]) -> tuple[str, ...]:
if not all(name.strip() for name in value):
raise ValueError("models must be non-empty model group names")
return value
@model_validator(mode="after")
def _at_least_one_target_at_most_hundred(self) -> "StartShadowEvalRequest":
total: Final = len(self.api_key_ids) + len(self.team_ids) + len(self.user_ids)
@ -387,6 +407,18 @@ class StartShadowEvalRequest(BaseModel):
raise ValueError("at most 100 targets per job across api_key_ids, team_ids, and user_ids")
return self
@model_validator(mode="after")
def _model_scope_is_forward_only(self) -> "StartShadowEvalRequest":
"""A reverse job admits exactly the requests its own router served, so every one of
them names that router and nothing else; any other scope would sample nothing and
the router itself is a no-op. Both readings are rejected rather than shipped as a
job that silently never samples."""
if self.models and self.direction == "reverse":
raise ValueError(
"models is only meaningful for a forward job; a reverse job samples its own router's traffic"
)
return self
@model_validator(mode="after")
def _baseline_model_matches_direction(self) -> "StartShadowEvalRequest":
if self.direction == "reverse" and self.baseline_model is None:
@ -599,6 +631,10 @@ class ShadowEvalJobResponse(BaseModel):
"traffic and judge every arm against the same real responses"
),
)
models: tuple[str, ...] = Field(
default=(),
description="Model groups the sampled traffic is narrowed to; empty means every model the targets use",
)
direction: ShadowEvalDirection = "forward"
baseline_model: str | None = None
judge_model: str

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

@ -1536,6 +1536,7 @@ model LiteLLM_ShadowEvalJob {
target_id String // hashed virtual key, team_id, or user_id whose traffic this leg shadows
router_name String // first (often only) auto-router under evaluation; router_names is the full set
router_names String[] @default([]) // all routers this job runs as shadow arms; empty on legacy rows, whose set is (router_name)
models String[] @default([]) // model groups the sampled traffic is narrowed to; empty samples every model
direction String @default("forward") // forward | reverse
baseline_model String? // reverse only: the fixed model the router is judged against
judge_model String

View file

@ -142,6 +142,8 @@ fi
lint_dashboard() {
(
trap 'exit 143' TERM
trap 'rm -f "${report:-}"' EXIT
rc=0
prettier_rel=()
eslint_rel=()
@ -168,7 +170,6 @@ EOF
report=$(mktemp)
npx eslint . -f json -o "$report" || true
node scripts/check-lint-budgets.mjs "$report" eslint-budgets.json || rc=1
rm -f "$report"
exit $rc
)
}

View file

@ -157,6 +157,9 @@ async def test_chat_completion_bad_model_with_spend_logs():
except json.JSONDecodeError:
print(f"Could not parse response body as JSON: {response.text}")
assert (
response.status_code == 400
), f"expected HTTP 400, got {response.status_code}: {response.text}"
assert (
litellm_call_id is not None
), "Failed to get LiteLLM Call ID from response headers"
@ -191,7 +194,7 @@ async def test_chat_completion_bad_model_with_spend_logs():
# Verify the structure of the log entry
assert log_entry["request_id"] == litellm_call_id
assert log_entry["model"] == "non-existent-model"
assert log_entry["model_group"] == "non-existent-model"
assert log_entry["model_group"] in ("", "non-existent-model")
assert log_entry["spend"] == 0.0
assert log_entry["total_tokens"] == 0
assert log_entry["prompt_tokens"] == 0
@ -206,8 +209,7 @@ async def test_chat_completion_bad_model_with_spend_logs():
error_info = log_entry["metadata"]["error_information"]
assert "traceback" in error_info
assert error_info["error_code"] == "400"
assert error_info["error_class"] == "BadRequestError"
assert "litellm.BadRequestError" in error_info["error_message"]
assert error_info["error_class"] in ("ProxyModelNotFoundError", "BadRequestError")
assert "non-existent-model" in error_info["error_message"]
# Verify request details

View file

@ -188,6 +188,62 @@ def test_streaming_span_carries_time_to_first_chunk():
assert span.attributes[GenAI.RESPONSE_TIME_TO_FIRST_CHUNK] == pytest.approx(0.75)
def test_llm_call_span_reports_the_server_spans_route():
"""``litellm.request.route`` is the anchored server span's own ``http.route``,
so an operator can group LLM spans by endpoint without joining to the parent."""
logger, exporter = _logger()
root = logger.tracer.start_span("POST /engines/{model:path}/chat/completions")
root.set_attribute("http.route", "/engines/{model:path}/chat/completions")
set_request_root_span(root)
_emit_llm(logger, ambient=root)
root.end()
llm_span = next(s for s in exporter.get_finished_spans() if s.kind is SpanKind.CLIENT)
assert llm_span.attributes[LiteLLM.REQUEST_ROUTE] == "/engines/{model:path}/chat/completions"
def test_llm_call_span_omits_the_route_without_a_server_span():
"""An SDK call has no server span, so the key is absent rather than empty."""
logger, exporter = _logger()
_emit_llm(logger)
(span,) = exporter.get_finished_spans()
assert LiteLLM.REQUEST_ROUTE not in span.attributes
def test_failed_llm_call_span_reports_the_server_spans_route():
"""The failure leg builds the same span data, so an errored call is still
attributable to the endpoint it came in on."""
logger, exporter = _logger()
root = logger.tracer.start_span("POST /v1/responses/{response_id}")
root.set_attribute("http.route", "/v1/responses/{response_id}")
set_request_root_span(root)
_emit_llm(logger, ambient=root, fail=True)
root.end()
llm_span = next(s for s in exporter.get_finished_spans() if s.kind is SpanKind.CLIENT)
assert llm_span.attributes[LiteLLM.REQUEST_ROUTE] == "/v1/responses/{response_id}"
def test_deferred_llm_call_span_reports_the_server_spans_route():
"""``pre_call`` driven from a thread pool sees no recordable parent, so the span
is created in the close callback instead. That branch has to carry the route
too, and it can: the worker context still holds the anchor."""
logger, exporter = _logger()
root = logger.tracer.start_span("POST /v1/messages")
root.set_attribute("http.route", "/v1/messages")
set_request_root_span(root)
# no ``ambient``: pre_call runs with no recordable span active, which is what
# defers creation to the close callback
_emit_llm(logger)
root.end()
llm_span = next(s for s in exporter.get_finished_spans() if s.kind is SpanKind.CLIENT)
assert llm_span.attributes[LiteLLM.REQUEST_ROUTE] == "/v1/messages"
def test_non_streaming_span_has_no_time_to_first_chunk():
logger, exporter = _logger()
kwargs = {

View file

@ -5,6 +5,8 @@ surface and the server-span + shared-provider behavior it produces.
"""
from datetime import datetime, timezone
import pytest
@ -18,6 +20,7 @@ from opentelemetry.sdk.trace.export import SimpleSpanProcessor # noqa: E402
from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( # noqa: E402
InMemorySpanExporter,
)
from opentelemetry import trace # noqa: E402
from opentelemetry.trace import SpanKind # noqa: E402
from litellm.integrations.otel.model.config import ( # noqa: E402
@ -30,6 +33,23 @@ from litellm.integrations.otel.mount import ( # noqa: E402
_passthrough_span_name_hook,
instrument_fastapi_app,
)
from litellm.integrations.otel.plumbing.context import ( # noqa: E402
request_root_http_route,
set_request_root_span,
)
@pytest.fixture(autouse=True)
def _reset_request_root_span():
"""Clear the root-span anchor around every test. Production gets a fresh
contextvar copy per request task; the test process shares one context."""
from litellm.integrations.otel.plumbing import context as _otel_context
_otel_context._request_root_span.set(None)
_otel_context._mcp_message_transport_span.set(None)
yield
_otel_context._request_root_span.set(None)
_otel_context._mcp_message_transport_span.set(None)
@pytest.fixture(autouse=True)
@ -128,6 +148,85 @@ def test_passthrough_hook_ignores_non_recording_span():
assert span.name is None
def test_llm_span_route_is_read_off_the_server_span(monkeypatch):
"""``request_root_http_route`` answers with the SERVER span's own ``http.route``.
Driven through ``instrument_fastapi_app`` and the same
``create_litellm_proxy_request_started_span`` call the proxy makes per request,
so breaking either the mount or the anchor capture fails this."""
monkeypatch.setenv("LITELLM_OTEL_V2", "1")
is_otel_v2_enabled.cache_clear()
app = fastapi.FastAPI()
seen = {}
def _anchor_then_read(key):
logger.create_litellm_proxy_request_started_span(start_time=datetime.now(timezone.utc), headers=None)
seen[key] = request_root_http_route()
@app.post("/engines/{model:path}/chat/completions")
async def engines(model: str):
_anchor_then_read("templated")
return {}
@app.post("/openai/{endpoint:path}")
async def openai_passthrough(endpoint: str):
_anchor_then_read("passthrough")
return {}
logger = OpenTelemetryV2(config=OpenTelemetryV2Config(exporter="in_memory"))
exporter = InMemorySpanExporter()
logger._tracer_provider.add_span_processor(SimpleSpanProcessor(exporter))
# instrument_fastapi_app passes no provider, so it binds to the OTel global the
# way the proxy does once proxy_startup_event publishes one. set_tracer_provider
# is a once-per-process door, so place it directly and let monkeypatch undo it.
monkeypatch.setattr(trace, "_TRACER_PROVIDER", logger._tracer_provider)
instrument_fastapi_app(app)
client = TestClient(app)
client.post("/engines/gpt-4o-mini/chat/completions")
client.post("/openai/v1/responses/resp_abc123")
routes = {
(s.attributes or {})["http.route"] for s in exporter.get_finished_spans() if s.kind is SpanKind.SERVER
}
# a parameterized route keeps its template; the passthrough hook rewrote the
# catch-all to the literal path, and both spans have to follow their own span
assert routes == {"/engines/{model:path}/chat/completions", "/openai/v1/responses/resp_abc123"}
assert seen["templated"] == "/engines/{model:path}/chat/completions"
assert seen["passthrough"] == "/openai/v1/responses/resp_abc123"
def test_server_span_route_survives_the_span_ending():
"""The LLM span closes in an async callback that can run after the server span
has ended, so the attribute has to still be readable then."""
from opentelemetry.sdk.trace import TracerProvider
span = TracerProvider().get_tracer("t").start_span("POST /v1/responses/{response_id}")
span.set_attribute("http.route", "/v1/responses/{response_id}")
set_request_root_span(span)
span.end()
assert request_root_http_route() == "/v1/responses/{response_id}"
def test_no_server_span_means_no_route():
"""An SDK call has no anchored server span, so the attribute is omitted rather
than reported as empty."""
assert request_root_http_route() is None
def test_blank_route_on_the_server_span_is_omitted():
"""An excluded or unmatched path leaves the server span without a usable route.
Report nothing rather than a span attribute whose value is the empty string."""
from opentelemetry.sdk.trace import TracerProvider
span = TracerProvider().get_tracer("t").start_span("GET")
span.set_attribute("http.route", "")
set_request_root_span(span)
assert request_root_http_route() is None
def test_known_passthrough_prefixes_present():
"""Guard the prefix set against accidental edits."""
assert {"openai", "anthropic", "vertex_ai", "bedrock"} <= PASSTHROUGH_PREFIXES

View file

@ -722,6 +722,39 @@ def test_request_identity_falls_back_to_legacy_team_keys():
assert ident.team_alias == "legacy"
def test_llm_span_carries_proxy_request_route():
"""The LLM span records the proxy route the request arrived on, so it can be
filtered by endpoint (``/v1/responses`` vs ``/v1/chat/completions``) without
joining back to the root SERVER span's ``http.route``. The value is that
span's ``http.route`` verbatim, so a parameterized route reports the template
the SERVER span reports and not the path the caller happened to send."""
data: Final = LLMCallSpanData.from_standard_logging_payload(
_sample_payload(metadata={"user_api_key_request_route": "/v1/responses/resp_abc123"}),
request_route="/v1/responses/{response_id}",
)
attrs: Final = GenAIMapper().map(data)
assert attrs[LiteLLM.REQUEST_ROUTE] == "/v1/responses/{response_id}"
def test_llm_span_falls_back_to_the_logged_route_without_a_server_span():
"""The route the proxy recorded at auth is the backstop for a deployment whose
FastAPI instrumentation never mounted: there is no server span to disagree with
there, and an endpoint name is worth more than an absent attribute."""
data: Final = LLMCallSpanData.from_standard_logging_payload(
_sample_payload(metadata={"user_api_key_request_route": "/v1/responses"})
)
assert GenAIMapper().map(data)[LiteLLM.REQUEST_ROUTE] == "/v1/responses"
def test_llm_span_omits_request_route_off_the_proxy():
"""An SDK call has no inbound route, so the key is absent rather than empty."""
attrs: Final = GenAIMapper().map(LLMCallSpanData.from_standard_logging_payload(_sample_payload(metadata={})))
assert LiteLLM.REQUEST_ROUTE not in attrs
def test_guardrail_span_data_block_carries_verdict_and_error():
from litellm.integrations.otel.model.payloads import GuardrailSpanData

View file

@ -72,6 +72,7 @@ def _job_record(job: ActiveShadowEvalJob, target_type="key", target_id="key-hash
target_id=target_id,
router_name=job.router_name,
router_names=job.router_names,
models=sorted(job.models),
direction=job.direction,
baseline_model=job.baseline_model,
shadow_percentage=job.shadow_percentage,
@ -247,6 +248,7 @@ def _success_kwargs(
request_metadata=None,
call_type="acompletion",
model="claude-opus",
model_group="opus-group",
response_cost=None,
cache_hit=None,
):
@ -255,6 +257,7 @@ def _success_kwargs(
"id": request_id,
"call_type": call_type,
"model": model,
"model_group": model_group,
"metadata": {"user_api_key_hash": api_key_hash},
"model_parameters": {"temperature": 0.5, "stream": True},
"response_cost": response_cost,
@ -1049,6 +1052,79 @@ class TestTargetMatching:
assert logger._job_starts == {"key-job": 1, "team-job": 1}
@pytest.mark.asyncio
class TestModelScope:
"""A job scoped to model groups samples a target's request only when the group the
caller asked for is one of them; an out-of-scope request is not the job's traffic at
all, so it records no funnel event, exactly like a direction mismatch."""
@pytest.mark.parametrize(
"requested,sampled",
[("sonnet-group", True), ("opus-group", False), ("", False)],
ids=["in-scope-group-samples", "other-group-skips", "unknown-group-fails-closed"],
)
async def test_scope_admits_only_the_named_groups_and_counts_nothing_else(self, requested, sampled):
prisma = _prisma()
logger = _logger(router=_router(), prisma=prisma, jobs=(_job(models=frozenset({"sonnet-group", "haiku-group"})),))
await logger.async_log_success_event(_success_kwargs(model_group=requested), RESPONSE, None, None)
await _drain(logger)
assert prisma.db.litellm_shadowevalattempt.create.await_count == (1 if sampled else 0)
assert logger._test_funnel == []
async def test_an_unscoped_job_samples_every_group(self):
prisma = _prisma()
logger = _logger(router=_router(), prisma=prisma, jobs=(_job(),))
await logger.async_log_success_event(_success_kwargs(model_group="anything"), RESPONSE, None, None)
await _drain(logger)
prisma.db.litellm_shadowevalattempt.create.assert_awaited_once()
@pytest.mark.parametrize(
"scoped_to,requested",
[("sonnet-group", "fast"), ("fast", "sonnet-group")],
ids=["job-names-the-target-request-uses-the-alias", "job-names-the-alias-request-uses-the-target"],
)
async def test_an_alias_and_its_target_are_one_group_on_both_sides(self, scoped_to, requested):
"""Both the job's scope and the request's group resolve through the router's alias
map at match time, so re-pointing an alias follows config rather than freezing at
job start."""
router = _router()
router.model_group_alias = {"fast": "sonnet-group"}
prisma = _prisma(jobs=[_job_record(_job(models=frozenset({scoped_to})))])
logger = _logger(router=router, prisma=prisma)
await logger.async_log_success_event(_success_kwargs(model_group=requested), RESPONSE, None, None)
await _drain(logger)
prisma.db.litellm_shadowevalattempt.create.assert_awaited_once()
assert prisma.db.litellm_shadowevaljob.find_many.await_count == 1
async def test_a_repointed_alias_applies_to_the_next_request_without_a_cache_refill(self):
router = _router()
router.model_group_alias = {"fast": "sonnet-group"}
prisma = _prisma(jobs=[_job_record(_job(models=frozenset({"fast"})))])
logger = _logger(router=router, prisma=prisma)
await logger.async_log_success_event(_success_kwargs(model_group="sonnet-group"), RESPONSE, None, None)
await _drain(logger)
assert prisma.db.litellm_shadowevalattempt.create.await_count == 1
router.model_group_alias = {"fast": "haiku-group"}
await logger.async_log_success_event(
_success_kwargs(request_id="req-2", model_group="sonnet-group"), RESPONSE, None, None
)
await logger.async_log_success_event(
_success_kwargs(request_id="req-3", model_group="haiku-group"), RESPONSE, None, None
)
await _drain(logger)
rows = [call.kwargs["data"]["request_id"] for call in prisma.db.litellm_shadowevalattempt.create.call_args_list]
assert rows == ["req-1", "req-3"]
assert prisma.db.litellm_shadowevaljob.find_many.await_count == 1
@pytest.mark.asyncio
class TestActiveJobsCache:
async def test_cache_miss_reads_db_once_then_serves_from_cache(self):

View file

@ -5,6 +5,7 @@ from datetime import datetime, timedelta, timezone
import pytest
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from jwt.utils import base64url_decode, base64url_encode
from pydantic import SecretStr
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import (
@ -52,6 +53,12 @@ def _refresh_token() -> str:
return minted.token.get_secret_value()
def _corrupt_signature(token: str) -> str:
unsigned, signature = token.rsplit(".", 1)
raw = base64url_decode(signature)
return f"{unsigned}.{base64url_encode(bytes((raw[0] ^ 0x01,)) + raw[1:]).decode()}"
def test_kdf_is_deterministic_and_key_length_is_256_bit():
again = session_keys_from_master_key(MASTER_KEY)
assert again.signing_key.get_secret_value() == KEYS.signing_key.get_secret_value()
@ -109,8 +116,7 @@ def test_resolve_fails_expired_token_closed_and_flags_expiry():
def test_resolve_fails_tampered_token_closed_without_expiry_flag():
token = _access_token()
tampered = token[:-2] + ("aa" if not token.endswith("aa") else "bb")
result = resolve_session_bearer(f"Bearer {tampered}", KEYS, NOW)
result = resolve_session_bearer(f"Bearer {_corrupt_signature(token)}", KEYS, NOW)
assert isinstance(result, SessionBearerInvalid)
assert result.expired is False

View file

@ -6,6 +6,7 @@ import jwt
import pytest
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from jwt.utils import base64url_decode, base64url_encode
from pydantic import SecretStr, ValidationError
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import (
@ -66,6 +67,12 @@ def _mint_refresh() -> str:
return minted.token.get_secret_value()
def _corrupt_signature(token: str) -> str:
unsigned, signature = token.rsplit(".", 1)
raw = base64url_decode(signature)
return f"{unsigned}.{base64url_encode(bytes((raw[0] ^ 0x01,)) + raw[1:]).decode()}"
def _sign_claims(payload: dict, prefix: str = SESSION_TOKEN_PREFIX, keys: SessionKeys = KEYS) -> str:
return prefix + jwt.encode(payload, keys.signing_key.get_secret_value(), algorithm="HS256")
@ -138,8 +145,7 @@ def test_still_valid_one_second_before_expiry():
def test_tampered_signature_is_bad_signature():
token = _mint_access()
tampered = token[:-2] + ("aa" if not token.endswith("aa") else "bb")
assert isinstance(open_session_token(tampered, KEYS, NOW), SessionBadSignature)
assert isinstance(open_session_token(_corrupt_signature(token), KEYS, NOW), SessionBadSignature)
def test_key_rotation_invalidates_outstanding_tokens():
@ -329,8 +335,7 @@ def test_rs256_tampered_signature_is_bad_signature():
minted = mint_session_token(PRINCIPAL, RSA_KEYS, NOW)
assert isinstance(minted, MintedSessionToken)
token = minted.token.get_secret_value()
tampered = token[:-2] + ("aa" if not token.endswith("aa") else "bb")
assert isinstance(open_session_token(tampered, RSA_KEYS, NOW), SessionBadSignature)
assert isinstance(open_session_token(_corrupt_signature(token), RSA_KEYS, NOW), SessionBadSignature)
def test_rs256_expired_token_is_expired():
@ -413,8 +418,7 @@ def test_rotation_window_still_enforces_expiry_and_tamper_on_the_previous_key():
)
after = NOW + timedelta(seconds=SESSION_TTL_SECONDS + 1)
assert isinstance(open_session_token(token, rotated, after), SessionExpired)
tampered = token[:-2] + ("aa" if not token.endswith("aa") else "bb")
assert isinstance(open_session_token(tampered, rotated, NOW), SessionBadSignature)
assert isinstance(open_session_token(_corrupt_signature(token), rotated, NOW), SessionBadSignature)
def test_weak_or_garbage_private_key_pem_rejected_at_construction():

View file

@ -1,5 +1,5 @@
import asyncio
from typing import Any, Dict, List, Tuple
from typing import Any, Dict, List, Optional, Tuple
from unittest.mock import patch
import click
@ -7,7 +7,8 @@ import pytest
import yaml
from click.testing import CliRunner
from InquirerPy.base.control import Choice
from prompt_toolkit.application import create_app_session
from InquirerPy.prompts.fuzzy import InquirerPyFuzzyControl
from prompt_toolkit.application import AppSession, create_app_session
from prompt_toolkit.input import create_pipe_input
from prompt_toolkit.output import DummyOutput
@ -283,27 +284,45 @@ class TestRunConfigureWizardNotInteractive:
assert not config_path.exists()
def _highlighted_choice(session: AppSession) -> Optional[str]:
if session.app is None:
return None
controls = [c for c in session.app.layout.find_all_controls() if isinstance(c, InquirerPyFuzzyControl)]
if not controls or controls[0].choice_count == 0:
return None
return controls[0].selection["name"]
async def _wait_until_highlighted(session: AppSession, name: str) -> None:
async def _poll() -> None:
while _highlighted_choice(session) != name:
await asyncio.sleep(0.01)
await asyncio.wait_for(_poll(), timeout=5)
def _drive_fuzzy_pick(
models: Tuple[DiscoveredModel, ...],
prompt_label: str,
multiselect: bool,
key_events: List[Tuple[str, float]],
key_events: List[Tuple[str, Optional[str]]],
) -> List[str]:
"""Drives the real InquirerPy fuzzy prompt through prompt_toolkit's own test input/output,
exercising the actual widget (filtering, tab-to-toggle, enter-to-confirm) rather than mocking
it away. asyncio.to_thread propagates the create_app_session context into the worker thread
running _fuzzy_pick's synchronous .execute() call."""
running _fuzzy_pick's synchronous .execute() call. Each key event names the choice the widget
must highlight before the next key is sent (None sends the next key immediately)."""
async def _run() -> List[str]:
with create_pipe_input() as pipe_input:
with create_app_session(input=pipe_input, output=DummyOutput()):
with create_app_session(input=pipe_input, output=DummyOutput()) as session:
task = asyncio.ensure_future(
asyncio.to_thread(wizard_module._fuzzy_pick, models, prompt_label, multiselect)
)
await asyncio.sleep(0.05)
for text, delay in key_events:
for text, highlighted in key_events:
pipe_input.send_text(text)
await asyncio.sleep(delay)
if highlighted is not None:
await _wait_until_highlighted(session, highlighted)
return await task
return asyncio.run(_run())
@ -315,13 +334,13 @@ class TestFuzzyPickWidget:
def test_single_select_filters_and_returns_highlighted_match(self):
result = _drive_fuzzy_pick(
self._models(), "test", multiselect=False, key_events=[("model-13", 0.3), ("\r", 0.1)]
self._models(), "test", multiselect=False, key_events=[("model-13", "model-13"), ("\r", None)]
)
assert result == ["model-13"]
def test_multiselect_requires_tab_to_toggle_before_enter(self):
result = _drive_fuzzy_pick(
self._models(), "test", multiselect=True, key_events=[("model-7", 0.3), ("\t", 0.1), ("\r", 0.1)]
self._models(), "test", multiselect=True, key_events=[("model-7", "model-7"), ("\t", None), ("\r", None)]
)
assert result == ["model-7"]
@ -331,12 +350,12 @@ class TestFuzzyPickWidget:
"test",
multiselect=True,
key_events=[
("model-3", 0.3),
("\t", 0.1),
*[("\x7f", 0.02) for _ in range("model-3".__len__())],
("model-15", 0.3),
("\t", 0.1),
("\r", 0.1),
("model-3", "model-3"),
("\t", None),
("\x7f" * len("model-3"), None),
("model-15", "model-15"),
("\t", None),
("\r", None),
],
)
assert set(result) == {"model-3", "model-15"}

View file

@ -0,0 +1,231 @@
import json
import os
import time
import pytest
import requests
import responses
from click.testing import CliRunner
from litellm.proxy.client.cli import cli
from litellm.proxy.client.cli.commands import debug as debug_module
from litellm.proxy.client.cli.commands.debug import (
SLASH_COMMAND_NAME,
detect_claude_session_id,
install_slash_command,
)
SESSION = "e96634a3-fa28-4083-b354-55542e2dca01"
OK_ROW = {
"request_id": "req-ok",
"startTime": "2026-09-02T10:00:00",
"endTime": "2026-09-02T10:00:02",
"model": "claude-opus-4-1",
"custom_llm_provider": "anthropic",
"status": "success",
"spend": 0.0125,
"prompt_tokens": 100,
"completion_tokens": 20,
"metadata": {"status": "success"},
}
FAILED_ROW = {
"request_id": "req-failed",
"startTime": "2026-09-02T10:01:00",
"endTime": "2026-09-02T10:01:01",
"model": "claude-opus-4-1",
"custom_llm_provider": "anthropic",
"status": "failure",
"spend": 0.0,
"prompt_tokens": 0,
"completion_tokens": 0,
"metadata": json.dumps(
{
"status": "failure",
"error_information": {
"error_code": "400",
"error_class": "BadRequestError",
"error_message": "`prompt` is required when `stop` is not true.",
},
}
),
}
PROXY = "http://localhost:4000"
def _mock_proxy(rows, payloads):
responses.get(
f"{PROXY}/spend/logs/session/ui",
json={"data": rows, "total": len(rows), "page": 1, "page_size": 100, "total_pages": 1},
match=[responses.matchers.query_param_matcher({"session_id": SESSION}, strict_match=False)],
)
for request_id, payload in payloads.items():
responses.get(f"{PROXY}/spend/logs/ui/{request_id}", json=payload)
def _called_paths():
return [c.request.path_url.split("?")[0] for c in responses.calls]
@pytest.fixture(autouse=True)
def env(monkeypatch, tmp_path):
monkeypatch.setenv("LITELLM_PROXY_URL", PROXY)
monkeypatch.setenv("LITELLM_PROXY_API_KEY", "sk-test")
monkeypatch.setattr(debug_module, "REPORT_DIR", tmp_path / "reports")
monkeypatch.setattr(debug_module, "CLAUDE_DIR", tmp_path / "claude")
@responses.activate
def test_report_includes_spend_error_and_bodies_for_failed_turn(tmp_path):
payloads = {
"req-failed": {
"proxy_server_request": {"body": {"model": "claude-opus-4-1", "messages": [{"role": "user"}]}},
"response": {"error": {"message": "`prompt` is required"}},
},
"req-ok": {"proxy_server_request": {"body": {"model": "claude-opus-4-1"}}, "response": {"id": "msg_1"}},
}
_mock_proxy([FAILED_ROW, OK_ROW], payloads)
result = CliRunner().invoke(cli, ["debug", "claude", "--session-id", SESSION, "--recent-bodies", "0"])
assert result.exit_code == 0, result.output
assert "turns: 2, failed: 1" in result.output
assert "total spend: $0.012500" in result.output
assert "### 1. ok claude-opus-4-1" in result.output
assert "### 2. FAILED claude-opus-4-1" in result.output
assert "`400` BadRequestError" in result.output
assert "`prompt` is required when `stop` is not true." in result.output
assert '"messages"' in result.output
assert "msg_1" not in result.output
assert _called_paths() == ["/spend/logs/session/ui", "/spend/logs/ui/req-failed"]
saved = tmp_path / "reports" / f"claude-{SESSION}.md"
assert result.stdout.startswith(saved.read_text())
assert "### 2. FAILED" in saved.read_text()
@responses.activate
def test_recent_bodies_fetches_latest_turns_even_when_successful():
_mock_proxy([OK_ROW], {"req-ok": {"proxy_server_request": {"body": {"x": 1}}, "response": {"id": "msg_1"}}})
result = CliRunner().invoke(cli, ["debug", "claude", "--session-id", SESSION, "--no-save"])
assert result.exit_code == 0, result.output
assert "msg_1" in result.output
assert _called_paths() == ["/spend/logs/session/ui", "/spend/logs/ui/req-ok"]
@responses.activate
def test_bodies_are_truncated_to_max_chars():
_mock_proxy([OK_ROW], {"req-ok": {"proxy_server_request": {"body": "a" * 5000}, "response": None}})
result = CliRunner().invoke(
cli, ["debug", "claude", "--session-id", SESSION, "--no-save", "--max-body-chars", "200"]
)
assert result.exit_code == 0, result.output
assert "truncated" in result.output
assert "a" * 300 not in result.output
@responses.activate
def test_no_rows_is_a_clear_error():
_mock_proxy([], {})
result = CliRunner().invoke(cli, ["debug", "claude", "--session-id", SESSION])
assert result.exit_code != 0
assert "No spend logs found for session" in result.output
def test_no_session_id_anywhere_is_a_clear_error(monkeypatch):
monkeypatch.delenv("CLAUDE_CODE_SESSION_ID", raising=False)
result = CliRunner().invoke(cli, ["debug", "claude"])
assert result.exit_code != 0
assert "Could not find a Claude Code session" in result.output
OLD_SESSION = "0f3c2b1a-1111-4222-8333-444455556666"
NEW_SESSION = "2d79c54d-4644-4708-b03e-95395ef9ecbd"
def test_detect_session_id_prefers_env_then_newest_session_transcript(tmp_path):
project = tmp_path / "projects" / "-Users-me-repo"
project.mkdir(parents=True)
old = project / f"{OLD_SESSION}.jsonl"
new = project / f"{NEW_SESSION}.jsonl"
subagent = project / "agent-a1b2c3d4.jsonl"
old.write_text("{}")
new.write_text("{}")
subagent.write_text("{}")
now = time.time()
os.utime(old, (now - 100, now - 100))
os.utime(new, (now - 50, now - 50))
os.utime(subagent, (now, now))
assert detect_claude_session_id({}, tmp_path) == NEW_SESSION
assert detect_claude_session_id({"CLAUDE_CODE_SESSION_ID": "from-env"}, tmp_path) == "from-env"
assert detect_claude_session_id({"CLAUDE_SESSION_ID": "stale-name"}, tmp_path) == NEW_SESSION
assert detect_claude_session_id({}, tmp_path / "missing") is None
def test_install_slash_command_writes_runnable_command_file(tmp_path):
path = install_slash_command(tmp_path)
assert path == tmp_path / "commands" / f"{SLASH_COMMAND_NAME}.md"
body = path.read_text()
assert body.startswith("---\n")
assert "allowed-tools: Bash(lite debug claude:*)" in body
assert "!`lite debug claude $ARGUMENTS`" in body
result = CliRunner().invoke(cli, ["debug", "install-claude-command"])
assert result.exit_code == 0, result.output
assert "/debug-lite" in result.output
@responses.activate
def test_rejected_key_is_a_clear_error_not_a_traceback():
responses.get(
f"{PROXY}/spend/logs/session/ui",
status=401,
json={"error": {"message": "Authentication Error, Invalid proxy server token passed", "code": "401"}},
)
result = CliRunner().invoke(cli, ["debug", "claude", "--session-id", SESSION])
assert isinstance(result.exception, SystemExit), result.exception
assert result.exit_code == 1
assert "401" in result.output
assert "Invalid proxy server token passed" in result.output
@responses.activate
def test_unreachable_proxy_is_a_clear_error_not_a_traceback():
responses.get(f"{PROXY}/spend/logs/session/ui", body=requests.ConnectionError("Connection refused"))
result = CliRunner().invoke(cli, ["debug", "claude", "--session-id", SESSION])
assert isinstance(result.exception, SystemExit), result.exception
assert result.exit_code == 1
assert "Connection refused" in result.output
@responses.activate
def test_non_json_proxy_response_is_a_clear_error_not_a_traceback():
responses.get(f"{PROXY}/spend/logs/session/ui", body="<html>502 Bad Gateway</html>")
result = CliRunner().invoke(cli, ["debug", "claude", "--session-id", SESSION])
assert isinstance(result.exception, SystemExit), result.exception
assert result.exit_code == 1
assert "/spend/logs/session/ui failed" in result.output
@responses.activate
def test_logged_content_with_code_fences_stays_inside_its_fence():
fenced_error_row = {
**FAILED_ROW,
"metadata": {
"status": "failure",
"error_information": {"error_code": "400", "error_message": "bad\n```\nrequest"},
},
}
_mock_proxy([fenced_error_row], {"req-failed": {"proxy_server_request": None, "response": "x\n````\ny"}})
result = CliRunner().invoke(cli, ["debug", "claude", "--session-id", SESSION, "--no-save"])
assert result.exit_code == 0, result.output
assert "````\nbad\n```\nrequest\n````\n" in result.output
assert "`````json\nx\n````\ny\n`````\n" in result.output

View file

@ -1,6 +1,7 @@
import asyncio
import json
import time
from typing import Final
from datetime import datetime, timedelta
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
@ -1190,27 +1191,23 @@ def test_health_liveliness_endpoint(proxy_client):
Test that /health/liveliness endpoint returns 200 OK with "I'm alive!" message.
This is a critical orchestration endpoint that must be simple and fast.
"""
# Measure the time taken for the health check call
start_time = time.perf_counter()
warm_up: Final = proxy_client.get("/health/liveliness")
assert warm_up.status_code == 200, f"Expected 200 OK, got {warm_up.status_code}: {warm_up.text}"
# Make GET request to /health/liveliness
response = proxy_client.get("/health/liveliness")
def _timed_poll() -> tuple[float, httpx.Response]:
start_time: Final = time.perf_counter()
response: Final = proxy_client.get("/health/liveliness")
return (time.perf_counter() - start_time) * 1000, response
end_time = time.perf_counter()
duration_ms = (end_time - start_time) * 1000
polls: Final = tuple(_timed_poll() for _ in range(5))
# Assert response status
assert response.status_code == 200, f"Expected 200 OK, got {response.status_code}: {response.text}"
for _, response in polls:
assert response.status_code == 200, f"Expected 200 OK, got {response.status_code}: {response.text}"
assert response.json() == "I'm alive!", f"Expected 'I'm alive!' message, got: {response.json()}"
# Assert response content (FastAPI JSON-encodes the string)
assert response.json() == "I'm alive!", f"Expected 'I'm alive!' message, got: {response.json()}"
# Verify response is fast (should be < 100ms for a simple endpoint)
# This is critical for orchestration systems that poll frequently
assert duration_ms < 100, f"Health check took {duration_ms:.2f}ms, expected < 100ms for a simple endpoint"
# Log the duration for visibility (useful for CI/CD monitoring)
print(f"\n/health/liveliness response time: {duration_ms:.2f}ms")
durations_ms: Final = tuple(sorted(duration_ms for duration_ms, _ in polls))
median_ms: Final = durations_ms[len(durations_ms) // 2]
assert median_ms < 100, f"Median of {len(polls)} health checks took {median_ms:.2f}ms, expected < 100ms"
def test_health_liveness_endpoint(proxy_client):

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

@ -882,6 +882,7 @@ def _leg_record(**overrides: object) -> MagicMock:
"target_id": "key-hash",
"router_name": "my-router",
"router_names": (),
"models": (),
"direction": "forward",
"baseline_model": None,
"judge_model": "anthropic/claude-sonnet-5",
@ -1033,6 +1034,7 @@ def _shadow_prisma(
"target_id",
"router_name",
"router_names",
"models",
"direction",
"baseline_model",
"judge_model",
@ -1533,6 +1535,87 @@ async def test_start_shadow_eval_forward_leaves_the_baseline_column_empty(monkey
assert rows[0]["baseline_model"] is None
@pytest.mark.asyncio
async def test_start_shadow_eval_writes_the_model_scope_on_every_leg_and_echoes_it(monkeypatch: pytest.MonkeyPatch):
"""A model scope is job config, so every leg carries the same copy and both the start
response and a later list read report it; an auto-router is a legitimate scope (a
forward job on one router may sample what another router serves today)."""
import litellm.proxy.proxy_server as proxy_server
_configure_anthropic_sdk_judge(monkeypatch)
prisma = _shadow_prisma()
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
monkeypatch.setattr(proxy_server, "llm_router", _shadow_router())
response = await start_shadow_eval(
_start_request(api_key_ids=("key-hash", "key-hash-2"), models=("cheap", "sonnet-router")), ADMIN
)
rows = prisma.db.litellm_shadowevaljob.create_many.call_args.kwargs["data"]
assert [row["models"] for row in rows] == [["cheap", "sonnet-router"], ["cheap", "sonnet-router"]]
assert response.models == ("cheap", "sonnet-router")
listed = _shadow_prisma(legs=[_leg_record(models=("cheap",)), _leg_record(id="leg-0", group_id="job-0")])
monkeypatch.setattr(proxy_server, "prisma_client", listed)
jobs = await list_shadow_eval_jobs(VIEWER, target_type=None, target_id=None, limit=50)
assert {job.job_id: job.models for job in jobs} == {"job-1": ("cheap",), "job-0": ()}
@pytest.mark.asyncio
async def test_start_shadow_eval_accepts_a_team_public_scope_for_a_user_target(monkeypatch: pytest.MonkeyPatch):
"""A user's traffic can arrive on any team's key, so a name only one team can ask for
is a legitimate scope for a user target even though it resolves for nobody unscoped."""
import litellm.proxy.proxy_server as proxy_server
_configure_anthropic_sdk_judge(monkeypatch)
prisma = _shadow_prisma(known_users={"dev-alice": "alice@example.com"})
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
monkeypatch.setattr(proxy_server, "llm_router", _shadow_router())
response = await start_shadow_eval(
_start_request(api_key_ids=(), user_ids=("dev-alice",), models=("house-judge",)), ADMIN
)
assert response.models == ("house-judge",)
assert prisma.db.litellm_shadowevaljob.create_many.call_args.kwargs["data"][0]["models"] == ["house-judge"]
@pytest.mark.asyncio
async def test_start_shadow_eval_rejects_a_model_scope_this_proxy_does_not_serve(monkeypatch: pytest.MonkeyPatch):
"""A typo'd model name would otherwise start a job that samples nothing. Only the
unresolvable names are reported, so the caller fixes them in one round."""
import litellm.proxy.proxy_server as proxy_server
_configure_anthropic_sdk_judge(monkeypatch)
prisma = _shadow_prisma()
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
monkeypatch.setattr(proxy_server, "llm_router", _shadow_router())
with pytest.raises(HTTPException) as exc:
await start_shadow_eval(_start_request(models=("cheap", "no-such-model-zzz")), ADMIN)
assert exc.value.status_code == 400
assert "'no-such-model-zzz'" in exc.value.detail
assert "'cheap'" not in exc.value.detail
prisma.db.litellm_shadowevaljob.create_many.assert_not_called()
def test_start_request_dedupes_the_model_scope_and_rejects_blank_names():
assert _start_request(models=("cheap", "mid", "cheap")).models == ("cheap", "mid")
assert _start_request().models == ()
with pytest.raises(ValidationError, match="non-empty model group names"):
_start_request(models=("cheap", " "))
def test_start_request_rejects_a_model_scope_on_a_reverse_job():
"""Reverse admission is the router's own traffic, whose requested group is always the
router, so a plain-model scope would sample nothing and the router itself is a no-op."""
with pytest.raises(ValidationError, match="only meaningful for a forward job"):
_start_request(direction="reverse", baseline_model="cheap", models=("mid",))
with pytest.raises(ValidationError, match="only meaningful for a forward job"):
_start_request(direction="reverse", baseline_model="cheap", models=("my-router",))
assert _start_request(direction="reverse", baseline_model="cheap").models == ()
@pytest.mark.asyncio
async def test_start_shadow_eval_rejects_keys_this_proxy_does_not_know(monkeypatch: pytest.MonkeyPatch):
"""A typo'd api_key_id would otherwise create a leg no traffic can ever match. Every

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

@ -12429,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

@ -2,15 +2,27 @@
# This tests litellm router
import pytest
import logging
from typing import Final
import pytest
import litellm
from litellm._logging import verbose_logger
async def _routed_model_ids(
router: litellm.Router, tags: list[str], remaining: frozenset[str], attempts: int = 100
) -> frozenset[str]:
if not remaining or attempts == 0:
return frozenset()
response: Final = await router.acompletion(
model="gpt-4", messages=[{"role": "user", "content": "hi"}], metadata={"tags": tags}, mock_response="hi"
)
seen: Final = frozenset({response._hidden_params["model_id"]})
return seen | await _routed_model_ids(router, tags, remaining - seen, attempts - 1)
@pytest.mark.asyncio()
async def test_router_free_paid_tier():
"""
@ -850,17 +862,10 @@ async def test_negation_regex_pattern_treated_as_literal():
# The regex-like string matches no deployment tag literally, so all
# candidates survive and both model IDs are reachable.
seen_ids = set()
for _ in range(10):
response = await router.acompletion(
model="gpt-4",
messages=[{"role": "user", "content": "hi"}],
metadata={"tags": ["!provider:(anthropic|openai)"]},
mock_response="hi",
)
seen_ids.add(response._hidden_params["model_id"])
expected: Final = frozenset({"anthropic-model", "openai-model"})
routed_ids: Final = await _routed_model_ids(router, ["!provider:(anthropic|openai)"], expected)
assert seen_ids == {"anthropic-model", "openai-model"}
assert routed_ids == expected
@pytest.mark.asyncio()
@ -1281,17 +1286,10 @@ async def test_chain_enable_tag_filtering_false_overrides_router_level_true():
enable_tag_filtering=True,
)
seen_ids = set()
for _ in range(10):
response = await router.acompletion(
model="gpt-4",
messages=[{"role": "user", "content": "hi"}],
metadata={"tags": ["teamA"]},
mock_response="hi",
)
seen_ids.add(response._hidden_params["model_id"])
expected: Final = frozenset({"team-a-deployment", "team-b-deployment"})
routed_ids: Final = await _routed_model_ids(router, ["teamA"], expected)
assert seen_ids == {"team-a-deployment", "team-b-deployment"}
assert routed_ids == expected
@pytest.mark.asyncio()

View file

@ -53,6 +53,12 @@ case "$*" in
"eslint --no-warn-ignored"*)
[ "${STUB_FAIL:-}" = "eslint" ] && exit 1
;;
"eslint . -f json"*)
if [ -n "${STUB_HANG_DIR:-}" ]; then
touch "$STUB_HANG_DIR/eslint_report.started"
sleep 60
fi
;;
esac
exit 0
"""
@ -340,11 +346,12 @@ def test_interrupt_kills_background_jobs_and_removes_logs(tmp_path: Path) -> Non
)
try:
assert _wait_until((hang_dir / "make.started").exists, 10)
assert _wait_until((hang_dir / "eslint_report.started").exists, 10)
os.killpg(proc.pid, signal.SIGINT)
assert proc.wait(timeout=10) != 0
make_pid = int((hang_dir / "make.pid").read_text())
assert _wait_until(lambda: _pid_gone(make_pid), 5)
assert list(tmp_dir.iterdir()) == []
assert _wait_until(lambda: not any(tmp_dir.iterdir()), 5), list(tmp_dir.iterdir())
finally:
with suppress(ProcessLookupError, PermissionError):
os.killpg(proc.pid, signal.SIGTERM)

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

@ -104,6 +104,7 @@ const job = (overrides: Partial<ShadowEvalJob> = {}): ShadowEvalJob => ({
status: "running",
router_name: "claude-auto",
router_names: ["claude-auto"],
models: [],
direction: "forward",
baseline_model: null,
judge_model: "anthropic/claude-sonnet-5",
@ -450,6 +451,7 @@ describe("ShadowEvalSection", () => {
api_key_ids: ["hash-alpha", "hash-beta"],
team_ids: [],
user_ids: [],
models: [],
router_names: ["gpt-auto"],
direction: "forward",
shadow_percentage: 10,
@ -479,6 +481,7 @@ describe("ShadowEvalSection", () => {
api_key_ids: [],
team_ids: ["team-eng"],
user_ids: [],
models: [],
router_names: ["gpt-auto"],
direction: "forward",
shadow_percentage: 10,
@ -489,15 +492,41 @@ describe("ShadowEvalSection", () => {
expect(start.mutate).toHaveBeenCalledWith(expectedBody);
});
it("narrows a job to the picked model groups and shows the scope on the job headline", async () => {
const user = userEvent.setup();
const { start } = mockHooks({});
render(<ShadowEvalSection />);
await user.click(screen.getByPlaceholderText("Search teams by alias"));
const teamList = await screen.findByTestId("paginated-multi-select-list");
await user.click(within(teamList).getByText("engineering"));
await chooseSelectOption(user, screen.getByPlaceholderText("Every model the targets use"), "prod-claude");
await chooseSelectOption(user, screen.getByPlaceholderText("Select up to 4 auto-routers"), "gpt-auto");
await user.click(screen.getByPlaceholderText("Select a judge model"));
await user.click(await screen.findByRole("option", { name: /anthropic\/claude-sonnet-5/ }));
await user.click(screen.getByText("Start shadow eval"));
expect(start.mutate).toHaveBeenCalledWith(
expect.objectContaining({ team_ids: ["team-eng"], models: ["prod-claude"] }),
);
const scoped = job({ models: ["prod-claude", "prod-haiku"] });
mockHooks({ jobs: [scoped], detailsById: { "job-1": scoped } });
render(<ShadowEvalSection />);
expect(screen.getByText("prod-claude, prod-haiku")).toBeInTheDocument();
});
it("requires a baseline model in reverse mode and submits it, while forward mode never shows the picker", async () => {
const user = userEvent.setup();
const { start } = mockHooks({});
render(<ShadowEvalSection />);
expect(screen.queryByPlaceholderText("Select a baseline model")).not.toBeInTheDocument();
expect(screen.getByPlaceholderText("Every model the targets use")).toBeInTheDocument();
await user.click(screen.getByText("Adoption check: key's traffic vs the router"));
await user.click(await screen.findByText("Regression check: router's picks vs a baseline"));
expect(screen.queryByPlaceholderText("Every model the targets use")).not.toBeInTheDocument();
await user.click(screen.getByPlaceholderText("Search keys by alias"));
const keyList = await screen.findByTestId("paginated-multi-select-list");
await user.click(within(keyList).getByText("prod-alpha"));
@ -516,6 +545,7 @@ describe("ShadowEvalSection", () => {
api_key_ids: ["hash-alpha"],
team_ids: [],
user_ids: [],
models: [],
router_names: ["gpt-auto"],
direction: "reverse",
baseline_model: "prod-claude",
@ -551,6 +581,7 @@ describe("ShadowEvalSection", () => {
api_key_ids: ["hash-alpha"],
team_ids: [],
user_ids: [],
models: [],
router_names: ["gpt-auto", "claude-auto"],
direction: "forward",
shadow_percentage: 10,

View file

@ -87,17 +87,25 @@ const targetStatus = (job: ShadowEvalJob, target: ShadowEvalJobTarget): string =
const jobRouters = (job: ShadowEvalJob): string => (job.router_names ?? [job.router_name]).join(", ");
const jobModelScope = (job: ShadowEvalJob): React.ReactNode =>
job.models && job.models.length > 0 ? (
<>
{" "}
on <span className="font-mono text-xs">{job.models.join(", ")}</span>
</>
) : null;
const jobHeadline = (job: ShadowEvalJob): React.ReactNode =>
job.direction === "reverse" ? (
<>
Comparing <span className="font-mono text-xs">{jobRouters(job)}</span> to{" "}
<span className="font-mono text-xs">{job.baseline_model}</span> on {job.shadow_percentage}% of{" "}
<span className="font-mono text-xs">{shadowedTargetsLabel(job)}</span> traffic
<span className="font-mono text-xs">{shadowedTargetsLabel(job)}</span> traffic{jobModelScope(job)}
</>
) : (
<>
Shadowing {job.shadow_percentage}% of <span className="font-mono text-xs">{shadowedTargetsLabel(job)}</span>{" "}
traffic via <span className="font-mono text-xs">{jobRouters(job)}</span>
traffic{jobModelScope(job)} via <span className="font-mono text-xs">{jobRouters(job)}</span>
</>
);

View file

@ -23,6 +23,7 @@ import { useStartShadowEval, type ShadowEvalJob } from "./useShadowEval";
type ShadowEvalDirection = ShadowEvalJob["direction"];
const MAX_ROUTERS = 4;
const MAX_MODELS = 100;
const RECOMMENDED_JUDGE_MODELS = ["anthropic/claude-sonnet-5", "openai/gpt-4o", "gemini/gemini-2.5-pro"] as const;
@ -206,6 +207,7 @@ interface StartFormValidityInputs {
apiKeyIds: string[];
teamIds: string[];
userIds: string[];
models: string[];
routerNames: string[];
direction: ShadowEvalDirection;
baselineModel: string;
@ -224,7 +226,8 @@ const startFormValidity = (inputs: StartFormValidityInputs) => {
const routerCountValid = inputs.routerNames.length >= 1 && inputs.routerNames.length <= MAX_ROUTERS;
const routersMatchDirection = inputs.direction === "forward" || inputs.routerNames.length === 1;
const routersValid = routerCountValid && routersMatchDirection;
const modelsPicked = routersValid && inputs.judgeModel !== "" && baselinePicked;
const scopeValid = routersValid && (inputs.direction === "reverse" || inputs.models.length <= MAX_MODELS);
const modelsPicked = scopeValid && inputs.judgeModel !== "" && baselinePicked;
const filled = targetsPicked && modelsPicked;
const boundsValid = percentageValid && maxBudgetValid;
const valid = Boolean(inputs.accessToken) && filled && boundsValid;
@ -235,6 +238,7 @@ interface StartBodyInputs {
apiKeyIds: string[];
teamIds: string[];
userIds: string[];
models: string[];
routerNames: string[];
direction: ShadowEvalDirection;
baselineModel: string;
@ -248,6 +252,7 @@ const buildStartBody = (inputs: StartBodyInputs) => ({
api_key_ids: inputs.apiKeyIds,
team_ids: inputs.teamIds,
user_ids: inputs.userIds,
models: inputs.direction === "forward" ? inputs.models : [],
router_names: inputs.routerNames,
direction: inputs.direction,
...(inputs.direction === "reverse" ? { baseline_model: inputs.baselineModel } : {}),
@ -262,6 +267,7 @@ export const StartForm: React.FC = () => {
const [apiKeyIds, setApiKeyIds] = useState<string[]>([]);
const [teamIds, setTeamIds] = useState<string[]>([]);
const [userIds, setUserIds] = useState<string[]>([]);
const [models, setModels] = useState<string[]>([]);
const [routerNames, setRouterNames] = useState<string[]>([]);
const [direction, setDirection] = useState<ShadowEvalDirection>("forward");
const [baselineModel, setBaselineModel] = useState("");
@ -272,6 +278,11 @@ export const StartForm: React.FC = () => {
const { data: autoRouters } = useAutoRouters();
const judgeModelOptions = useJudgeModelOptions();
const baselineModelOptions = useBaselineModelOptions();
const configuredGroups = usePlainModelGroups();
const modelOptions = useMemo<SearchSelectOption[]>(
() => [...configuredGroups].toSorted((a, b) => a.localeCompare(b)).map((name) => ({ label: name, value: name })),
[configuredGroups],
);
const start = useStartShadowEval();
const routerOptions = useMemo<SearchSelectOption[]>(() => {
@ -286,6 +297,7 @@ export const StartForm: React.FC = () => {
apiKeyIds,
teamIds,
userIds,
models,
routerNames,
direction,
baselineModel,
@ -299,6 +311,7 @@ export const StartForm: React.FC = () => {
apiKeyIds,
teamIds,
userIds,
models,
routerNames,
direction,
baselineModel,
@ -344,6 +357,22 @@ export const StartForm: React.FC = () => {
<Field label="Users to shadow" htmlFor="shadow-eval-user">
<UserSelect value={userIds} onChange={setUserIds} />
</Field>
{direction === "forward" && (
<Field label="Only on models">
<MultiSelect
options={modelOptions}
value={models}
onValueChange={setModels}
placeholder="Every model the targets use"
emptyText="No models configured"
/>
{models.length > MAX_MODELS ? (
<p className="text-xs text-destructive">Pick at most {MAX_MODELS} models</p>
) : (
<p className="text-xs text-muted-foreground">Narrows every target above to requests for these models</p>
)}
</Field>
)}
<RouterField
options={routerOptions}
routerNames={routerNames}

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

@ -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", () => {

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 {
@ -324,7 +330,7 @@ export const getSemanticConfigError = ({
};
interface CustomTierWireFieldInputs {
classifierLlmConfig: ClassifierLLMConfig | undefined;
classifierLlmConfig: ClassifierLLMConfigWire | undefined;
planModeMinTierId: string | undefined;
classificationPrompt: string | undefined;
classificationExamples: string | undefined;
@ -356,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,

View file

@ -1264,7 +1264,10 @@ export interface paths {
* A target is a virtual key, a team, or a user. Team and user targets match on the
* identity every request resolves to at auth time, so they cover JWT-authenticated
* traffic, which presents no virtual key; a user target samples that user's traffic
* across all their teams, whether it arrives on a JWT or a key they own.
* across all their teams, whether it arrives on a JWT or a key they own. models narrows
* every target to requests for those model groups, so a user plus one model samples that
* user's traffic on that model across every key they own; it is forward-only, since a
* reverse job already samples exactly the traffic its own router served.
*
* A forward job answers whether the targets should adopt router_name: it samples the
* requests the router did not serve and duplicates them through it. A reverse job
@ -22764,22 +22767,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 +25326,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
@ -35763,6 +35808,12 @@ export interface components {
* @description Most recent attempt error; detail endpoint only
*/
last_error?: string | null;
/**
* Models
* @description Model groups the sampled traffic is narrowed to; empty means every model the targets use
* @default []
*/
models: string[];
/** @description Stratified verdicts; detail endpoint only */
results?: components["schemas"]["ShadowEvalResult"] | null;
/**
@ -36186,6 +36237,12 @@ export interface components {
* @default 10
*/
max_budget: number;
/**
* Models
* @description Model groups to narrow the sampled traffic to, matched on the group the caller requested and resolved through model_group_alias, so an alias and its target are one name. Empty samples every model the targets use. This ANDs with the targets: a job over a user and one model samples that user's requests on that model across every key they own, and none of their other traffic. Forward jobs only: a reverse job samples exactly the traffic its own router served, which no other model group can name
* @default []
*/
models: string[];
/**
* Router Name
* @description The auto-router under evaluation, in either direction: the single-router spelling of router_names. Provide exactly one of the two fields