mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_fix_realtime_reasoning_double_bill
This commit is contained in:
commit
6294367920
51 changed files with 2941 additions and 346 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ from openai.types.responses import ResponseReasoningItem
|
|||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.llms.azure.chat.gpt_5_transformation import AzureOpenAIGPT5Config
|
||||
from litellm.llms.azure.common_utils import BaseAzureLLM
|
||||
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
|
||||
from litellm.types.llms.openai import *
|
||||
|
|
@ -29,6 +30,14 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
def custom_llm_provider(self) -> LlmProviders:
|
||||
return LlmProviders.AZURE
|
||||
|
||||
@staticmethod
|
||||
def _supports_reasoning_effort_none(model: str) -> bool:
|
||||
return AzureOpenAIGPT5Config._supports_reasoning_effort_level(model, "none")
|
||||
|
||||
@staticmethod
|
||||
def _effort_resolves_to_none(model: str, effort: str | None) -> bool:
|
||||
return AzureOpenAIGPT5Config.effort_resolves_to_none(model, effort)
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list:
|
||||
"""
|
||||
Azure Responses API does not support context_management (compaction).
|
||||
|
|
|
|||
|
|
@ -7157,6 +7157,53 @@
|
|||
"supports_xhigh_reasoning_effort": true,
|
||||
"supports_minimal_reasoning_effort": false
|
||||
},
|
||||
"azure/gpt-6-astra": {
|
||||
"cache_creation_input_token_cost": 1.25e-05,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-05,
|
||||
"cache_read_input_token_cost": 1e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 2e-06,
|
||||
"input_cost_per_token": 1e-05,
|
||||
"input_cost_per_token_above_272k_tokens": 2e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 922000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 7.5e-05,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_high": 0.01,
|
||||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_max_reasoning_effort": true,
|
||||
"supports_minimal_reasoning_effort": false,
|
||||
"supports_native_streaming": true,
|
||||
"supports_none_reasoning_effort": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_xhigh_reasoning_effort": true
|
||||
},
|
||||
"azure/us/gpt-5.6": {
|
||||
"cache_creation_input_token_cost": 6.875e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1.375e-05,
|
||||
|
|
@ -7376,6 +7423,53 @@
|
|||
"supports_xhigh_reasoning_effort": true,
|
||||
"supports_minimal_reasoning_effort": false
|
||||
},
|
||||
"azure/us/gpt-6-astra": {
|
||||
"cache_creation_input_token_cost": 1.375e-05,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 2.75e-05,
|
||||
"cache_read_input_token_cost": 1.1e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 2.2e-06,
|
||||
"input_cost_per_token": 1.1e-05,
|
||||
"input_cost_per_token_above_272k_tokens": 2.2e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 922000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5.5e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 8.25e-05,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_high": 0.01,
|
||||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_max_reasoning_effort": true,
|
||||
"supports_minimal_reasoning_effort": false,
|
||||
"supports_native_streaming": true,
|
||||
"supports_none_reasoning_effort": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_xhigh_reasoning_effort": true
|
||||
},
|
||||
"azure/eu/gpt-5.6": {
|
||||
"cache_creation_input_token_cost": 6.875e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1.375e-05,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -3326,9 +3326,6 @@ async def team_member_delete(
|
|||
data=data,
|
||||
)
|
||||
|
||||
if not removed_team_members:
|
||||
raise HTTPException(status_code=400, detail={"error": "User not found in team"})
|
||||
|
||||
existing_team_row.members_with_roles = new_team_members
|
||||
|
||||
_db_new_team_members: Final[list[dict]] = [m.model_dump() for m in new_team_members]
|
||||
|
|
@ -3336,17 +3333,27 @@ async def team_member_delete(
|
|||
## DELETE TEAM ID from USER ROW, IF EXISTS ##
|
||||
# get user row
|
||||
removed_user_ids: Final = frozenset(m.user_id for m in removed_team_members if m.user_id is not None)
|
||||
addressed_user_ids: Final = (
|
||||
removed_user_ids if removed_team_members else frozenset((data.user_id,) if data.user_id is not None else ())
|
||||
)
|
||||
key_val: Final[Mapping[str, object]] = (
|
||||
{"user_id": {"in": sorted(removed_user_ids)}} if removed_user_ids else {"user_email": data.user_email}
|
||||
{"user_id": {"in": sorted(addressed_user_ids)}} if addressed_user_ids else {"user_email": data.user_email}
|
||||
)
|
||||
member_tx: Final[_MemberDeleteTx] = tx
|
||||
existing_user_rows: Final = await member_tx.litellm_usertable.find_many(where=key_val)
|
||||
|
||||
# Also clean up any existing team membership rows for this user and team
|
||||
user_ids_to_delete: Final = removed_user_ids.union(
|
||||
(data.user_id,) if data.user_id is not None else (),
|
||||
(user.user_id for user in existing_user_rows if user.user_id),
|
||||
)
|
||||
# A user row can outlive its roster entry, and until the team is off user.teams the user
|
||||
# still sees it and still fails key creation against it, so removal has to clear it too
|
||||
stale_user_rows: Final = tuple(user for user in existing_user_rows if data.team_id in user.teams)
|
||||
|
||||
# Also clean up any existing team membership rows for this user and team. An email can
|
||||
# match several user rows, so with no roster entry to name the member, only the rows
|
||||
# actually carrying the team are the ones this request is allowed to touch
|
||||
cleanup_user_rows: Final = existing_user_rows if removed_team_members else stale_user_rows
|
||||
user_ids_to_delete: Final = addressed_user_ids.union(user.user_id for user in cleanup_user_rows if user.user_id)
|
||||
|
||||
if not removed_team_members and not stale_user_rows:
|
||||
raise HTTPException(status_code=400, detail={"error": "User not found in team"})
|
||||
|
||||
## DELETE KEYS CREATED BY USER FOR THIS TEAM
|
||||
# Fetch keys before deletion so their audit records can be persisted alongside the delete.
|
||||
|
|
@ -3358,17 +3365,17 @@ async def team_member_delete(
|
|||
}
|
||||
)
|
||||
|
||||
await _team_tx_db(tx).update(
|
||||
where={"team_id": data.team_id},
|
||||
data={"members_with_roles": json.dumps(_db_new_team_members)},
|
||||
)
|
||||
if removed_team_members:
|
||||
await _team_tx_db(tx).update(
|
||||
where={"team_id": data.team_id},
|
||||
data={"members_with_roles": json.dumps(_db_new_team_members)},
|
||||
)
|
||||
|
||||
for existing_user in existing_user_rows:
|
||||
if data.team_id in existing_user.teams:
|
||||
await tx.litellm_usertable.update(
|
||||
where={"user_id": existing_user.user_id},
|
||||
data={"teams": {"set": [team for team in existing_user.teams if team != data.team_id]}},
|
||||
)
|
||||
for existing_user in stale_user_rows:
|
||||
await tx.litellm_usertable.update(
|
||||
where={"user_id": existing_user.user_id},
|
||||
data={"teams": {"set": [team for team in existing_user.teams if team != data.team_id]}},
|
||||
)
|
||||
|
||||
for _uid in sorted(user_ids_to_delete):
|
||||
await tx.litellm_teammembership.delete_many(where={"team_id": data.team_id, "user_id": _uid})
|
||||
|
|
|
|||
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})
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -7157,6 +7157,53 @@
|
|||
"supports_xhigh_reasoning_effort": true,
|
||||
"supports_minimal_reasoning_effort": false
|
||||
},
|
||||
"azure/gpt-6-astra": {
|
||||
"cache_creation_input_token_cost": 1.25e-05,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-05,
|
||||
"cache_read_input_token_cost": 1e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 2e-06,
|
||||
"input_cost_per_token": 1e-05,
|
||||
"input_cost_per_token_above_272k_tokens": 2e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 922000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 7.5e-05,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_high": 0.01,
|
||||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_max_reasoning_effort": true,
|
||||
"supports_minimal_reasoning_effort": false,
|
||||
"supports_native_streaming": true,
|
||||
"supports_none_reasoning_effort": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_xhigh_reasoning_effort": true
|
||||
},
|
||||
"azure/us/gpt-5.6": {
|
||||
"cache_creation_input_token_cost": 6.875e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1.375e-05,
|
||||
|
|
@ -7376,6 +7423,53 @@
|
|||
"supports_xhigh_reasoning_effort": true,
|
||||
"supports_minimal_reasoning_effort": false
|
||||
},
|
||||
"azure/us/gpt-6-astra": {
|
||||
"cache_creation_input_token_cost": 1.375e-05,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 2.75e-05,
|
||||
"cache_read_input_token_cost": 1.1e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 2.2e-06,
|
||||
"input_cost_per_token": 1.1e-05,
|
||||
"input_cost_per_token_above_272k_tokens": 2.2e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 922000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5.5e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 8.25e-05,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_high": 0.01,
|
||||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_max_reasoning_effort": true,
|
||||
"supports_minimal_reasoning_effort": false,
|
||||
"supports_native_streaming": true,
|
||||
"supports_none_reasoning_effort": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_xhigh_reasoning_effort": true
|
||||
},
|
||||
"azure/eu/gpt-5.6": {
|
||||
"cache_creation_input_token_cost": 6.875e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1.375e-05,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -2008,6 +2008,49 @@ def test_generic_cost_per_token_azure_gpt56(_local_model_cost_map,
|
|||
assert round(completion_cost, 10) == round(output_cost * completion_tokens, 10)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model,zone_multiplier", [("azure/gpt-6-astra", 1.0), ("azure/us/gpt-6-astra", 1.1)])
|
||||
@pytest.mark.parametrize(
|
||||
"prompt_tokens,input_side_multiplier,output_multiplier",
|
||||
[(100000, 1.0, 1.0), (300000, 2.0, 1.5)],
|
||||
)
|
||||
def test_generic_cost_per_token_azure_gpt_6_astra_foundry_price_sheet(
|
||||
_local_model_cost_map,
|
||||
model,
|
||||
zone_multiplier,
|
||||
prompt_tokens,
|
||||
input_side_multiplier,
|
||||
output_multiplier,
|
||||
):
|
||||
"""Microsoft Foundry sells gpt-6-astra at the OpenAI rates: $10 input, $1 cache read, $12.50 cache write,
|
||||
$50 output per 1M tokens on Standard Global, with the input side doubling and output 1.5x above 272K
|
||||
prompt tokens. Standard US Data Zone carries the usual 10% uplift on every rate.
|
||||
"""
|
||||
cached_tokens = 50000
|
||||
cache_write_tokens = 40000
|
||||
text_tokens = prompt_tokens - cached_tokens - cache_write_tokens
|
||||
completion_tokens = 1000
|
||||
usage = Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=prompt_tokens + completion_tokens,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
cached_tokens=cached_tokens, cache_write_tokens=cache_write_tokens
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, completion_cost = generic_cost_per_token(
|
||||
model=model,
|
||||
usage=usage,
|
||||
custom_llm_provider="azure",
|
||||
)
|
||||
|
||||
input_side = zone_multiplier * input_side_multiplier
|
||||
assert prompt_cost == pytest.approx(
|
||||
input_side * (text_tokens * 1e-5 + cached_tokens * 1e-6 + cache_write_tokens * 1.25e-5)
|
||||
)
|
||||
assert completion_cost == pytest.approx(zone_multiplier * output_multiplier * completion_tokens * 5e-5)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model,expected_none,expected_xhigh,expected_minimal",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -348,3 +348,31 @@ def test_azure_gpt_6_astra_takes_the_reasoning_series_request_shape():
|
|||
assert params["max_completion_tokens"] == 100
|
||||
assert "max_tokens" not in params
|
||||
assert params["reasoning_effort"] == "max"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", ["azure/gpt-6-astra", "azure/us/gpt-6-astra"])
|
||||
def test_azure_gpt6_astra_reasoning_effort_none_unlocks_temperature(config: AzureOpenAIGPT5Config, model: str):
|
||||
"""Foundry's gpt-6-astra accepts reasoning_effort='none' and, only then, a non-default
|
||||
temperature (verified live against a Foundry deployment), unlike OpenAI's gpt-6-astra."""
|
||||
params = config.map_openai_params(
|
||||
non_default_params={"temperature": 0.2, "reasoning_effort": "none"},
|
||||
optional_params={},
|
||||
model=model,
|
||||
drop_params=False,
|
||||
api_version="2025-04-01-preview",
|
||||
)
|
||||
assert params["temperature"] == 0.2
|
||||
assert params["reasoning_effort"] == "none"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", ["azure/gpt-6-astra", "azure/us/gpt-6-astra"])
|
||||
def test_azure_gpt6_astra_rejects_reasoning_effort_minimal(config: AzureOpenAIGPT5Config, model: str):
|
||||
"""Foundry's gpt-6-astra lists none, low, medium, high, xhigh and max but not minimal."""
|
||||
with pytest.raises(litellm.utils.UnsupportedParamsError):
|
||||
config.map_openai_params(
|
||||
non_default_params={"reasoning_effort": "minimal"},
|
||||
optional_params={},
|
||||
model=model,
|
||||
drop_params=False,
|
||||
api_version="2025-04-01-preview",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,11 +1,10 @@
|
|||
from copy import deepcopy
|
||||
from unittest.mock import patch
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map
|
||||
from litellm.llms.azure.responses.o_series_transformation import (
|
||||
AzureOpenAIOSeriesResponsesAPIConfig,
|
||||
)
|
||||
|
|
@ -613,3 +612,39 @@ class TestAzureResponsesAPIConfig:
|
|||
|
||||
assert result["tools"][0] is tool
|
||||
assert "anyOf" in result["tools"][0]["parameters"]
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def local_model_cost_map(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Pin the bundled cost map: the published map lags a key added in this repo."""
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
monkeypatch.setattr(litellm, "model_cost", get_model_cost_map(url=litellm.model_cost_map_url))
|
||||
litellm.add_known_models(model_cost_map=litellm.model_cost)
|
||||
|
||||
|
||||
def test_azure_responses_gpt6_astra_reasoning_effort_none_unlocks_temperature(local_model_cost_map: None):
|
||||
"""Foundry's gpt-6-astra accepts reasoning.effort='none' with a non-default temperature
|
||||
while OpenAI's gpt-6-astra does not, so the gate must read the azure/ cost-map entry
|
||||
for the bare deployment name rather than OpenAI's."""
|
||||
params = AzureOpenAIResponsesAPIConfig().map_openai_params(
|
||||
response_api_optional_params=ResponsesAPIOptionalRequestParams(
|
||||
temperature=0.2,
|
||||
reasoning={"effort": "none"},
|
||||
),
|
||||
model="gpt-6-astra",
|
||||
drop_params=False,
|
||||
)
|
||||
assert params["temperature"] == 0.2
|
||||
assert params["reasoning"] == {"effort": "none"}
|
||||
|
||||
|
||||
def test_azure_responses_gpt6_astra_rejects_temperature_while_reasoning(local_model_cost_map: None):
|
||||
with pytest.raises(litellm.UnsupportedParamsError):
|
||||
AzureOpenAIResponsesAPIConfig().map_openai_params(
|
||||
response_api_optional_params=ResponsesAPIOptionalRequestParams(
|
||||
temperature=0.2,
|
||||
reasoning={"effort": "low"},
|
||||
),
|
||||
model="gpt-6-astra",
|
||||
drop_params=False,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -4863,6 +4863,302 @@ async def test_team_member_delete_by_email_the_user_row_does_not_carry(
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_member_delete_clears_team_left_on_the_user_row_without_a_roster_entry(
|
||||
mock_db_client, mock_admin_auth
|
||||
):
|
||||
"""
|
||||
A user row can keep a team (several times over, from older duplicate-prone adds) after the
|
||||
roster entry is gone, which leaves the team listed on the user, offered in the key creation
|
||||
dropdown, and rejected by key creation itself. Reporting "User not found in team" left that
|
||||
residue unremovable, so the delete now cleans every copy of the team off the user row.
|
||||
"""
|
||||
from litellm.proxy._types import TeamMemberDeleteRequest
|
||||
from litellm.proxy.management_endpoints.team_endpoints import team_member_delete
|
||||
|
||||
test_team_id = "team-del-orphan-123"
|
||||
test_user_id = "user-del-orphan-123"
|
||||
|
||||
mock_team_row = MagicMock()
|
||||
mock_team_row.model_dump.return_value = {
|
||||
"team_id": test_team_id,
|
||||
"members_with_roles": [],
|
||||
"team_member_permissions": [],
|
||||
"metadata": {},
|
||||
"models": [],
|
||||
"spend": 0.0,
|
||||
}
|
||||
mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(
|
||||
return_value=mock_team_row
|
||||
)
|
||||
mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_team_row)
|
||||
|
||||
mock_user_row = MagicMock()
|
||||
mock_user_row.user_id = test_user_id
|
||||
mock_user_row.user_email = None
|
||||
mock_user_row.teams = [test_team_id, "other-team", test_team_id]
|
||||
mock_db_client.db.litellm_usertable.find_many = AsyncMock(
|
||||
return_value=[mock_user_row]
|
||||
)
|
||||
mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock())
|
||||
|
||||
mock_db_client.db.litellm_teammembership = MagicMock()
|
||||
mock_db_client.db.litellm_teammembership.delete_many = AsyncMock(
|
||||
return_value=MagicMock()
|
||||
)
|
||||
|
||||
mock_db_client.db.litellm_verificationtoken = MagicMock()
|
||||
mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
|
||||
mock_db_client.db.litellm_verificationtoken.delete_many = AsyncMock(
|
||||
return_value=MagicMock()
|
||||
)
|
||||
|
||||
_wire_member_delete_tx(mock_db_client)
|
||||
|
||||
await team_member_delete(
|
||||
data=TeamMemberDeleteRequest(team_id=test_team_id, user_id=test_user_id),
|
||||
user_api_key_dict=mock_admin_auth,
|
||||
)
|
||||
|
||||
mock_db_client.db.litellm_usertable.update.assert_awaited_once_with(
|
||||
where={"user_id": test_user_id},
|
||||
data={"teams": {"set": ["other-team"]}},
|
||||
)
|
||||
mock_db_client.db.litellm_teammembership.delete_many.assert_awaited_once_with(
|
||||
where={"team_id": test_team_id, "user_id": test_user_id}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_member_delete_still_rejects_a_user_the_team_has_no_trace_of(
|
||||
mock_db_client, mock_admin_auth
|
||||
):
|
||||
from litellm.proxy._types import TeamMemberDeleteRequest
|
||||
from litellm.proxy.management_endpoints.team_endpoints import team_member_delete
|
||||
|
||||
test_team_id = "team-del-absent-123"
|
||||
test_user_id = "user-del-absent-123"
|
||||
|
||||
mock_team_row = MagicMock()
|
||||
mock_team_row.model_dump.return_value = {
|
||||
"team_id": test_team_id,
|
||||
"members_with_roles": [],
|
||||
"team_member_permissions": [],
|
||||
"metadata": {},
|
||||
"models": [],
|
||||
"spend": 0.0,
|
||||
}
|
||||
mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(
|
||||
return_value=mock_team_row
|
||||
)
|
||||
mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_team_row)
|
||||
|
||||
mock_user_row = MagicMock()
|
||||
mock_user_row.user_id = test_user_id
|
||||
mock_user_row.user_email = None
|
||||
mock_user_row.teams = ["other-team"]
|
||||
mock_db_client.db.litellm_usertable.find_many = AsyncMock(
|
||||
return_value=[mock_user_row]
|
||||
)
|
||||
mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock())
|
||||
|
||||
mock_db_client.db.litellm_teammembership = MagicMock()
|
||||
mock_db_client.db.litellm_teammembership.delete_many = AsyncMock(
|
||||
return_value=MagicMock()
|
||||
)
|
||||
|
||||
_wire_member_delete_tx(mock_db_client)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await team_member_delete(
|
||||
data=TeamMemberDeleteRequest(team_id=test_team_id, user_id=test_user_id),
|
||||
user_api_key_dict=mock_admin_auth,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert exc_info.value.detail == {"error": "User not found in team"}
|
||||
mock_db_client.db.litellm_usertable.update.assert_not_awaited()
|
||||
mock_db_client.db.litellm_teammembership.delete_many.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_member_delete_leaves_a_bystander_named_by_a_conflicting_user_id_alone(
|
||||
mock_db_client, mock_admin_auth
|
||||
):
|
||||
"""
|
||||
A request can carry a user_id and a user_email that point at two different people, and only the
|
||||
email matches a roster entry. Cleaning up both ids would strip the team, the membership row and
|
||||
the keys off the bystander the roster never listed, so the user_id only widens the cleanup when
|
||||
the roster came back empty.
|
||||
"""
|
||||
from litellm.proxy._types import TeamMemberDeleteRequest
|
||||
from litellm.proxy.management_endpoints.team_endpoints import team_member_delete
|
||||
|
||||
test_team_id = "team-del-conflict-123"
|
||||
roster_user_id = "user-del-conflict-roster"
|
||||
bystander_user_id = "user-del-conflict-bystander"
|
||||
roster_email = "roster@example.com"
|
||||
|
||||
mock_team_row = MagicMock()
|
||||
mock_team_row.model_dump.return_value = {
|
||||
"team_id": test_team_id,
|
||||
"members_with_roles": [
|
||||
{"user_id": roster_user_id, "user_email": roster_email, "role": "user"}
|
||||
],
|
||||
"team_member_permissions": [],
|
||||
"metadata": {},
|
||||
"models": [],
|
||||
"spend": 0.0,
|
||||
}
|
||||
mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(
|
||||
return_value=mock_team_row
|
||||
)
|
||||
mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_team_row)
|
||||
|
||||
roster_user_row = MagicMock()
|
||||
roster_user_row.user_id = roster_user_id
|
||||
roster_user_row.user_email = roster_email
|
||||
roster_user_row.teams = [test_team_id]
|
||||
|
||||
bystander_user_row = MagicMock()
|
||||
bystander_user_row.user_id = bystander_user_id
|
||||
bystander_user_row.user_email = "bystander@example.com"
|
||||
bystander_user_row.teams = [test_team_id]
|
||||
|
||||
rows_by_user_id = {
|
||||
roster_user_id: roster_user_row,
|
||||
bystander_user_id: bystander_user_row,
|
||||
}
|
||||
|
||||
async def find_user_rows(where):
|
||||
user_id_filter = where.get("user_id")
|
||||
if isinstance(user_id_filter, dict):
|
||||
return [
|
||||
rows_by_user_id[uid]
|
||||
for uid in user_id_filter.get("in", [])
|
||||
if uid in rows_by_user_id
|
||||
]
|
||||
return [
|
||||
row
|
||||
for row in rows_by_user_id.values()
|
||||
if row.user_email == where.get("user_email")
|
||||
]
|
||||
|
||||
mock_db_client.db.litellm_usertable.find_many = AsyncMock(
|
||||
side_effect=find_user_rows
|
||||
)
|
||||
mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock())
|
||||
|
||||
mock_db_client.db.litellm_teammembership = MagicMock()
|
||||
mock_db_client.db.litellm_teammembership.delete_many = AsyncMock(
|
||||
return_value=MagicMock()
|
||||
)
|
||||
|
||||
mock_db_client.db.litellm_verificationtoken = MagicMock()
|
||||
mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
|
||||
mock_db_client.db.litellm_verificationtoken.delete_many = AsyncMock(
|
||||
return_value=MagicMock()
|
||||
)
|
||||
|
||||
_wire_member_delete_tx(mock_db_client)
|
||||
|
||||
await team_member_delete(
|
||||
data=TeamMemberDeleteRequest(
|
||||
team_id=test_team_id,
|
||||
user_id=bystander_user_id,
|
||||
user_email=roster_email,
|
||||
),
|
||||
user_api_key_dict=mock_admin_auth,
|
||||
)
|
||||
|
||||
mock_db_client.db.litellm_usertable.update.assert_awaited_once_with(
|
||||
where={"user_id": roster_user_id},
|
||||
data={"teams": {"set": []}},
|
||||
)
|
||||
mock_db_client.db.litellm_teammembership.delete_many.assert_awaited_once_with(
|
||||
where={"team_id": test_team_id, "user_id": roster_user_id}
|
||||
)
|
||||
mock_db_client.db.litellm_verificationtoken.delete_many.assert_awaited_once_with(
|
||||
where={"user_id": {"in": [roster_user_id]}, "team_id": test_team_id}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_member_delete_by_email_only_touches_the_row_carrying_the_stale_team(
|
||||
mock_db_client, mock_admin_auth
|
||||
):
|
||||
"""
|
||||
user_email is not unique, so an email delete against an empty roster can match several user
|
||||
rows. Only the row that actually carries the team is stale; the namesake keeps its team, its
|
||||
membership row and its keys.
|
||||
"""
|
||||
from litellm.proxy._types import TeamMemberDeleteRequest
|
||||
from litellm.proxy.management_endpoints.team_endpoints import team_member_delete
|
||||
|
||||
test_team_id = "team-del-shared-email-123"
|
||||
stale_user_id = "user-del-shared-email-stale"
|
||||
namesake_user_id = "user-del-shared-email-namesake"
|
||||
shared_email = "shared@example.com"
|
||||
|
||||
mock_team_row = MagicMock()
|
||||
mock_team_row.model_dump.return_value = {
|
||||
"team_id": test_team_id,
|
||||
"members_with_roles": [],
|
||||
"team_member_permissions": [],
|
||||
"metadata": {},
|
||||
"models": [],
|
||||
"spend": 0.0,
|
||||
}
|
||||
mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(
|
||||
return_value=mock_team_row
|
||||
)
|
||||
mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_team_row)
|
||||
|
||||
stale_user_row = MagicMock()
|
||||
stale_user_row.user_id = stale_user_id
|
||||
stale_user_row.user_email = shared_email
|
||||
stale_user_row.teams = [test_team_id]
|
||||
|
||||
namesake_user_row = MagicMock()
|
||||
namesake_user_row.user_id = namesake_user_id
|
||||
namesake_user_row.user_email = shared_email
|
||||
namesake_user_row.teams = ["other-team"]
|
||||
|
||||
mock_db_client.db.litellm_usertable.find_many = AsyncMock(
|
||||
return_value=[stale_user_row, namesake_user_row]
|
||||
)
|
||||
mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock())
|
||||
|
||||
mock_db_client.db.litellm_teammembership = MagicMock()
|
||||
mock_db_client.db.litellm_teammembership.delete_many = AsyncMock(
|
||||
return_value=MagicMock()
|
||||
)
|
||||
|
||||
mock_db_client.db.litellm_verificationtoken = MagicMock()
|
||||
mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
|
||||
mock_db_client.db.litellm_verificationtoken.delete_many = AsyncMock(
|
||||
return_value=MagicMock()
|
||||
)
|
||||
|
||||
_wire_member_delete_tx(mock_db_client)
|
||||
|
||||
await team_member_delete(
|
||||
data=TeamMemberDeleteRequest(team_id=test_team_id, user_email=shared_email),
|
||||
user_api_key_dict=mock_admin_auth,
|
||||
)
|
||||
|
||||
mock_db_client.db.litellm_usertable.update.assert_awaited_once_with(
|
||||
where={"user_id": stale_user_id},
|
||||
data={"teams": {"set": []}},
|
||||
)
|
||||
mock_db_client.db.litellm_teammembership.delete_many.assert_awaited_once_with(
|
||||
where={"team_id": test_team_id, "user_id": stale_user_id}
|
||||
)
|
||||
mock_db_client.db.litellm_verificationtoken.delete_many.assert_awaited_once_with(
|
||||
where={"user_id": {"in": [stale_user_id]}, "team_id": test_team_id}
|
||||
)
|
||||
|
||||
|
||||
class _InjectedMemberDeleteFailure(Exception):
|
||||
pass
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -388,3 +388,21 @@ class TestGpt6AstraAdvertisesItsDocumentedLevels:
|
|||
"xhigh",
|
||||
"max",
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("model", ["azure/gpt-6-astra", "azure/us/gpt-6-astra"])
|
||||
def test_a_foundry_deployment_also_advertises_none(self, local_model_cost_map, model):
|
||||
"""Microsoft Foundry serves the same model but its API accepts reasoning_effort none
|
||||
(verified live: 200 with zero reasoning tokens, and it unlocks temperature), which
|
||||
OpenAI's rejects, so an Azure deployment offers none on top of low through max."""
|
||||
from litellm.utils import _get_model_info_helper
|
||||
|
||||
model_info = dict(_get_model_info_helper(model=model, custom_llm_provider="azure"))
|
||||
|
||||
assert resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=True) == (
|
||||
"none",
|
||||
"low",
|
||||
"medium",
|
||||
"high",
|
||||
"xhigh",
|
||||
"max",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
42
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
42
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -22764,22 +22764,40 @@ export interface components {
|
|||
/** Spend */
|
||||
spend?: number | null;
|
||||
};
|
||||
/**
|
||||
* AccessGroupResource
|
||||
* @description A resource referenced by an access group. `name` is null when the id no longer resolves or has no alias.
|
||||
*/
|
||||
AccessGroupResource: {
|
||||
/** Id */
|
||||
id: string;
|
||||
/** Name */
|
||||
name: string | null;
|
||||
};
|
||||
/** AccessGroupResponse */
|
||||
AccessGroupResponse: {
|
||||
/** Access Agent Ids */
|
||||
access_agent_ids: string[];
|
||||
/** Access Agents */
|
||||
access_agents: components["schemas"]["AccessGroupResource"][];
|
||||
/** Access Group Id */
|
||||
access_group_id: string;
|
||||
/** Access Group Name */
|
||||
access_group_name: string;
|
||||
/** Access Mcp Server Ids */
|
||||
access_mcp_server_ids: string[];
|
||||
/** Access Mcp Servers */
|
||||
access_mcp_servers: components["schemas"]["AccessGroupResource"][];
|
||||
/** Access Model Names */
|
||||
access_model_names: string[];
|
||||
/** Assigned Key Ids */
|
||||
assigned_key_ids: string[];
|
||||
/** Assigned Keys */
|
||||
assigned_keys: components["schemas"]["AccessGroupResource"][];
|
||||
/** Assigned Team Ids */
|
||||
assigned_team_ids: string[];
|
||||
/** Assigned Teams */
|
||||
assigned_teams: components["schemas"]["AccessGroupResource"][];
|
||||
/**
|
||||
* Created At
|
||||
* Format: date-time
|
||||
|
|
@ -25305,6 +25323,30 @@ export interface components {
|
|||
* @default 3000
|
||||
*/
|
||||
timeout_ms: number;
|
||||
/** @description Whether the classifier sees images on the request, and how many */
|
||||
vision?: components["schemas"]["ClassifierVisionConfig"];
|
||||
};
|
||||
/**
|
||||
* ClassifierVisionConfig
|
||||
* @description Whether the LLM classifier sees the images on the request it is classifying.
|
||||
*
|
||||
* Off by default because images cost far more than the text ask they arrive with, and the
|
||||
* classifier runs on every request. A turn whose complexity lives in the image ("what is wrong in
|
||||
* this stack trace screenshot") is invisible to a text-only classifier, which is what this buys.
|
||||
*/
|
||||
ClassifierVisionConfig: {
|
||||
/**
|
||||
* Enabled
|
||||
* @description Forward image content to the classifier. Requires a classifier model declared supports_vision, on the deployment's model_info or in the model cost map; images stay stripped otherwise, so a classifier that cannot read them is never sent one. Declare model_info.supports_vision on the deployment to enable a model the cost map does not describe. Only inline data: URIs are forwarded. A request whose images are http(s) URLs still classifies on its text alone, because some providers fetch such a URL from the proxy rather than the provider, which would let a caller aim a proxy-side request at an address of their choosing.
|
||||
* @default false
|
||||
*/
|
||||
enabled: boolean;
|
||||
/**
|
||||
* Max Images
|
||||
* @description How many images from the newest user turn to forward, in wire order. Bounds the added cost of a turn that attaches many images. Images on earlier turns are never forwarded.
|
||||
* @default 1
|
||||
*/
|
||||
max_images: number;
|
||||
};
|
||||
/**
|
||||
* CloudZeroExportRequest
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue