mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_shadow_eval_judge_output_cap
This commit is contained in:
commit
955baf8a5c
54 changed files with 2637 additions and 332 deletions
|
|
@ -0,0 +1 @@
|
|||
ALTER TABLE "LiteLLM_ShadowEvalJob" ADD COLUMN IF NOT EXISTS "models" TEXT[] NOT NULL DEFAULT ARRAY[]::TEXT[];
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]] = {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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``,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
],
|
||||
|
|
|
|||
372
litellm/proxy/client/cli/commands/debug.py
Normal file
372
litellm/proxy/client/cli/commands/debug.py
Normal 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.")
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
61
litellm/proxy/management_helpers/resource_display_names.py
Normal file
61
litellm/proxy/management_helpers/resource_display_names.py
Normal 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})
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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=(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
231
tests/test_litellm/proxy/client/cli/test_debug_commands.py
Normal file
231
tests/test_litellm/proxy/client/cli/test_debug_commands.py
Normal 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
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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})
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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", () => {
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
</>
|
||||
);
|
||||
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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) ────────────────────────────
|
||||
|
||||
|
|
|
|||
|
|
@ -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", () => {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
59
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
59
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue