mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
chore: merge litellm_internal_staging into litellm_lit_6348_fireworks_responses_api
This commit is contained in:
commit
525d9fb14c
124 changed files with 8035 additions and 653 deletions
2
.github/workflows/_test-unit-base.yml
vendored
2
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -116,7 +116,7 @@ jobs:
|
|||
if: steps.changes.outputs.decision != 'skip'
|
||||
timeout-minutes: 8
|
||||
run: |
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra mongodb
|
||||
uv run --no-sync python -c 'import os, sys; print(sys.version); assert f"{sys.version_info.major}.{sys.version_info.minor}" == os.environ["UV_PYTHON"]'
|
||||
|
||||
- name: Cache Prisma binaries
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"reportAny": {
|
||||
"limit": 14074
|
||||
"limit": 14072
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"limit": 2206
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 19
|
||||
},
|
||||
"reportExplicitAny": {
|
||||
"limit": 4124
|
||||
"limit": 4121
|
||||
},
|
||||
"reportFunctionMemberAccess": {
|
||||
"limit": 7
|
||||
|
|
@ -108,7 +108,7 @@
|
|||
"limit": 38307
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19621
|
||||
"limit": 19620
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 29838
|
||||
|
|
|
|||
|
|
@ -0,0 +1 @@
|
|||
ALTER TABLE "LiteLLM_ShadowEvalJob" ADD COLUMN IF NOT EXISTS "models" TEXT[] NOT NULL DEFAULT ARRAY[]::TEXT[];
|
||||
|
|
@ -1536,6 +1536,7 @@ model LiteLLM_ShadowEvalJob {
|
|||
target_id String // hashed virtual key, team_id, or user_id whose traffic this leg shadows
|
||||
router_name String // first (often only) auto-router under evaluation; router_names is the full set
|
||||
router_names String[] @default([]) // all routers this job runs as shadow arms; empty on legacy rows, whose set is (router_name)
|
||||
models String[] @default([]) // model groups the sampled traffic is narrowed to; empty samples every model
|
||||
direction String @default("forward") // forward | reverse
|
||||
baseline_model String? // reverse only: the fixed model the router is judged against
|
||||
judge_model String
|
||||
|
|
|
|||
|
|
@ -257,7 +257,7 @@ class DualCache(BaseCache):
|
|||
self,
|
||||
current_time: float,
|
||||
keys: list[str],
|
||||
result: Sequence[Any],
|
||||
result: Sequence[object],
|
||||
) -> tuple[list[str], dict[str, float | None]]:
|
||||
"""
|
||||
Atomically choose keys to fetch from Redis and reserve their access time.
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ import httpx
|
|||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.constants import REDACTED_BY_LITELLM
|
||||
from litellm.constants import REDACTED_BY_LITELLM, REDACTED_BY_LITELM_STRING
|
||||
from litellm.integrations.custom_batch_logger import CustomBatchLogger
|
||||
from litellm.integrations.datadog.datadog_handler import (
|
||||
get_datadog_base_url_from_env,
|
||||
|
|
@ -46,9 +46,10 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
from litellm.proxy.spend_tracking.savings import extract_cache_creation_tokens, extract_cache_read_tokens
|
||||
from litellm.types.integrations.datadog_llm_obs import *
|
||||
from litellm.types.utils import (
|
||||
AUDIT_GUARDRAIL_FIELDS,
|
||||
PROMPT_CARRYING_GUARDRAIL_FIELDS,
|
||||
PROMPT_QUOTING_ROUTING_DECISION_FIELDS,
|
||||
CallTypes,
|
||||
StandardLoggingGuardrailInformation,
|
||||
StandardLoggingPayload,
|
||||
StandardLoggingPayloadErrorInformation,
|
||||
)
|
||||
|
|
@ -60,6 +61,8 @@ _SAFE_REDACTED_MESSAGE_ROLES: Final = frozenset(
|
|||
{"agent", "assistant", "developer", "function", "model", "system", "tool", "user"}
|
||||
)
|
||||
|
||||
_CLASSIFIED_GUARDRAIL_FIELDS: Final = AUDIT_GUARDRAIL_FIELDS | PROMPT_CARRYING_GUARDRAIL_FIELDS
|
||||
|
||||
_PROMPT_CARRYING_METADATA_FIELDS: Final = frozenset(
|
||||
{
|
||||
"routing_decision",
|
||||
|
|
@ -108,6 +111,49 @@ def _router_span_fields(
|
|||
)
|
||||
|
||||
|
||||
def _guardrail_entries(guardrail_information: object) -> tuple[Mapping[str, object], ...]:
|
||||
"""The guardrail records as a sequence, whatever shape the payload carries.
|
||||
|
||||
`guardrail_information` is typed as a list, but a guardrail that writes the metadata key itself
|
||||
can leave a single record there; Prometheus normalizes the same shape at
|
||||
`_guardrail_overhead_seconds`.
|
||||
"""
|
||||
if isinstance(guardrail_information, Mapping):
|
||||
return (guardrail_information,)
|
||||
if isinstance(guardrail_information, (list, tuple)):
|
||||
return tuple(entry for entry in guardrail_information if isinstance(entry, Mapping))
|
||||
return ()
|
||||
|
||||
|
||||
def _guardrail_entry_without_prompt_carriers(entry: Mapping[str, object]) -> Mapping[str, object]:
|
||||
"""One guardrail record kept as its audit fields, with the prompt-quoting ones marked redacted.
|
||||
|
||||
Built as an allow-list rather than a deny-list: a key neither set classifies is dropped, so a
|
||||
guardrail that records its own extra detail cannot put the caller's prompt on a redacted span.
|
||||
"""
|
||||
return { # mutable-ok: a fresh record built per entry, handed straight to the span serializer
|
||||
field: REDACTED_BY_LITELM_STRING if field in PROMPT_CARRYING_GUARDRAIL_FIELDS else value
|
||||
for field, value in entry.items()
|
||||
if field in _CLASSIFIED_GUARDRAIL_FIELDS
|
||||
}
|
||||
|
||||
|
||||
def _guardrail_information_without_prompt_carriers(
|
||||
guardrail_information: object,
|
||||
) -> tuple[Mapping[str, object], ...] | None:
|
||||
"""The guardrail records reduced to what a redacted span may carry.
|
||||
|
||||
Redaction removes the prompt, not the record that a guardrail ran: the name, mode, status,
|
||||
timings and masked-entity counts are what an operator reads to answer whether a guardrail
|
||||
caught anything on a request, and none of them reproduce the prompt. Field-level rather than
|
||||
dropping the list, which is what `_sanitize_guardrail_information_for_spend_logs` already does
|
||||
for spend logs.
|
||||
"""
|
||||
if guardrail_information is None:
|
||||
return None
|
||||
return tuple(_guardrail_entry_without_prompt_carriers(entry) for entry in _guardrail_entries(guardrail_information))
|
||||
|
||||
|
||||
def _metadata_without_prompt_carriers(standard_logging_metadata: Mapping[str, Any]) -> Mapping[str, Any]:
|
||||
"""The metadata minus the records that quote prompts, tool arguments, tool results, or retrieved text."""
|
||||
return MappingProxyType(
|
||||
|
|
@ -872,7 +918,9 @@ class DataDogLLMObsLogger(CustomBatchLogger):
|
|||
"cache_key": standard_logging_payload.get("cache_key", "unknown"),
|
||||
"saved_cache_cost": standard_logging_payload.get("saved_cache_cost", 0),
|
||||
"guardrail_information": (
|
||||
None if redact_prompt_text else standard_logging_payload.get("guardrail_information", None)
|
||||
_guardrail_information_without_prompt_carriers(standard_logging_payload.get("guardrail_information"))
|
||||
if redact_prompt_text
|
||||
else standard_logging_payload.get("guardrail_information", None)
|
||||
),
|
||||
"is_streamed_request": self._get_stream_value_from_payload(standard_logging_payload),
|
||||
"latency_metrics": dict(self._get_latency_metrics(standard_logging_payload)),
|
||||
|
|
@ -904,14 +952,12 @@ class DataDogLLMObsLogger(CustomBatchLogger):
|
|||
latency_metrics["litellm_overhead_time_ms"] = litellm_overhead_ms
|
||||
|
||||
# Guardrail overhead latency
|
||||
guardrail_info: Final[list[StandardLoggingGuardrailInformation] | None] = standard_logging_payload.get(
|
||||
"guardrail_information"
|
||||
)
|
||||
if guardrail_info is not None:
|
||||
guardrail_info: Final = _guardrail_entries(standard_logging_payload.get("guardrail_information"))
|
||||
if guardrail_info:
|
||||
total_duration = 0.0
|
||||
for info in guardrail_info:
|
||||
_guardrail_duration_seconds: float | None = info.get("duration")
|
||||
if _guardrail_duration_seconds is not None:
|
||||
_guardrail_duration_seconds = info.get("duration")
|
||||
if isinstance(_guardrail_duration_seconds, (int, float, str)):
|
||||
total_duration += float(_guardrail_duration_seconds)
|
||||
|
||||
if total_duration > 0:
|
||||
|
|
|
|||
|
|
@ -35,6 +35,15 @@ span orphaned into its own trace). The anchor — a contextvar inherited by thos
|
|||
child tasks — gives a stable parent in both cases. DB/service spans keep ambient
|
||||
parenting so an auth DB lookup still nests under `auth`.
|
||||
|
||||
The anchor is also what `litellm.request.route` is read from: `request_root_http_route`
|
||||
returns the server span's own `http.route`, so the LLM call span cannot disagree with
|
||||
its parent about which endpoint served the request. That means the route template on a
|
||||
normal route and the literal path on a passthrough prefix, because the passthrough hook
|
||||
rewrote the attribute; an MCP call anchors the same server span, so it reports the
|
||||
`/mcp` mount point. Attributes stay readable after a span ends, so the async close
|
||||
callback reads the same value. Where no server span was anchored at all, the route the
|
||||
proxy recorded at auth (`metadata.user_api_key_request_route`) is the backstop.
|
||||
|
||||
**Which service calls become spans (`spans.span_role_for_service`).** LiteLLM's
|
||||
service-logging layer instruments many internal functions, but only some are
|
||||
traceable units of work:
|
||||
|
|
|
|||
|
|
@ -49,6 +49,7 @@ from litellm.integrations.otel.model.utils import to_ns
|
|||
from litellm.integrations.otel.plumbing.context import (
|
||||
is_recordable_span,
|
||||
mcp_message_transport_span,
|
||||
request_root_http_route,
|
||||
request_root_span,
|
||||
resolve_mcp_span_context,
|
||||
resolve_parent_context,
|
||||
|
|
@ -541,6 +542,7 @@ class OpenTelemetryV2(CustomLogger):
|
|||
payload,
|
||||
capture_content=self.config.capture_span_content,
|
||||
time_to_first_chunk_seconds=call.time_to_first_chunk_seconds,
|
||||
request_route=request_root_http_route(),
|
||||
)
|
||||
end_time_ns: Final = to_ns(end_time)
|
||||
if carrier is not None and carrier.span is not None:
|
||||
|
|
|
|||
|
|
@ -89,6 +89,7 @@ class GenAIMapper:
|
|||
f"{LiteLLM.COST_PREFIX}margin_percent": lambda d: d.cost.margin_percent,
|
||||
f"{LiteLLM.COST_PREFIX}margin_total_amount": lambda d: d.cost.margin_total_amount,
|
||||
LiteLLM.REQUEST_STREAMING: lambda d: d.is_streaming,
|
||||
LiteLLM.REQUEST_ROUTE: lambda d: d.request_route,
|
||||
}
|
||||
|
||||
_TOOL_ATTRS: dict[str, Callable[[ToolDefinition], AttrValue | None]] = {
|
||||
|
|
|
|||
|
|
@ -64,6 +64,7 @@ class RequestIdentity:
|
|||
# completes (routing has picked a deployment), so it's absent from the
|
||||
# auth-time seed and filled only from the payload.
|
||||
provider_model: str | None = None
|
||||
request_route: str | None = None
|
||||
metadata: Mapping[str, str] = field(default_factory=dict)
|
||||
|
||||
@classmethod
|
||||
|
|
@ -87,6 +88,7 @@ class RequestIdentity:
|
|||
key_hash=as_str(raw_meta.get("user_api_key_hash")),
|
||||
end_user=as_str(payload.get("end_user")) or as_str(raw_meta.get("user_api_key_end_user_id")),
|
||||
provider_model=resolve_provider_model(payload),
|
||||
request_route=as_str(raw_meta.get("user_api_key_request_route")),
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -386,6 +386,7 @@ class LLMCallSpanData:
|
|||
# keeps routes the convention folds into one operation distinguishable.
|
||||
output_type: GenAIOutputType | None = None
|
||||
call_type: str | None = None
|
||||
request_route: str | None = None
|
||||
|
||||
@classmethod
|
||||
def from_standard_logging_payload(
|
||||
|
|
@ -393,6 +394,7 @@ class LLMCallSpanData:
|
|||
payload: StandardLoggingPayload,
|
||||
capture_content: bool = False,
|
||||
time_to_first_chunk_seconds: float | None = None,
|
||||
request_route: str | None = None,
|
||||
) -> LLMCallSpanData:
|
||||
params: Final = cast(Mapping[str, object], payload.get("model_parameters") or {})
|
||||
# The single parse of the request's metadata — the request-vs-provider
|
||||
|
|
@ -433,6 +435,7 @@ class LLMCallSpanData:
|
|||
time_to_first_chunk_seconds=time_to_first_chunk_seconds,
|
||||
output_type=resolve_output_type(call_type),
|
||||
call_type=call_type or None,
|
||||
request_route=request_route or context.identity.request_route,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -295,6 +295,7 @@ class LiteLLM:
|
|||
# ``litellm_params.model``), distinct from the user-facing ``gen_ai.request.model``.
|
||||
PROVIDER_MODEL: Final = "litellm.provider.model"
|
||||
REQUEST_STREAMING: Final = "litellm.request.streaming"
|
||||
REQUEST_ROUTE: Final = "litellm.request.route"
|
||||
TOOLS_DECLARED: Final = "litellm.request.tools.declared"
|
||||
GUARDRAIL_NAME: Final = "litellm.guardrail.name"
|
||||
GUARDRAIL_MODE: Final = "litellm.guardrail.mode"
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ from typing import Final
|
|||
|
||||
from opentelemetry import baggage
|
||||
from opentelemetry.context import Context, get_current
|
||||
from opentelemetry.sdk.trace import ReadableSpan
|
||||
from opentelemetry.trace import (
|
||||
Link,
|
||||
NonRecordingSpan,
|
||||
|
|
@ -18,6 +19,8 @@ from opentelemetry.trace.propagation.tracecontext import (
|
|||
TraceContextTextMapPropagator,
|
||||
)
|
||||
|
||||
from litellm.integrations.otel.model.semconv import HTTP
|
||||
|
||||
_PROPAGATOR: Final = TraceContextTextMapPropagator()
|
||||
|
||||
# The request's root span — the FastAPI-owned SERVER span — captured ONCE when the
|
||||
|
|
@ -55,6 +58,25 @@ def request_root_span() -> "Span | None":
|
|||
return span if is_recordable_span(span) else None
|
||||
|
||||
|
||||
def request_root_http_route() -> str | None:
|
||||
"""``http.route`` exactly as the request's root SERVER span reports it.
|
||||
|
||||
Read off the span rather than re-derived, so the LLM call span cannot disagree
|
||||
with its own parent about which endpoint served the request: the template the
|
||||
instrumentation matched, or the literal path where
|
||||
``mount._passthrough_span_name_hook`` rewrote it, are already in the attribute.
|
||||
An MCP call anchors that same server span, so it reports the ``/mcp`` mount
|
||||
point the instrumentation matched. Attributes stay readable after a span ends,
|
||||
so this answers just as well from the async logging callback.
|
||||
|
||||
None when no server span is anchored, which is the SDK path and any deployment
|
||||
where the FastAPI instrumentation did not mount.
|
||||
"""
|
||||
span: Final = request_root_span()
|
||||
route: Final = span.attributes.get(HTTP.ROUTE) if isinstance(span, ReadableSpan) and span.attributes else None
|
||||
return route if isinstance(route, str) and route else None
|
||||
|
||||
|
||||
# The W3C trace-context carrier (``traceparent``/``tracestate``/``baggage``) the
|
||||
# MCP client propagated in the current request's ``params._meta``. The MCP gateway
|
||||
# sets it per message so the MCP span can record the client's span as a span
|
||||
|
|
|
|||
|
|
@ -37,6 +37,7 @@ from litellm.litellm_core_utils.llm_judge import (
|
|||
)
|
||||
from litellm.litellm_core_utils.redact_messages import should_redact_message_logging
|
||||
from litellm.llms.base_llm.base_utils import type_to_response_format_param
|
||||
from litellm.router_utils.common_utils import resolve_model_group_alias
|
||||
from litellm.types.management_endpoints.auto_router_endpoints import ShadowEvalDirection
|
||||
from litellm.types.utils import SHADOW_EVAL_JUDGE_CALL_ORIGIN, SHADOW_EVAL_ROUTER_CALL_ORIGIN
|
||||
|
||||
|
|
@ -165,28 +166,91 @@ def _chat_request_from_responses(
|
|||
)
|
||||
|
||||
|
||||
def _chat_final_text(response_obj: object) -> str:
|
||||
"""The assistant's text, or empty when the turn carries tool calls: only text-final
|
||||
turns produce a judgeable A/B comparison."""
|
||||
def _chat_choice(response_obj: object) -> object | None:
|
||||
"""The response's first choice, from a payload mapping or a duck-typed ModelResponse."""
|
||||
try:
|
||||
message: Final = (
|
||||
response_obj["choices"][0]["message"]
|
||||
if isinstance(response_obj, Mapping)
|
||||
else response_obj.choices[0].message # pyright: ignore[reportAttributeAccessIssue] # duck-typed ModelResponse
|
||||
)
|
||||
if isinstance(response_obj, Mapping):
|
||||
return response_obj["choices"][0]
|
||||
return response_obj.choices[0] # pyright: ignore[reportAttributeAccessIssue] # duck-typed ModelResponse
|
||||
except (AttributeError, KeyError, IndexError, TypeError):
|
||||
return None
|
||||
|
||||
|
||||
def _field_reader(obj: object) -> Callable[[str], object]:
|
||||
return obj.get if isinstance(obj, Mapping) else lambda key: getattr(obj, key, None)
|
||||
|
||||
|
||||
def _chat_message_reader(response_obj: object) -> Callable[[str], object] | None:
|
||||
"""Field access over the assistant message of a chat response, or None for a payload
|
||||
with no readable message."""
|
||||
choice: Final = _chat_choice(response_obj)
|
||||
if choice is None:
|
||||
return None
|
||||
message: Final = _field_reader(choice)("message")
|
||||
return _field_reader(message) if message is not None else None
|
||||
|
||||
|
||||
def _chat_final_text(response_obj: object) -> str:
|
||||
"""The turn's judgeable text: prose, or every tool call serialized alongside it as
|
||||
`[tool call] name(arguments)` when the assistant chose to act instead of, or as well
|
||||
as, answering directly. A tool call is a real turn, not a gap, so this is what both
|
||||
the real arm's sampling decision and the shadow arm's reply compare against."""
|
||||
read: Final = _chat_message_reader(response_obj)
|
||||
if read is None:
|
||||
return ""
|
||||
read: Final = message.get if isinstance(message, Mapping) else lambda key: getattr(message, key, None)
|
||||
if read("tool_calls") or read("function_call"):
|
||||
return ""
|
||||
return extract_text_from_content(read("content"))
|
||||
prose: Final = extract_text_from_content(read("content"))
|
||||
if not (read("tool_calls") or read("function_call")):
|
||||
return prose
|
||||
serialized: Final = _serialize_tool_calls(read)
|
||||
return f"{prose} {serialized}".strip() if prose else serialized
|
||||
|
||||
|
||||
def _chat_finish_reason(response_obj: object) -> str:
|
||||
choice: Final = _chat_choice(response_obj)
|
||||
raw: Final = _field_reader(choice)("finish_reason") if choice is not None else None
|
||||
return str(raw) if raw else "unknown"
|
||||
|
||||
|
||||
_RESPONSES_TOOL_CALL_TYPES: Final = frozenset(("function_call", "custom_tool_call"))
|
||||
|
||||
|
||||
def _tool_calls_list(read: Callable[[str], object]) -> tuple[object, ...]:
|
||||
calls: Final = read("tool_calls")
|
||||
listed: Final = tuple(calls) if isinstance(calls, Sequence) and not isinstance(calls, str) else ()
|
||||
single: Final = read("function_call")
|
||||
return listed if listed else ((single,) if single is not None else ())
|
||||
|
||||
|
||||
def _tool_call_invocation(call: object) -> str:
|
||||
"""One tool call as `name(arguments)`. Custom tool calls name themselves and carry their
|
||||
arguments under `custom` rather than `function`."""
|
||||
read_call: Final = _field_reader(call)
|
||||
payload: Final = read_call("function") or read_call("custom") or call
|
||||
read_payload: Final = _field_reader(payload)
|
||||
name: Final = read_payload("name")
|
||||
arguments: Final = read_payload("arguments") or read_payload("input") or ""
|
||||
return f"{name or 'unnamed'}({arguments})"
|
||||
|
||||
|
||||
def _serialize_tool_calls(read: Callable[[str], object]) -> str:
|
||||
"""Every tool call in a reply as text a judge built for prose can still read."""
|
||||
return ", ".join(f"[tool call] {_tool_call_invocation(call)}" for call in _tool_calls_list(read))
|
||||
|
||||
|
||||
def _shadow_empty_reply_error(response_obj: object, routed_model: str) -> str:
|
||||
"""Why a shadow reply yielded no judgeable text at all: no prose, and no tool call to
|
||||
serialize either. The stable sentence comes first and every varying part after the
|
||||
semicolon, so grouping rows by error still yields one row per cause."""
|
||||
detail: Final = f"finish_reason={_chat_finish_reason(response_obj)}, model={routed_model or 'unknown'}"
|
||||
return f"shadow router returned an empty response; {detail}"
|
||||
|
||||
|
||||
def _responses_final_text(response_obj: object) -> str:
|
||||
"""The turn's aggregated output text, or empty when the turn carries tool calls. A
|
||||
dict-shaped payload is validated into the owner type first, because ``output_text``
|
||||
is a derived property rather than a serialized field, so it never exists on a dict;
|
||||
a dict the owner type rejects is unjudgeable and skipped."""
|
||||
"""The turn's judgeable text: the aggregated output plus any tool call serialized
|
||||
alongside it, the same way the chat surface renders one. A dict-shaped payload is
|
||||
validated into the owner type first, because ``output_text`` is a derived property
|
||||
rather than a serialized field, so it never exists on a dict; a dict the owner type
|
||||
rejects is unjudgeable and skipped."""
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
|
||||
try:
|
||||
|
|
@ -199,11 +263,16 @@ def _responses_final_text(response_obj: object) -> str:
|
|||
if not isinstance(output, Sequence):
|
||||
return ""
|
||||
items: Final = tuple(item.model_dump() if isinstance(item, BaseModel) else item for item in output)
|
||||
if any(
|
||||
not isinstance(item, Mapping) or item.get("type") in ("function_call", "custom_tool_call") for item in items
|
||||
):
|
||||
if any(not isinstance(item, Mapping) for item in items):
|
||||
return ""
|
||||
return str(getattr(response, "output_text", "") or "")
|
||||
calls: Final = tuple(
|
||||
item for item in items if isinstance(item, Mapping) and item.get("type") in _RESPONSES_TOOL_CALL_TYPES
|
||||
)
|
||||
prose: Final = str(getattr(response, "output_text", "") or "")
|
||||
if not calls:
|
||||
return prose
|
||||
serialized: Final = ", ".join(f"[tool call] {_tool_call_invocation(call)}" for call in calls)
|
||||
return f"{prose} {serialized}".strip() if prose else serialized
|
||||
|
||||
|
||||
class _SurfaceOps:
|
||||
|
|
@ -273,8 +342,8 @@ def _judgeable_sample(
|
|||
response_obj: object,
|
||||
) -> tuple[tuple[Mapping[str, object], ...], Mapping[str, object], str] | None:
|
||||
"""The normalized chat conversation, the forwardable generation params, and the
|
||||
judgeable final text; None when this request's shapes cannot be sampled (tool-final
|
||||
turn, empty text, or a shape the owner transformations reject)."""
|
||||
judgeable final text; None when this request's shapes cannot be sampled (no text and no
|
||||
tool call to serialize, or a shape the owner transformations reject)."""
|
||||
try:
|
||||
request: Final = ops.chat_request(kwargs, model_parameters)
|
||||
items: Final = _MESSAGE_ITEMS_ADAPTER.validate_python(request.get("messages"))
|
||||
|
|
@ -307,6 +376,11 @@ PAIRWISE_JUDGE_SYSTEM_PROMPT: Final = """You are an impartial quality judge comp
|
|||
|
||||
The responses are labeled A and B in random order. You do not know which system produced which.
|
||||
|
||||
A response may be prose, or a tool call shown as `[tool call] name(arguments)` if the
|
||||
assistant chose to act instead of answering directly. A tool call is not a defect: judge
|
||||
whether calling that tool was the right response to the conversation, the same as you
|
||||
would judge prose.
|
||||
|
||||
Criteria: correctness, completeness, clarity, conciseness.
|
||||
|
||||
Return ONLY valid JSON in this exact format, no other text:
|
||||
|
|
@ -376,14 +450,37 @@ def _unmask_preference(raw_preference: str, real_is_a: bool) -> str:
|
|||
return "tie"
|
||||
|
||||
|
||||
def _judge_user_prompt(conversation: str, response_a: str, response_b: str) -> str:
|
||||
_MAX_JUDGE_TOOL_DEFS_CHARS: Final = 2_000
|
||||
|
||||
|
||||
def _tool_definitions_text(tools: object) -> str:
|
||||
"""The tools available to both arms, name and description only: enough for the judge
|
||||
to tell whether the chosen tool, and not some other one, was the right call, without
|
||||
forwarding parameter schemas it does not need to score that."""
|
||||
if not isinstance(tools, Sequence) or isinstance(tools, str):
|
||||
return ""
|
||||
entries: Final = tuple(
|
||||
_field_reader(t)("function") or _field_reader(t)("custom") or t for t in tools if not isinstance(t, str)
|
||||
)
|
||||
lines: Final = tuple(
|
||||
f"- {_field_reader(e)('name') or 'unnamed'}: {_field_reader(e)('description') or 'no description'}"
|
||||
for e in entries
|
||||
)
|
||||
if not lines:
|
||||
return ""
|
||||
return ("Tools available to both responses:\n" + "\n".join(lines))[:_MAX_JUDGE_TOOL_DEFS_CHARS]
|
||||
|
||||
|
||||
def _judge_user_prompt(conversation: str, response_a: str, response_b: str, tool_definitions: str = "") -> str:
|
||||
"""The judge prompt under one total character budget: each response is capped, and
|
||||
the conversation tail gets whatever budget the responses left over."""
|
||||
the conversation tail gets whatever budget the responses and tool definitions left
|
||||
over."""
|
||||
a: Final = response_a[:_MAX_JUDGE_RESPONSE_CHARS]
|
||||
b: Final = response_b[:_MAX_JUDGE_RESPONSE_CHARS]
|
||||
conversation_budget: Final = _MAX_JUDGE_PROMPT_CHARS - len(a) - len(b)
|
||||
prefix: Final = f"{tool_definitions}\n\n" if tool_definitions else ""
|
||||
conversation_budget: Final = _MAX_JUDGE_PROMPT_CHARS - len(a) - len(b) - len(prefix)
|
||||
return (
|
||||
f"Conversation:\n{conversation[-conversation_budget:]}\n\n"
|
||||
f"{prefix}Conversation:\n{conversation[-conversation_budget:]}\n\n"
|
||||
f"Response A:\n{a}\n\n"
|
||||
f"Response B:\n{b}\n\n"
|
||||
"Which response is better?"
|
||||
|
|
@ -554,6 +651,7 @@ class ActiveShadowEvalJob(BaseModel):
|
|||
id: str
|
||||
router_name: str
|
||||
router_names: tuple[str, ...] = ()
|
||||
models: frozenset[str] = frozenset()
|
||||
direction: ShadowEvalDirection = "forward"
|
||||
baseline_model: str | None = None
|
||||
shadow_percentage: float
|
||||
|
|
@ -596,6 +694,21 @@ class ActiveShadowEvalJob(BaseModel):
|
|||
return self.baseline_model or arm_router
|
||||
|
||||
|
||||
def _canonical_group(router: "Router | None", model_group: str) -> str:
|
||||
"""A model group in the one spelling both a job's scope and a request's model compare
|
||||
under: an alias resolves to its target so the two never fail to match on spelling."""
|
||||
return (
|
||||
resolve_model_group_alias(router.model_group_alias, model_group) if router is not None else None
|
||||
) or model_group
|
||||
|
||||
|
||||
def _scope_admits(router: "Router | None", job: "ActiveShadowEvalJob", model_group: str) -> bool:
|
||||
"""Whether the request's group is in the job's model scope. Both sides resolve through
|
||||
the router's alias map at match time, so a re-pointed alias applies to the next request
|
||||
rather than after the jobs cache rolls."""
|
||||
return not job.models or any(_canonical_group(router, name) == model_group for name in job.models)
|
||||
|
||||
|
||||
def _as_active_job(record: object, attempts: int, spend: float) -> ActiveShadowEvalJob | None:
|
||||
"""The sampling path's view of one job row, or None for a row it cannot sample: an
|
||||
unknown direction, or a reverse job with no baseline model to duplicate against.
|
||||
|
|
@ -618,7 +731,8 @@ class ShadowEvalLogger(CustomLogger):
|
|||
A job targets a virtual key, a team, or a user; a request qualifies for a job when
|
||||
any of its resolved identities (key hash, team id, user id) matches the job's
|
||||
target, so team and user jobs cover JWT-authenticated traffic, which carries no
|
||||
key hash at all."""
|
||||
key hash at all. A job scoped to model groups further requires the request's
|
||||
requested group to be one of them."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -705,19 +819,24 @@ class ShadowEvalLogger(CustomLogger):
|
|||
active_jobs: Sequence[ActiveShadowEvalJob],
|
||||
request_metadata: Mapping[str, object],
|
||||
request_id: str,
|
||||
model_group: str,
|
||||
) -> tuple[ActiveShadowEvalJob, ...]:
|
||||
"""The jobs that sample this request. A key can hold one job per direction, and a
|
||||
request routed by one job's router while bypassing the other's qualifies for both;
|
||||
each is separately budgeted, so both fire. An admitting job that loses the sampling
|
||||
dice is counted, so results can weigh judged rows against the traffic they stand for."""
|
||||
dice is counted, so results can weigh judged rows against the traffic they stand for.
|
||||
A request outside a job's direction or model scope is not that job's traffic and
|
||||
goes uncounted, so the funnel stays a fraction of the traffic the job admits."""
|
||||
eligible: list[ActiveShadowEvalJob] = [] # mutable-ok: bucketed per-job admission
|
||||
now: Final = datetime.now(timezone.utc)
|
||||
router: Final = self._router_provider()
|
||||
for job in active_jobs:
|
||||
if (
|
||||
now >= job.ends_at
|
||||
or job.attempts + self._job_starts.get(job.id, 0) >= job.max_turns
|
||||
or (job.max_budget is not None and job.spend >= job.max_budget)
|
||||
or not _direction_admits(request_metadata, job)
|
||||
or not _scope_admits(router, job, model_group)
|
||||
):
|
||||
continue
|
||||
if not _sample_hits(request_id, job.id, job.shadow_percentage):
|
||||
|
|
@ -772,6 +891,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
tuple(job for target in targets for job in active_jobs.get(target, ())),
|
||||
request_metadata,
|
||||
request_id,
|
||||
_canonical_group(self._router_provider(), str(payload.get("model_group") or "")),
|
||||
)
|
||||
if not eligible:
|
||||
return
|
||||
|
|
@ -942,6 +1062,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
messages=messages,
|
||||
real_text=real_text,
|
||||
shadow_text=shadow.text,
|
||||
tools=shadow_params.get("tools"),
|
||||
parent_metadata=parent_metadata,
|
||||
)
|
||||
if isinstance(verdict, _CallFailure):
|
||||
|
|
@ -1080,15 +1201,18 @@ class ShadowEvalLogger(CustomLogger):
|
|||
classifier_cost=_decision_classifier_cost(shadow_metadata),
|
||||
)
|
||||
text: Final = _chat_final_text(response)
|
||||
routed_model: Final = str(
|
||||
getattr(response, "model", None) or _routing_decision(shadow_metadata).get("routed_model") or ""
|
||||
)
|
||||
if not text:
|
||||
return _CallFailure(
|
||||
"shadow router returned an empty response",
|
||||
_shadow_empty_reply_error(response, routed_model),
|
||||
cost=_call_cost(response),
|
||||
classifier_cost=_decision_classifier_cost(shadow_metadata),
|
||||
)
|
||||
return _ShadowResponse(
|
||||
text=text,
|
||||
model=str(getattr(response, "model", None) or _routing_decision(shadow_metadata).get("routed_model") or ""),
|
||||
model=routed_model,
|
||||
tier=_routed_tier(shadow_metadata),
|
||||
cost=_call_cost(response),
|
||||
classifier_cost=_decision_classifier_cost(shadow_metadata),
|
||||
|
|
@ -1100,9 +1224,12 @@ class ShadowEvalLogger(CustomLogger):
|
|||
messages: Sequence[Mapping[str, object]],
|
||||
real_text: str,
|
||||
shadow_text: str,
|
||||
tools: object,
|
||||
parent_metadata: Mapping[str, object],
|
||||
) -> "_JudgeVerdict | _CallFailure":
|
||||
"""Blind pairwise judge with A/B labels randomized to cancel position bias."""
|
||||
"""Blind pairwise judge with A/B labels randomized to cancel position bias. Both
|
||||
arms were offered the same tools, so the judge is shown their definitions too: a
|
||||
tool call is only assessable against what else was available to call instead."""
|
||||
real_is_a: Final = random.random() < 0.5
|
||||
response_a: Final = real_text if real_is_a else shadow_text
|
||||
response_b: Final = shadow_text if real_is_a else real_text
|
||||
|
|
@ -1117,7 +1244,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
{"role": "system", "content": PAIRWISE_JUDGE_SYSTEM_PROMPT}, # mutable-ok: SDK message
|
||||
{
|
||||
"role": "user",
|
||||
"content": _judge_user_prompt(conversation, response_a, response_b),
|
||||
"content": _judge_user_prompt(conversation, response_a, response_b, _tool_definitions_text(tools)),
|
||||
}, # mutable-ok: SDK message
|
||||
]
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -3821,7 +3821,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
def record_streamed_anthropic_message_id(self, message_id: str) -> None:
|
||||
self.streamed_anthropic_message_id = message_id
|
||||
|
||||
def _anthropic_messages_logged_response(self, result: Any) -> ModelResponse:
|
||||
def _anthropic_messages_logged_response(self, result: object) -> ModelResponse:
|
||||
"""
|
||||
The ModelResponse a /v1/messages spend_logs row is built from.
|
||||
|
||||
|
|
|
|||
|
|
@ -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``,
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
from collections.abc import Mapping, Sequence
|
||||
from collections.abc import Set as AbstractSet
|
||||
from typing import Any, Final
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -6,38 +7,45 @@ from pydantic import BaseModel
|
|||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH, DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER
|
||||
from litellm.litellm_core_utils.secret_redaction import REDACTED
|
||||
|
||||
_DEFAULT_SENSITIVE_PATTERNS: Final = frozenset(
|
||||
(
|
||||
"password",
|
||||
"secret",
|
||||
"key",
|
||||
"token",
|
||||
"auth",
|
||||
"authorization",
|
||||
"credential",
|
||||
# Plural form: Vertex uses ``vertex_credentials``; segment-exact
|
||||
# matching otherwise misses it because "credential" != "credentials".
|
||||
"credentials",
|
||||
"access",
|
||||
"private",
|
||||
"certificate",
|
||||
"fingerprint",
|
||||
"tenancy",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class SensitiveDataMasker:
|
||||
def __init__(
|
||||
self,
|
||||
sensitive_patterns: set[str] | None = None,
|
||||
non_sensitive_overrides: set[str] | None = None,
|
||||
sensitive_patterns: AbstractSet[str] | None = None,
|
||||
non_sensitive_overrides: AbstractSet[str] | None = None,
|
||||
visible_prefix: int = 4,
|
||||
visible_suffix: int = 4,
|
||||
mask_char: str = "*",
|
||||
mask_short_values: bool = True,
|
||||
extra_sensitive_patterns: AbstractSet[str] | None = None,
|
||||
):
|
||||
self.sensitive_patterns = sensitive_patterns or {
|
||||
"password",
|
||||
"secret",
|
||||
"key",
|
||||
"token",
|
||||
"auth",
|
||||
"authorization",
|
||||
"credential",
|
||||
# Plural form: Vertex uses ``vertex_credentials``; segment-exact
|
||||
# matching otherwise misses it because "credential" != "credentials".
|
||||
"credentials",
|
||||
"access",
|
||||
"private",
|
||||
"certificate",
|
||||
"fingerprint",
|
||||
"tenancy",
|
||||
}
|
||||
self.sensitive_patterns = (sensitive_patterns or _DEFAULT_SENSITIVE_PATTERNS) | (
|
||||
extra_sensitive_patterns or frozenset()
|
||||
)
|
||||
# If any key segment matches one of these, the key is not considered sensitive
|
||||
# even if it also matches a sensitive pattern. For example, "input_cost_per_token"
|
||||
# contains "token" but "cost" overrides that — it's a pricing field, not a secret.
|
||||
self.non_sensitive_overrides = non_sensitive_overrides or {"cost"}
|
||||
self.non_sensitive_overrides = non_sensitive_overrides or frozenset(("cost",))
|
||||
|
||||
self.visible_prefix = visible_prefix
|
||||
self.visible_suffix = visible_suffix
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -89,11 +89,24 @@ def _extract_converse_texts(
|
|||
top-level ``text`` blocks this scans the arbitrary-JSON fields a caller can
|
||||
hide prompt content in -- ``toolUse.input`` and
|
||||
``toolResult.content[].json`` (alongside ``toolResult.content[].text``) --
|
||||
as well as the request-level fields still forwarded to Bedrock that a caller
|
||||
can route blocked content through: ``toolConfig.tools`` (tool names,
|
||||
descriptions and input schemas) and ``additionalModelRequestFields``. Tool
|
||||
message blocks are skipped when tool messages are excluded, but tool
|
||||
definitions are always scanned to match the chat-completions guardrail path.
|
||||
as well as ``additionalModelRequestFields``, a free-form model-parameter bag
|
||||
with no schema that a caller can route blocked content through.
|
||||
|
||||
``toolConfig.tools`` is deliberately NOT scanned. Tool definitions are
|
||||
app-authored config, so their names, descriptions and JSON-schema strings
|
||||
("object", property names, titles, type names, enum values) would each reach
|
||||
the guardrail as a separate INPUT item, producing false positives and
|
||||
inflating guardrail usage for a request whose only prompt is one user
|
||||
message. No other guardrail translation handler puts tool definitions in
|
||||
``texts``; the chat and messages handlers carry them in the structured
|
||||
``tools`` input instead, which this handler does not populate because a
|
||||
Bedrock ``toolSpec`` is not the OpenAI tool shape those consumers expect.
|
||||
|
||||
``additionalModelRequestFields`` is treated differently on purpose. Bedrock
|
||||
gives ``toolConfig.tools`` a fixed schema whose contents are tool metadata by
|
||||
contract, while ``additionalModelRequestFields`` is free-form and defined by
|
||||
the target model, so what it carries cannot be classified without knowing
|
||||
that model. Scanning it stays the fail-closed default.
|
||||
"""
|
||||
holders: Final[list[_StringHolder]] = []
|
||||
|
||||
|
|
@ -121,10 +134,6 @@ def _extract_converse_texts(
|
|||
_collect_block_text(inner, holders)
|
||||
_collect_strings(inner.get("json"), holders)
|
||||
|
||||
tool_config: Final = body.get("toolConfig")
|
||||
if isinstance(tool_config, dict):
|
||||
_collect_strings(tool_config.get("tools"), holders)
|
||||
|
||||
_collect_strings(body.get("additionalModelRequestFields"), holders)
|
||||
|
||||
texts: Final = [container[key] for container, key in holders]
|
||||
|
|
|
|||
|
|
@ -272,11 +272,15 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig):
|
|||
)
|
||||
|
||||
# Only add tool_choice for models that explicitly support it
|
||||
if supports_tool_choice(model=model, custom_llm_provider="fireworks_ai"):
|
||||
if self._get_model_cost_capability_exact(
|
||||
model=model, capability="supports_tool_choice"
|
||||
) or supports_tool_choice(model=model, custom_llm_provider="fireworks_ai"):
|
||||
supported_params.append("tool_choice")
|
||||
|
||||
# Only add reasoning params for models that support it
|
||||
if supports_reasoning(model=model, custom_llm_provider="fireworks_ai"):
|
||||
if self._get_model_cost_capability_exact(model=model, capability="supports_reasoning") or supports_reasoning(
|
||||
model=model, custom_llm_provider="fireworks_ai"
|
||||
):
|
||||
supported_params.append("reasoning_effort")
|
||||
supported_params.append("reasoning_history")
|
||||
supported_params.append("thinking")
|
||||
|
|
|
|||
0
litellm/llms/mongodb/__init__.py
Normal file
0
litellm/llms/mongodb/__init__.py
Normal file
303
litellm/llms/mongodb/common_utils.py
Normal file
303
litellm/llms/mongodb/common_utils.py
Normal file
|
|
@ -0,0 +1,303 @@
|
|||
"""Shared helpers for the MongoDB integrations. pymongo lives in the optional ``mongodb`` extra,
|
||||
so every import of it is deferred to call time."""
|
||||
|
||||
import asyncio
|
||||
import threading
|
||||
import weakref
|
||||
from asyncio import AbstractEventLoop
|
||||
from collections import OrderedDict
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, TypeAlias, TypeVar
|
||||
|
||||
from litellm.exceptions import BadRequestError, ServiceUnavailableError, Timeout
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pymongo import AsyncMongoClient, MongoClient
|
||||
|
||||
PYMONGO_INSTALL_HINT: Final = (
|
||||
"The MongoDB vector store requires the 'pymongo' package. "
|
||||
"Run 'pip install litellm[mongodb]' (or 'pip install pymongo') to install it."
|
||||
)
|
||||
|
||||
MONGODB_PROVIDER: Final = "mongodb"
|
||||
|
||||
|
||||
def config_error(message: str) -> BadRequestError:
|
||||
"""400 rather than the 500 a bare ValueError becomes once litellm.exception_type wraps it."""
|
||||
return BadRequestError(message=message, model=None, llm_provider=MONGODB_PROVIDER)
|
||||
|
||||
|
||||
def timeout_error(message: str) -> Timeout:
|
||||
return Timeout(message=message, model=None, llm_provider=MONGODB_PROVIDER)
|
||||
|
||||
|
||||
def unavailable_error(message: str) -> ServiceUnavailableError:
|
||||
"""litellm only retries 408, 409, 429 and 5xx, so a 400 here would make a failover permanent."""
|
||||
return ServiceUnavailableError(message=message, model=None, llm_provider=MONGODB_PROVIDER)
|
||||
|
||||
|
||||
DEFAULT_CONNECT_TIMEOUT_MS: Final = 10_000
|
||||
DEFAULT_SOCKET_TIMEOUT_MS: Final = 30_000
|
||||
DEFAULT_SERVER_SELECTION_TIMEOUT_MS: Final = 10_000
|
||||
|
||||
_MAX_CACHED_CLIENTS: Final = 32
|
||||
|
||||
_APP_NAME: Final = "litellm"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MongoClientKey:
|
||||
connection_string: str
|
||||
connect_timeout_ms: int
|
||||
socket_timeout_ms: int
|
||||
server_selection_timeout_ms: int
|
||||
|
||||
|
||||
SyncClientFactory: TypeAlias = Callable[..., "MongoClient"]
|
||||
AsyncClientFactory: TypeAlias = Callable[..., "AsyncMongoClient"]
|
||||
|
||||
_K = TypeVar("_K")
|
||||
_V = TypeVar("_V")
|
||||
|
||||
_AsyncClientCacheKey: TypeAlias = tuple[MongoClientKey, int]
|
||||
# CPython recycles id() aggressively, so the id alone would hand a new loop a closed loop's client
|
||||
_AsyncClientEntry: TypeAlias = tuple["weakref.ref[AbstractEventLoop]", "AsyncMongoClient"]
|
||||
|
||||
_SyncClientCache: TypeAlias = "OrderedDict[MongoClientKey, MongoClient]"
|
||||
_AsyncClientCache: TypeAlias = "OrderedDict[_AsyncClientCacheKey, _AsyncClientEntry]"
|
||||
|
||||
_sync_clients: Final[_SyncClientCache] = OrderedDict() # mutable-ok: process-level client cache
|
||||
_async_clients: Final[_AsyncClientCache] = OrderedDict() # mutable-ok: same cache, per loop
|
||||
# async searches reach the sync client through executor threads, so both caches are shared state
|
||||
_cache_lock: Final = threading.Lock()
|
||||
|
||||
|
||||
def _store_bounded(cache: "OrderedDict[_K, _V]", cache_key: "_K", value: "_V") -> None:
|
||||
"""Eviction only drops this cache's reference; an in-flight search keeps its client alive."""
|
||||
with _cache_lock:
|
||||
cache[cache_key] = value # mutable-ok: an LRU cache is mutable state by definition
|
||||
cache.move_to_end(cache_key)
|
||||
while len(cache) > _MAX_CACHED_CLIENTS:
|
||||
cache.popitem(last=False)
|
||||
|
||||
|
||||
def _mark_used(cache: "OrderedDict[_K, _V]", cache_key: "_K") -> None:
|
||||
with _cache_lock:
|
||||
if cache_key in cache:
|
||||
cache.move_to_end(cache_key)
|
||||
|
||||
|
||||
def import_sync_mongo_client() -> "type[MongoClient]":
|
||||
try:
|
||||
from pymongo import MongoClient as SyncMongoClient
|
||||
except ImportError as e:
|
||||
raise config_error(PYMONGO_INSTALL_HINT) from e
|
||||
return SyncMongoClient
|
||||
|
||||
|
||||
def import_async_mongo_client() -> "type[AsyncMongoClient]":
|
||||
try:
|
||||
from pymongo import AsyncMongoClient as AsyncMongoClientClass
|
||||
except ImportError as e:
|
||||
raise config_error(PYMONGO_INSTALL_HINT) from e
|
||||
return AsyncMongoClientClass
|
||||
|
||||
|
||||
def _client_kwargs(key: MongoClientKey) -> Mapping[str, object]:
|
||||
return MappingProxyType(
|
||||
{
|
||||
"connectTimeoutMS": key.connect_timeout_ms,
|
||||
"socketTimeoutMS": key.socket_timeout_ms,
|
||||
"serverSelectionTimeoutMS": key.server_selection_timeout_ms,
|
||||
"appname": _APP_NAME,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def get_sync_client(key: MongoClientKey, client_class: SyncClientFactory | None = None) -> "MongoClient":
|
||||
cached: Final = _sync_clients.get(key)
|
||||
if cached is not None:
|
||||
_mark_used(_sync_clients, key)
|
||||
return cached
|
||||
build: Final = client_class if client_class is not None else import_sync_mongo_client()
|
||||
client: Final = build(key.connection_string, **_client_kwargs(key))
|
||||
_store_bounded(_sync_clients, key, client)
|
||||
return client
|
||||
|
||||
|
||||
def _purge_dead_loops() -> None:
|
||||
"""A cached client holds its loop alive, so a closed loop's entry would pin that client and its
|
||||
sockets for the life of the process."""
|
||||
with _cache_lock:
|
||||
for stale in tuple(
|
||||
cache_key
|
||||
for cache_key, (loop_ref, _) in _async_clients.items()
|
||||
if (cached_loop := loop_ref()) is None or cached_loop.is_closed()
|
||||
):
|
||||
del _async_clients[stale]
|
||||
|
||||
|
||||
def get_async_client(key: MongoClientKey, client_class: AsyncClientFactory | None = None) -> "AsyncMongoClient":
|
||||
"""Async clients bind to the loop that created them, so the cache is keyed per loop."""
|
||||
loop: Final = asyncio.get_running_loop()
|
||||
loop_key: Final = (key, id(loop))
|
||||
cached: Final = _async_clients.get(loop_key)
|
||||
if cached is not None and cached[0]() is loop:
|
||||
_mark_used(_async_clients, loop_key)
|
||||
return cached[1]
|
||||
_purge_dead_loops()
|
||||
build: Final = client_class if client_class is not None else import_async_mongo_client()
|
||||
client: Final = build(key.connection_string, **_client_kwargs(key))
|
||||
_store_bounded(_async_clients, loop_key, (weakref.ref(loop), client))
|
||||
return client
|
||||
|
||||
|
||||
def reset_client_cache() -> None:
|
||||
with _cache_lock:
|
||||
_sync_clients.clear()
|
||||
_async_clients.clear()
|
||||
|
||||
|
||||
_AUTHENTICATION_FAILED_CODE: Final = 18
|
||||
_UNAUTHORIZED_CODE: Final = 13
|
||||
# Atlas reports a rejected user as code 8000 "AtlasError" where a self-managed mongod reports 18
|
||||
_AUTHENTICATION_MESSAGE_MARKERS: Final = ("bad auth", "authentication failed", "not authorized")
|
||||
_RESOLUTION_TIMEOUT_MARKERS: Final = ("resolution lifetime expired", "dns operation timed out")
|
||||
_UNKNOWN_HOSTNAME_MARKERS: Final = ("dns query name does not exist", "name or service not known")
|
||||
_CREDENTIAL_ESCAPING_MARKERS: Final = ("must be escaped according to rfc 3986", "bad database name")
|
||||
|
||||
|
||||
def _index_hint(index_name: str, database: str, collection: str) -> str:
|
||||
return (
|
||||
f"No queryable MongoDB Vector Search index named '{index_name}' was found on "
|
||||
f"'{database}.{collection}'. Confirm the index exists on that exact collection, that its "
|
||||
"status is READY rather than still building, and that the vector store id matches the index name."
|
||||
)
|
||||
|
||||
|
||||
def missing_index_error(index_name: str, database: str, collection: str) -> BadRequestError:
|
||||
"""$vectorSearch against a missing index, database or collection returns zero documents rather
|
||||
than failing, so an empty result set is checked against the catalogue and reported as this."""
|
||||
return config_error(
|
||||
f"{_index_hint(index_name, database, collection)} A vector search against a database, "
|
||||
"collection or index that does not exist returns no results rather than an error, so this "
|
||||
"was reported as an empty result set by MongoDB."
|
||||
)
|
||||
|
||||
|
||||
def index_not_ready_error(index_name: str, database: str, collection: str, status: str) -> BadRequestError:
|
||||
return config_error(
|
||||
f"The MongoDB Vector Search index '{index_name}' on '{database}.{collection}' is not queryable "
|
||||
f"yet; its status is {status}. Searches against it return no results until the build finishes."
|
||||
)
|
||||
|
||||
|
||||
def translate_mongo_error(error: Exception, index_name: str, database: str, collection: str) -> Exception:
|
||||
"""Returns the exception to raise, so callers keep the driver error as ``__cause__``."""
|
||||
try:
|
||||
from pymongo.errors import (
|
||||
ConfigurationError,
|
||||
ConnectionFailure,
|
||||
ExecutionTimeout,
|
||||
InvalidOperation,
|
||||
NetworkTimeout,
|
||||
OperationFailure,
|
||||
ServerSelectionTimeoutError,
|
||||
)
|
||||
except ImportError:
|
||||
return error
|
||||
|
||||
if isinstance(error, ServerSelectionTimeoutError):
|
||||
return timeout_error(
|
||||
"Could not reach the MongoDB deployment before the timeout. On Atlas this is usually the "
|
||||
"project's IP access list not containing this host, or a paused cluster. On a self-managed "
|
||||
"deployment it is usually the host or port in the URI, or a firewall between this process "
|
||||
f"and mongod. Either way it can also be an unresolvable hostname. Driver detail: {error}"
|
||||
)
|
||||
# ExecutionTimeout subclasses OperationFailure, so it has to be matched before it
|
||||
if isinstance(error, (NetworkTimeout, ExecutionTimeout)):
|
||||
return timeout_error(
|
||||
f"The MongoDB vector search against '{database}.{collection}' timed out before returning. "
|
||||
f"Driver detail: {error}"
|
||||
)
|
||||
# ServerSelectionTimeoutError and NetworkTimeout also subclass ConnectionFailure, so this only
|
||||
# sees what those branches left
|
||||
if isinstance(error, ConnectionFailure):
|
||||
return unavailable_error(
|
||||
f"The connection to '{database}.{collection}' was dropped or refused. That is usually a "
|
||||
"replica set failover or a restarted node, so the search is worth retrying. If it keeps "
|
||||
"happening: on Atlas the usual cause is a connection string with no username and password, "
|
||||
"or a TLS failure, so confirm the URI is the one Atlas shows under Connect, Drivers; on a "
|
||||
"self-managed deployment, check that mongod is listening on the host and port in the URI. "
|
||||
f"Driver detail: {error}"
|
||||
)
|
||||
if isinstance(error, OperationFailure):
|
||||
code: Final = error.code
|
||||
detail: Final = str(error).lower()
|
||||
if code in (_AUTHENTICATION_FAILED_CODE, _UNAUTHORIZED_CODE) or any(
|
||||
marker in detail for marker in _AUTHENTICATION_MESSAGE_MARKERS
|
||||
):
|
||||
return config_error(
|
||||
"MongoDB rejected the credentials in mongodb_connection_string, or the database user "
|
||||
f"lacks read access to '{database}.{collection}'. Driver detail: {error.details}"
|
||||
)
|
||||
if "dimension" in detail:
|
||||
return config_error(
|
||||
"The query embedding does not match the vector dimensions the index was built for. "
|
||||
"litellm_embedding_model must be the same model that produced the stored vectors. "
|
||||
f"Driver detail: {error}"
|
||||
)
|
||||
if "is not indexed as vector" in detail:
|
||||
return config_error(
|
||||
"mongodb_embedding_field names a field the MongoDB Vector Search index does not cover. "
|
||||
f"It must match the 'path' the index '{index_name}' was created on. Driver detail: {error}"
|
||||
)
|
||||
if "index" in detail and ("not found" in detail or "does not exist" in detail or "unknown" in detail):
|
||||
return config_error(f"{_index_hint(index_name, database, collection)} Driver detail: {error}")
|
||||
return config_error(
|
||||
f"MongoDB rejected the vector search against '{database}.{collection}' using index "
|
||||
f"'{index_name}'. Driver detail: {error}"
|
||||
)
|
||||
if isinstance(error, ConfigurationError):
|
||||
configuration_detail: Final = str(error).lower()
|
||||
if any(marker in configuration_detail for marker in _RESOLUTION_TIMEOUT_MARKERS):
|
||||
return timeout_error(
|
||||
"The DNS lookup for the cluster in mongodb_connection_string did not finish in time. "
|
||||
"A mongodb+srv:// URI needs an SRV lookup before any connection is attempted, so this "
|
||||
f"is DNS or the configured timeout, not MongoDB. Driver detail: {error}"
|
||||
)
|
||||
if any(marker in configuration_detail for marker in _UNKNOWN_HOSTNAME_MARKERS):
|
||||
return config_error(
|
||||
"The hostname in mongodb_connection_string does not exist in DNS. On Atlas, check the "
|
||||
"cluster name against the URI shown under Connect, Drivers. On a self-managed deployment, "
|
||||
f"check that the hostname resolves from this process. Driver detail: {error}"
|
||||
)
|
||||
if any(marker in configuration_detail for marker in _CREDENTIAL_ESCAPING_MARKERS):
|
||||
return config_error(
|
||||
"mongodb_connection_string could not be parsed. A username or password containing "
|
||||
"'@', '/', ':' or '%' has to be percent-encoded per RFC 3986, so 'p@ss/word' becomes "
|
||||
"'p%40ss%2Fword'. If the credentials are already encoded, check the database name in "
|
||||
f"the URI path instead. Driver detail: {error}"
|
||||
)
|
||||
return config_error(
|
||||
f"mongodb_connection_string is not a usable MongoDB connection string. Driver detail: {error}"
|
||||
)
|
||||
if isinstance(error, InvalidOperation):
|
||||
return config_error(f"The MongoDB client was already closed or is unusable. Driver detail: {error}")
|
||||
# An unreadable tlsCAFile or tlsCertificateKeyFile raises OSError, not a PyMongoError
|
||||
if isinstance(error, OSError) and error.filename:
|
||||
return config_error(
|
||||
f"'{error.filename}', named by a TLS option in mongodb_connection_string, could not be read. "
|
||||
"Check that tlsCAFile and tlsCertificateKeyFile point at files this process can open; inside "
|
||||
f"a container that is the path in the container, not on the host. Driver detail: {error}"
|
||||
)
|
||||
# pymongo raises a plain ValueError, not a PyMongoError, for an unusable port
|
||||
if isinstance(error, ValueError):
|
||||
return config_error(
|
||||
"The host and port in mongodb_connection_string could not be parsed. If the port is a "
|
||||
"number between 0 and 65535, the cause is usually an unescaped ':' in the password, which "
|
||||
f"has to be percent-encoded per RFC 3986 as '%3A'. Driver detail: {error}"
|
||||
)
|
||||
return error
|
||||
0
litellm/llms/mongodb/vector_stores/__init__.py
Normal file
0
litellm/llms/mongodb/vector_stores/__init__.py
Normal file
431
litellm/llms/mongodb/vector_stores/transformation.py
Normal file
431
litellm/llms/mongodb/vector_stores/transformation.py
Normal file
|
|
@ -0,0 +1,431 @@
|
|||
"""MongoDB Vector Search has no HTTP query API, so this is a direct provider that runs the
|
||||
``$vectorSearch`` aggregation through pymongo. ``vector_store_id`` is the search index name."""
|
||||
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, NoReturn
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
BaseDirectVectorStoreConfig,
|
||||
LiteLLMVectorStoreEmbeddingExecutor,
|
||||
VectorStoreEmbeddingExecutor,
|
||||
)
|
||||
from litellm.llms.mongodb.common_utils import (
|
||||
DEFAULT_CONNECT_TIMEOUT_MS,
|
||||
DEFAULT_SERVER_SELECTION_TIMEOUT_MS,
|
||||
DEFAULT_SOCKET_TIMEOUT_MS,
|
||||
MongoClientKey,
|
||||
config_error,
|
||||
get_async_client,
|
||||
get_sync_client,
|
||||
index_not_ready_error,
|
||||
missing_index_error,
|
||||
translate_mongo_error,
|
||||
)
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
from litellm.types.vector_stores import (
|
||||
VectorStoreCreateOptionalRequestParams,
|
||||
VectorStoreResultContent,
|
||||
VectorStoreSearchOptionalRequestParams,
|
||||
VectorStoreSearchResponse,
|
||||
VectorStoreSearchResult,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
DEFAULT_EMBEDDING_FIELD_NAME: Final = "embedding"
|
||||
DEFAULT_TEXT_FIELD_NAME: Final = "text"
|
||||
SCORE_FIELD_NAME: Final = "score"
|
||||
|
||||
DEFAULT_MAX_NUM_RESULTS: Final = 10
|
||||
MIN_MAX_NUM_RESULTS: Final = 1
|
||||
MAX_MAX_NUM_RESULTS: Final = 50
|
||||
|
||||
NUM_CANDIDATES_MULTIPLIER: Final = 10
|
||||
MIN_NUM_CANDIDATES: Final = 100
|
||||
MAX_NUM_CANDIDATES: Final = 10_000
|
||||
|
||||
MAX_QUERY_CHARACTERS: Final = 32_000
|
||||
|
||||
_EMPTY_EMBEDDING_CONFIG: Final = MappingProxyType({})
|
||||
|
||||
_SEARCH_ONLY_MESSAGE: Final = (
|
||||
"MongoDB vector store is search-only. Create the collection and its MongoDB Vector Search "
|
||||
"index in MongoDB directly, then register it here by index name."
|
||||
)
|
||||
|
||||
|
||||
class _MongoDBSearchParams(BaseModel):
|
||||
"""Typed view over the vector store's litellm_params; unrelated keys are ignored."""
|
||||
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
litellm_embedding_model: str | None = None
|
||||
litellm_embedding_config: Mapping[str, object] | None = None
|
||||
mongodb_connection_string: str | None = None
|
||||
mongodb_database: str | None = None
|
||||
mongodb_collection: str | None = None
|
||||
mongodb_text_field: str | None = None
|
||||
mongodb_embedding_field: str | None = None
|
||||
mongodb_num_candidates: int | None = None
|
||||
|
||||
@property
|
||||
def text_field(self) -> str:
|
||||
return self.mongodb_text_field or DEFAULT_TEXT_FIELD_NAME
|
||||
|
||||
@property
|
||||
def embedding_field(self) -> str:
|
||||
return self.mongodb_embedding_field or DEFAULT_EMBEDDING_FIELD_NAME
|
||||
|
||||
def require_embedding_model(self) -> str:
|
||||
if not self.litellm_embedding_model:
|
||||
raise config_error(
|
||||
"litellm_embedding_model is required in litellm_params for the MongoDB vector store. "
|
||||
"It must be the same model that produced the vectors stored in "
|
||||
f"'{self.mongodb_collection or '<collection>'}.{self.embedding_field}', or search results "
|
||||
"will be meaningless. Example: litellm_embedding_model: openai/text-embedding-3-small"
|
||||
)
|
||||
return self.litellm_embedding_model
|
||||
|
||||
def require_connection_string(self) -> str:
|
||||
if not self.mongodb_connection_string:
|
||||
raise config_error(
|
||||
"mongodb_connection_string is required in litellm_params for the MongoDB vector store. "
|
||||
"Example: mongodb+srv://<user>:<password>@<cluster>.mongodb.net for Atlas, or "
|
||||
"mongodb://<user>:<password>@<host>:27017 for a self-managed deployment"
|
||||
)
|
||||
scheme: Final = self.mongodb_connection_string.split("://", 1)[0].lower()
|
||||
if scheme not in ("mongodb", "mongodb+srv"):
|
||||
raise config_error(
|
||||
"mongodb_connection_string must start with 'mongodb://' or 'mongodb+srv://', "
|
||||
f"got '{self.mongodb_connection_string.split('://', 1)[0]}://'"
|
||||
)
|
||||
return self.mongodb_connection_string
|
||||
|
||||
def require_database(self) -> str:
|
||||
if not self.mongodb_database:
|
||||
raise config_error(
|
||||
"mongodb_database is required in litellm_params for the MongoDB vector store. "
|
||||
"Example: mongodb_database: sample_mflix"
|
||||
)
|
||||
return self.mongodb_database
|
||||
|
||||
def require_collection(self) -> str:
|
||||
if not self.mongodb_collection:
|
||||
raise config_error(
|
||||
"mongodb_collection is required in litellm_params for the MongoDB vector store. "
|
||||
"Example: mongodb_collection: embedded_movies"
|
||||
)
|
||||
return self.mongodb_collection
|
||||
|
||||
|
||||
_MONGODB_PARAM_PREFIX: Final = "mongodb_"
|
||||
_KNOWN_MONGODB_PARAMS: Final = frozenset(
|
||||
name for name in _MongoDBSearchParams.model_fields if name.startswith(_MONGODB_PARAM_PREFIX)
|
||||
)
|
||||
|
||||
|
||||
class MongoDBVectorStoreConfig(BaseDirectVectorStoreConfig):
|
||||
def __init__(
|
||||
self,
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
sync_client_factory: Callable[[MongoClientKey], object] | None = None,
|
||||
async_client_factory: Callable[[MongoClientKey], object] | None = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.embedding_executor: Final[VectorStoreEmbeddingExecutor] = (
|
||||
embedding_executor if embedding_executor is not None else LiteLLMVectorStoreEmbeddingExecutor()
|
||||
)
|
||||
self.sync_client_factory: Final[Callable[[MongoClientKey], object]] = (
|
||||
sync_client_factory if sync_client_factory is not None else get_sync_client
|
||||
)
|
||||
self.async_client_factory: Final[Callable[[MongoClientKey], object]] = (
|
||||
async_client_factory if async_client_factory is not None else get_async_client
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _reject_unknown_params(litellm_params: Mapping[str, object]) -> None:
|
||||
"""Without this a mistyped mongodb_collection reads as 'mongodb_collection is required',
|
||||
naming a key the reader can see they have set."""
|
||||
unknown: Final = sorted(
|
||||
key for key in litellm_params if key.startswith(_MONGODB_PARAM_PREFIX) and key not in _KNOWN_MONGODB_PARAMS
|
||||
)
|
||||
if unknown:
|
||||
raise config_error(
|
||||
f"Unrecognised MongoDB vector store parameter(s): {', '.join(unknown)}. "
|
||||
f"Supported: {', '.join(sorted(_KNOWN_MONGODB_PARAMS))}."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _query_text(query: str | Sequence[str]) -> str:
|
||||
text: Final = query if isinstance(query, str) else " ".join(query)
|
||||
if not text.strip():
|
||||
raise config_error("query must not be empty")
|
||||
if len(text) > MAX_QUERY_CHARACTERS:
|
||||
raise config_error(f"query must be at most {MAX_QUERY_CHARACTERS} characters, got {len(text)}")
|
||||
return text
|
||||
|
||||
@staticmethod
|
||||
def _limit(vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams) -> int:
|
||||
requested: Final = vector_store_search_optional_params.get("max_num_results")
|
||||
if requested is None:
|
||||
return DEFAULT_MAX_NUM_RESULTS
|
||||
if not MIN_MAX_NUM_RESULTS <= requested <= MAX_MAX_NUM_RESULTS:
|
||||
raise config_error(
|
||||
f"max_num_results must be between {MIN_MAX_NUM_RESULTS} and {MAX_MAX_NUM_RESULTS}, got {requested}"
|
||||
)
|
||||
return requested
|
||||
|
||||
@staticmethod
|
||||
def _num_candidates(limit: int, configured: int | None) -> int:
|
||||
if configured is not None:
|
||||
if not limit <= configured <= MAX_NUM_CANDIDATES:
|
||||
raise config_error(
|
||||
f"mongodb_num_candidates must be between max_num_results ({limit}) and "
|
||||
f"{MAX_NUM_CANDIDATES}, got {configured}"
|
||||
)
|
||||
return configured
|
||||
return min(max(limit * NUM_CANDIDATES_MULTIPLIER, MIN_NUM_CANDIDATES), MAX_NUM_CANDIDATES)
|
||||
|
||||
@staticmethod
|
||||
def _timeout_ms(timeout: float | httpx.Timeout | None) -> tuple[int, int]:
|
||||
"""The connect and socket budgets pymongo is built with, in that order."""
|
||||
if isinstance(timeout, httpx.Timeout):
|
||||
return (
|
||||
int((timeout.connect or DEFAULT_CONNECT_TIMEOUT_MS / 1000) * 1000),
|
||||
int((timeout.read or DEFAULT_SOCKET_TIMEOUT_MS / 1000) * 1000),
|
||||
)
|
||||
if timeout is None:
|
||||
return DEFAULT_CONNECT_TIMEOUT_MS, DEFAULT_SOCKET_TIMEOUT_MS
|
||||
return min(int(float(timeout) * 1000), DEFAULT_CONNECT_TIMEOUT_MS), int(float(timeout) * 1000)
|
||||
|
||||
@classmethod
|
||||
def _client_key(cls, params: _MongoDBSearchParams, timeout: float | httpx.Timeout | None) -> MongoClientKey:
|
||||
connect_ms, socket_ms = cls._timeout_ms(timeout)
|
||||
return MongoClientKey(
|
||||
connection_string=params.require_connection_string(),
|
||||
connect_timeout_ms=connect_ms,
|
||||
socket_timeout_ms=socket_ms,
|
||||
server_selection_timeout_ms=min(connect_ms, DEFAULT_SERVER_SELECTION_TIMEOUT_MS),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _pipeline(
|
||||
cls,
|
||||
vector_store_id: str,
|
||||
query_vector: Sequence[float],
|
||||
params: _MongoDBSearchParams,
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
) -> Sequence[Mapping[str, object]]:
|
||||
if vector_store_search_optional_params.get("filters") is not None:
|
||||
raise config_error(
|
||||
"MongoDB vector store does not support the filters parameter yet. "
|
||||
"Restrict the collection or the MongoDB Vector Search index definition instead."
|
||||
)
|
||||
if vector_store_search_optional_params.get("ranking_options") is not None:
|
||||
raise config_error(
|
||||
"MongoDB vector store does not support the ranking_options parameter yet. "
|
||||
"Every result already carries the vectorSearchScore, so filter or re-rank "
|
||||
"on that rather than having the threshold silently ignored."
|
||||
)
|
||||
if vector_store_search_optional_params.get("rewrite_query") is not None:
|
||||
raise config_error(
|
||||
"MongoDB vector store does not support the rewrite_query parameter. The query is "
|
||||
"embedded exactly as sent; rewrite it before calling if you need that."
|
||||
)
|
||||
limit: Final = cls._limit(vector_store_search_optional_params)
|
||||
search: Final = MappingProxyType(
|
||||
{
|
||||
"index": vector_store_id,
|
||||
"path": params.embedding_field,
|
||||
"queryVector": tuple(query_vector),
|
||||
"numCandidates": cls._num_candidates(limit, params.mongodb_num_candidates),
|
||||
"limit": limit,
|
||||
}
|
||||
)
|
||||
projection: Final = MappingProxyType(
|
||||
{params.text_field: 1, SCORE_FIELD_NAME: MappingProxyType({"$meta": "vectorSearchScore"})}
|
||||
)
|
||||
return [ # mutable-ok: pymongo rejects any non-list pipeline in common.validate_list
|
||||
MappingProxyType({"$vectorSearch": search}),
|
||||
MappingProxyType({"$project": projection}),
|
||||
]
|
||||
|
||||
@classmethod
|
||||
def _field_value(cls, document: Mapping[str, object], dotted_path: str) -> str | None:
|
||||
"""None means absent, which is what separates a mistyped field from genuinely empty text."""
|
||||
head, _, rest = dotted_path.partition(".")
|
||||
if head not in document:
|
||||
return None
|
||||
value: Final = document[head]
|
||||
if not rest:
|
||||
return None if value is None else str(value)
|
||||
return cls._field_value(value, rest) if isinstance(value, Mapping) else None
|
||||
|
||||
@classmethod
|
||||
def _to_result(cls, document: Mapping[str, object], text_field: str) -> VectorStoreSearchResult:
|
||||
document_id: Final = document.get("_id")
|
||||
identifier: Final = None if document_id is None else str(document_id)
|
||||
content: Final = [ # mutable-ok: VectorStoreSearchResult declares a list of content parts
|
||||
VectorStoreResultContent(text=cls._field_value(document, text_field) or "", type="text")
|
||||
]
|
||||
raw_score: Final = document.get(SCORE_FIELD_NAME)
|
||||
return VectorStoreSearchResult(
|
||||
score=float(raw_score) if isinstance(raw_score, (int, float)) else None,
|
||||
content=content,
|
||||
file_id=identifier,
|
||||
filename=identifier,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _raise_for_missing_text_field(
|
||||
cls, documents: Sequence[Mapping[str, object]], text_field: str, database: str, collection: str
|
||||
) -> None:
|
||||
"""$vectorSearch matches documents carrying no text, so a mistyped mongodb_text_field
|
||||
returns well-scored results with empty content instead of failing."""
|
||||
if documents and all(cls._field_value(document, text_field) is None for document in documents):
|
||||
raise config_error(
|
||||
f"None of the {len(documents)} matched documents in '{database}.{collection}' has a "
|
||||
f"'{text_field}' field, so every result would carry empty text. Set mongodb_text_field "
|
||||
"to the field holding the readable text; it accepts a dotted path such as metadata.body."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _to_response(
|
||||
cls, documents: Sequence[Mapping[str, object]], query_text: str, text_field: str
|
||||
) -> VectorStoreSearchResponse:
|
||||
return VectorStoreSearchResponse(
|
||||
object="vector_store.search_results.page",
|
||||
search_query=query_text,
|
||||
data=[ # mutable-ok: VectorStoreSearchResponse declares data as a list
|
||||
cls._to_result(document, text_field) for document in documents
|
||||
],
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _raise_for_unusable_index(
|
||||
catalogue: Sequence[Mapping[str, object]], index_name: str, database: str, collection: str
|
||||
) -> None:
|
||||
"""mongod returns zero documents both for a query that matched nothing and for a missing
|
||||
database, collection or index, so the catalogue decides which one happened."""
|
||||
if not catalogue:
|
||||
raise missing_index_error(index_name, database, collection)
|
||||
entry: Final = catalogue[0]
|
||||
if not entry.get("queryable"):
|
||||
raise index_not_ready_error(index_name, database, collection, str(entry.get("status") or "unknown"))
|
||||
|
||||
@staticmethod
|
||||
def _embedding_vector(embedding_response: EmbeddingResponse) -> Sequence[float]:
|
||||
data: Final = embedding_response.data
|
||||
if not data:
|
||||
raise config_error(
|
||||
"The embedding model returned no embedding for the search query, so there is nothing "
|
||||
"to search MongoDB with. Check the embedding deployment named by litellm_embedding_model."
|
||||
)
|
||||
return data[0]["embedding"]
|
||||
|
||||
def execute_search_vector_store_request(
|
||||
self,
|
||||
vector_store_id: str,
|
||||
query: str | Sequence[str],
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj",
|
||||
litellm_params: Mapping[str, object],
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
) -> VectorStoreSearchResponse:
|
||||
self._reject_unknown_params(litellm_params)
|
||||
params: Final = _MongoDBSearchParams.model_validate(litellm_params)
|
||||
query_text: Final = self._query_text(query)
|
||||
key: Final = self._client_key(params, timeout)
|
||||
database: Final = params.require_database()
|
||||
collection: Final = params.require_collection()
|
||||
|
||||
embedding_response: Final = (embedding_executor or self.embedding_executor).embed(
|
||||
params.require_embedding_model(),
|
||||
query_text,
|
||||
params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG,
|
||||
)
|
||||
pipeline: Final = self._pipeline(
|
||||
vector_store_id, self._embedding_vector(embedding_response), params, vector_store_search_optional_params
|
||||
)
|
||||
|
||||
try:
|
||||
client: Final = self.sync_client_factory(key)
|
||||
target: Final = client[database][collection] # pyright: ignore[reportIndexIssue] # factory is typed as returning object so injected doubles are accepted
|
||||
documents: Final = tuple(target.aggregate(pipeline))
|
||||
except Exception as e:
|
||||
raise translate_mongo_error(e, index_name=vector_store_id, database=database, collection=collection) from e
|
||||
if not documents:
|
||||
try:
|
||||
catalogue: Final = tuple(target.list_search_indexes(vector_store_id))
|
||||
except Exception as e:
|
||||
raise translate_mongo_error(
|
||||
e, index_name=vector_store_id, database=database, collection=collection
|
||||
) from e
|
||||
self._raise_for_unusable_index(catalogue, vector_store_id, database, collection)
|
||||
self._raise_for_missing_text_field(documents, params.text_field, database, collection)
|
||||
return self._to_response(documents, query_text, params.text_field)
|
||||
|
||||
async def aexecute_search_vector_store_request(
|
||||
self,
|
||||
vector_store_id: str,
|
||||
query: str | Sequence[str],
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj",
|
||||
litellm_params: Mapping[str, object],
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
) -> VectorStoreSearchResponse:
|
||||
self._reject_unknown_params(litellm_params)
|
||||
params: Final = _MongoDBSearchParams.model_validate(litellm_params)
|
||||
query_text: Final = self._query_text(query)
|
||||
key: Final = self._client_key(params, timeout)
|
||||
database: Final = params.require_database()
|
||||
collection: Final = params.require_collection()
|
||||
|
||||
embedding_response: Final = await (embedding_executor or self.embedding_executor).aembed(
|
||||
params.require_embedding_model(),
|
||||
query_text,
|
||||
params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG,
|
||||
)
|
||||
pipeline: Final = self._pipeline(
|
||||
vector_store_id, self._embedding_vector(embedding_response), params, vector_store_search_optional_params
|
||||
)
|
||||
|
||||
try:
|
||||
client: Final = self.async_client_factory(key)
|
||||
target: Final = client[database][collection] # pyright: ignore[reportIndexIssue] # factory is typed as returning object so injected doubles are accepted
|
||||
cursor: Final = await target.aggregate(pipeline)
|
||||
documents: Final = [ # mutable-ok: an async comprehension cannot build a tuple directly
|
||||
document async for document in cursor
|
||||
]
|
||||
except Exception as e:
|
||||
raise translate_mongo_error(e, index_name=vector_store_id, database=database, collection=collection) from e
|
||||
if not documents:
|
||||
try:
|
||||
index_cursor: Final = await target.list_search_indexes(vector_store_id)
|
||||
catalogue: Final = [ # mutable-ok: an async comprehension cannot build a tuple directly
|
||||
entry async for entry in index_cursor
|
||||
]
|
||||
except Exception as e:
|
||||
raise translate_mongo_error(
|
||||
e, index_name=vector_store_id, database=database, collection=collection
|
||||
) from e
|
||||
self._raise_for_unusable_index(catalogue, vector_store_id, database, collection)
|
||||
self._raise_for_missing_text_field(documents, params.text_field, database, collection)
|
||||
return self._to_response(documents, query_text, params.text_field)
|
||||
|
||||
def transform_create_vector_store_request(
|
||||
self,
|
||||
vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams,
|
||||
api_base: str,
|
||||
) -> NoReturn:
|
||||
raise config_error(_SEARCH_ONLY_MESSAGE)
|
||||
|
||||
def transform_create_vector_store_response(self, response: httpx.Response) -> NoReturn:
|
||||
raise config_error(_SEARCH_ONLY_MESSAGE)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -2,11 +2,11 @@ import json
|
|||
import re
|
||||
from collections.abc import Collection, Mapping
|
||||
from types import MappingProxyType, UnionType
|
||||
from typing import Any, Final, Union, get_args, get_origin
|
||||
from typing import Annotated, Any, Final, Union, get_args, get_origin
|
||||
|
||||
import orjson
|
||||
from fastapi import Request, UploadFile, status
|
||||
from typing_extensions import ReadOnly
|
||||
from typing_extensions import NotRequired, ReadOnly, Required
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB
|
||||
|
|
@ -18,6 +18,8 @@ from litellm.types.router import Deployment
|
|||
|
||||
_FORM_CONTENT_TYPES: Final[frozenset[str]] = frozenset({"application/x-www-form-urlencoded", "multipart/form-data"})
|
||||
|
||||
_ANNOTATION_QUALIFIERS: Final[frozenset[object]] = frozenset({Annotated, NotRequired, ReadOnly, Required})
|
||||
|
||||
|
||||
def _normalize_media_type(content_type: str) -> str:
|
||||
"""Return the bare media type per RFC 7231: strip params, trim, lowercase."""
|
||||
|
|
@ -42,9 +44,17 @@ def _is_json_content_type(content_type: str) -> bool:
|
|||
return _normalize_media_type(content_type) == "application/json"
|
||||
|
||||
|
||||
def _unqualified(annotation: object) -> object:
|
||||
"""Which qualifiers ``get_type_hints`` already stripped varies by interpreter version, so peel them all."""
|
||||
if get_origin(annotation) not in _ANNOTATION_QUALIFIERS:
|
||||
return annotation
|
||||
qualified: Final[tuple[object, ...]] = get_args(annotation)
|
||||
return _unqualified(qualified[0])
|
||||
|
||||
|
||||
def _numeric_form_type(annotation: object) -> type[int] | type[float] | None:
|
||||
"""The scalar to parse an ``int``/``float``-typed field as, else ``None``."""
|
||||
unwrapped: Final = get_args(annotation)[0] if get_origin(annotation) is ReadOnly else annotation
|
||||
unwrapped: Final = _unqualified(annotation)
|
||||
candidates: Final = (
|
||||
tuple(arg for arg in get_args(unwrapped) if arg is not type(None))
|
||||
if get_origin(unwrapped) in (Union, UnionType)
|
||||
|
|
|
|||
|
|
@ -38,6 +38,7 @@ from litellm.proxy.common_utils.timezone_utils import (
|
|||
get_budget_reset_settings,
|
||||
)
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
end_user_cache_key,
|
||||
model_access_group_cache_key,
|
||||
model_access_group_spend_counter_key,
|
||||
tag_cache_key,
|
||||
|
|
@ -177,6 +178,21 @@ def _model_access_group_cache_keys(row: _ModelAccessGroupRow) -> tuple[str, ...]
|
|||
return (model_access_group_cache_key(row.access_group_name),)
|
||||
|
||||
|
||||
def _enduser_counter_key(row: _EndUserRow) -> str:
|
||||
return f"spend:end_user:{row.user_id}"
|
||||
|
||||
|
||||
def _enduser_cache_keys(row: _EndUserRow) -> tuple[str, ...]:
|
||||
return (end_user_cache_key(row.user_id),)
|
||||
|
||||
|
||||
def _enduser_carried_spend(row: _EndUserRow, caps: Mapping[str, float]) -> float:
|
||||
if not caps:
|
||||
return 0.0
|
||||
effective_budget_id: Final[str | None] = row.budget_id or litellm.max_end_user_budget_id
|
||||
return _carried_spend(row.spend, caps.get(effective_budget_id) if effective_budget_id is not None else None)
|
||||
|
||||
|
||||
def _budget_link_where(
|
||||
budget_ids: Sequence[str],
|
||||
extra: Mapping[str, object] = MappingProxyType({}),
|
||||
|
|
@ -650,6 +666,7 @@ class ResetBudgetJob:
|
|||
if _rollover_enabled()
|
||||
else {} # mutable-ok: empty sentinel immediately frozen by MappingProxyType
|
||||
)
|
||||
endusers: Final[tuple[_EndUserRow, ...]] = await self._collect_endusers_to_reset(budget_ids)
|
||||
return _BudgetCascade(
|
||||
budgets=tuple(budgets_to_reset),
|
||||
budget_ids=budget_ids,
|
||||
|
|
@ -661,7 +678,7 @@ class ResetBudgetJob:
|
|||
for b in budgets_to_reset
|
||||
if b.budget_id is not None and b.budget_duration is not None
|
||||
),
|
||||
endusers=await self._collect_endusers_to_reset(budget_ids),
|
||||
endusers=endusers,
|
||||
counter_resets=(
|
||||
*(
|
||||
(_team_membership_counter_key(row), _row_carried_spend(row, rollover_caps))
|
||||
|
|
@ -674,6 +691,7 @@ class ResetBudgetJob:
|
|||
(_model_access_group_counter_key(row), _row_carried_spend(row, rollover_caps))
|
||||
for row in model_access_groups
|
||||
),
|
||||
*((_enduser_counter_key(row), _enduser_carried_spend(row, rollover_caps)) for row in endusers),
|
||||
),
|
||||
rollover_caps=rollover_caps,
|
||||
cache_keys=(
|
||||
|
|
@ -682,6 +700,7 @@ class ResetBudgetJob:
|
|||
*(key for row in orgs for key in _org_cache_keys(row)),
|
||||
*(key for row in tags for key in _tag_cache_keys(row)),
|
||||
*(key for row in model_access_groups for key in _model_access_group_cache_keys(row)),
|
||||
*(key for row in endusers for key in _enduser_cache_keys(row)),
|
||||
),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ from litellm.proxy._types import Litellm_EntityType
|
|||
from litellm.repositories.organization_repository import OrganizationRepository
|
||||
from litellm.repositories.table_repositories import (
|
||||
BudgetWindowSpendRepository,
|
||||
EndUserRepository,
|
||||
SpendLogsRepository,
|
||||
TeamMembershipRepository,
|
||||
)
|
||||
|
|
@ -36,6 +37,8 @@ from litellm.repositories.verification_token_repository import (
|
|||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma.types import LiteLLM_EndUserTableWhereUniqueInput
|
||||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
|
|
@ -47,6 +50,8 @@ _WINDOW_SPEND_ENTITY_TYPES: Final[Mapping[str, str]] = MappingProxyType(
|
|||
}
|
||||
)
|
||||
|
||||
END_USER_COUNTER_PREFIX: Final = "spend:end_user:"
|
||||
|
||||
_WINDOW_SPEND_LOG_FIELDS: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{
|
||||
"Key": "api_key",
|
||||
|
|
@ -74,6 +79,10 @@ class SpendCounterReseed:
|
|||
End-user and tag spend counters intentionally do not reseed here. Their
|
||||
auth paths already load the corresponding objects via get_end_user_object()
|
||||
and get_tag_objects_batch(); callers pass those values as fallback_spend.
|
||||
end_user_from_db is the one end-user read, used only as the budget floor when
|
||||
a counter sits below that cached spend: a worker that did not run the budget
|
||||
reset still caches the pre-reset end-user object, and LiteLLM_EndUserTable
|
||||
is the row the reset zeroed.
|
||||
"""
|
||||
|
||||
_locks: ClassVar["OrderedDict[str, asyncio.Lock]"] = OrderedDict()
|
||||
|
|
@ -129,7 +138,7 @@ class SpendCounterReseed:
|
|||
elif counter_key.startswith("spend:user:"):
|
||||
user_id = counter_key[len("spend:user:") :]
|
||||
row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id})
|
||||
elif counter_key.startswith("spend:end_user:") or counter_key.startswith("spend:tag:"):
|
||||
elif counter_key.startswith(END_USER_COUNTER_PREFIX) or counter_key.startswith("spend:tag:"):
|
||||
return None
|
||||
elif counter_key.startswith("spend:org:"):
|
||||
org_id: Final = counter_key[len("spend:org:") :]
|
||||
|
|
@ -143,6 +152,20 @@ class SpendCounterReseed:
|
|||
return None
|
||||
return float(getattr(row, "spend", 0.0) or 0.0)
|
||||
|
||||
@staticmethod
|
||||
async def end_user_from_db(prisma_client: Optional["PrismaClient"], counter_key: str) -> float | None:
|
||||
if prisma_client is None or not counter_key.startswith(END_USER_COUNTER_PREFIX):
|
||||
return None
|
||||
where: Final[LiteLLM_EndUserTableWhereUniqueInput] = {"user_id": counter_key[len(END_USER_COUNTER_PREFIX) :]}
|
||||
try:
|
||||
row: Final = await EndUserRepository(prisma_client).table.find_unique(where=where)
|
||||
except Exception: # noqa: BLE001 # a failed floor read falls back to the cached spend, like from_db
|
||||
verbose_proxy_logger.exception("SpendCounterReseed.end_user_from_db: failed for %s", counter_key)
|
||||
return None
|
||||
if row is None:
|
||||
return None
|
||||
return float(row.spend or 0.0)
|
||||
|
||||
@staticmethod
|
||||
def _is_key_or_team_window_counter(counter_key: str) -> bool:
|
||||
for prefix in ("spend:key:", "spend:team:"):
|
||||
|
|
|
|||
|
|
@ -1095,10 +1095,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
add_guardrail_to_applied_guardrails_header,
|
||||
)
|
||||
|
||||
# Collect all chunks
|
||||
all_chunks: Final[list[Any]] = []
|
||||
async for chunk in response:
|
||||
all_chunks.append(chunk)
|
||||
all_chunks: Final[Sequence[object]] = tuple([chunk async for chunk in response])
|
||||
|
||||
if not all_chunks or self._is_terminal_error_stream(all_chunks):
|
||||
for chunk in all_chunks:
|
||||
|
|
|
|||
|
|
@ -1,16 +1,20 @@
|
|||
from collections.abc import Mapping, Sequence
|
||||
import asyncio
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Protocol
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
from litellm.proxy._types import (
|
||||
CommonProxyErrors,
|
||||
LiteLLM_AccessGroupTable,
|
||||
LitellmUserRoles,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
_cache_access_object,
|
||||
_cache_key_object,
|
||||
|
|
@ -20,10 +24,16 @@ from litellm.proxy.auth.auth_checks import (
|
|||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
from litellm.proxy.management_helpers.access_group_team_sync import invalidate_access_group_cache
|
||||
from litellm.proxy.utils import get_prisma_client_or_throw
|
||||
from litellm.proxy.management_helpers.resource_display_names import (
|
||||
agent_display_names,
|
||||
key_display_names,
|
||||
mcp_server_display_names,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient, get_prisma_client_or_throw
|
||||
from litellm.repositories.table_repositories import AccessGroupRepository, TeamRepository
|
||||
from litellm.types.access_group import (
|
||||
AccessGroupCreateRequest,
|
||||
AccessGroupResource,
|
||||
AccessGroupResponse,
|
||||
AccessGroupUpdateRequest,
|
||||
)
|
||||
|
|
@ -37,6 +47,12 @@ class _AccessGroupRecord(Protocol):
|
|||
@property
|
||||
def access_group_id(self) -> str: ...
|
||||
|
||||
@property
|
||||
def access_mcp_server_ids(self) -> Sequence[str] | None: ...
|
||||
|
||||
@property
|
||||
def access_agent_ids(self) -> Sequence[str] | None: ...
|
||||
|
||||
@property
|
||||
def assigned_team_ids(self) -> Sequence[str] | None: ...
|
||||
|
||||
|
|
@ -50,6 +66,9 @@ class _TeamRecord(Protocol):
|
|||
@property
|
||||
def team_id(self) -> str: ...
|
||||
|
||||
@property
|
||||
def team_alias(self) -> str | None: ...
|
||||
|
||||
@property
|
||||
def access_group_ids(self) -> Sequence[str] | None: ...
|
||||
|
||||
|
|
@ -120,16 +139,75 @@ def _require_admin_view(user_api_key_dict: UserAPIKeyAuth) -> None:
|
|||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _ResourceNames:
|
||||
mcp_servers: Mapping[str, str]
|
||||
agents: Mapping[str, str]
|
||||
teams: Mapping[str, str | None]
|
||||
keys: Mapping[str, str]
|
||||
|
||||
|
||||
def _label(ids: Sequence[str], names: Mapping[str, str | None]) -> tuple[AccessGroupResource, ...]:
|
||||
return tuple(AccessGroupResource(id=resource_id, name=names.get(resource_id)) for resource_id in ids)
|
||||
|
||||
|
||||
def _record_to_response(
|
||||
record: _AccessGroupRecord, *, assigned_team_ids: Sequence[str] | None = None
|
||||
record: _AccessGroupRecord, *, assigned_team_ids: Sequence[str], names: _ResourceNames
|
||||
) -> AccessGroupResponse:
|
||||
stored: Final = record.dict()
|
||||
payload: Final = (
|
||||
stored if assigned_team_ids is None else MappingProxyType({**stored, "assigned_team_ids": assigned_team_ids})
|
||||
payload: Final = MappingProxyType(
|
||||
{
|
||||
**record.dict(),
|
||||
"assigned_team_ids": assigned_team_ids,
|
||||
"access_mcp_servers": _label(record.access_mcp_server_ids or (), names.mcp_servers),
|
||||
"access_agents": _label(record.access_agent_ids or (), names.agents),
|
||||
"assigned_teams": _label(assigned_team_ids, names.teams),
|
||||
"assigned_keys": _label(record.assigned_key_ids or (), names.keys),
|
||||
}
|
||||
)
|
||||
return AccessGroupResponse.model_validate(payload)
|
||||
|
||||
|
||||
def _ids_across(
|
||||
records: Sequence[_AccessGroupRecord], pick: Callable[[_AccessGroupRecord], Sequence[str] | None]
|
||||
) -> tuple[str, ...]:
|
||||
return tuple(dict.fromkeys(resource_id for record in records for resource_id in (pick(record) or ())))
|
||||
|
||||
|
||||
async def _responses_for(
|
||||
prisma_client: PrismaClient, records: Sequence[_AccessGroupRecord]
|
||||
) -> tuple[AccessGroupResponse, ...]:
|
||||
if not records:
|
||||
return ()
|
||||
teams: Final = await _teams_touching(TeamRepository(prisma_client).table, records)
|
||||
mcp_servers, agents, keys = await asyncio.gather(
|
||||
mcp_server_display_names(
|
||||
prisma_client,
|
||||
_ids_across(records, lambda record: record.access_mcp_server_ids),
|
||||
global_mcp_server_manager.config_mcp_servers,
|
||||
),
|
||||
agent_display_names(
|
||||
prisma_client, _ids_across(records, lambda record: record.access_agent_ids), global_agent_registry
|
||||
),
|
||||
key_display_names(prisma_client, _ids_across(records, lambda record: record.assigned_key_ids)),
|
||||
)
|
||||
names: Final = _ResourceNames(
|
||||
mcp_servers=mcp_servers,
|
||||
agents=agents,
|
||||
teams=MappingProxyType({team.team_id: team.team_alias for team in teams}),
|
||||
keys=keys,
|
||||
)
|
||||
attached: Final = _attached_team_ids_by_group(records, teams)
|
||||
return tuple(
|
||||
_record_to_response(record, assigned_team_ids=attached[record.access_group_id], names=names)
|
||||
for record in records
|
||||
)
|
||||
|
||||
|
||||
async def _response_for(prisma_client: PrismaClient, record: _AccessGroupRecord) -> AccessGroupResponse:
|
||||
(response,) = await _responses_for(prisma_client, (record,))
|
||||
return response
|
||||
|
||||
|
||||
def _attached_team_ids_by_group(
|
||||
records: Sequence[_AccessGroupRecord], teams: Sequence[_TeamRecord]
|
||||
) -> Mapping[str, tuple[str, ...]]:
|
||||
|
|
@ -144,19 +222,21 @@ def _attached_team_ids_by_group(
|
|||
return MappingProxyType({record.access_group_id: attached(record) for record in records})
|
||||
|
||||
|
||||
async def _teams_touching(team_table: _TeamTable, records: Sequence[_AccessGroupRecord]) -> Sequence[_TeamRecord]:
|
||||
"""Team rows listed on any of the groups or carrying any of them in access_group_ids."""
|
||||
group_ids: Final = tuple(record.access_group_id for record in records)
|
||||
stored_team_ids: Final = _ids_across(records, lambda record: record.assigned_team_ids)
|
||||
carrying: Final = {"access_group_ids": {"hasSome": group_ids}} # mutable-ok: prisma where is a dict
|
||||
listed: Final = {"team_id": {"in": stored_team_ids}} # mutable-ok: prisma where is a dict
|
||||
return await team_table.find_many(where={"OR": (carrying, listed)}) # mutable-ok: prisma where is a dict
|
||||
|
||||
|
||||
async def _attached_team_ids_for(
|
||||
team_table: _TeamTable, records: Sequence[_AccessGroupRecord]
|
||||
) -> Mapping[str, tuple[str, ...]]:
|
||||
if not records:
|
||||
return MappingProxyType({})
|
||||
group_ids: Final = tuple(record.access_group_id for record in records)
|
||||
stored_team_ids: Final = tuple(
|
||||
dict.fromkeys(team_id for record in records for team_id in (record.assigned_team_ids or ()))
|
||||
)
|
||||
carrying: Final = {"access_group_ids": {"hasSome": group_ids}} # mutable-ok: prisma where is a dict
|
||||
listed: Final = {"team_id": {"in": stored_team_ids}} # mutable-ok: prisma where is a dict
|
||||
teams: Final = await team_table.find_many(where={"OR": (carrying, listed)}) # mutable-ok: prisma where is a dict
|
||||
return _attached_team_ids_by_group(records, teams)
|
||||
return _attached_team_ids_by_group(records, await _teams_touching(team_table, records))
|
||||
|
||||
|
||||
async def _require_teams_exist(tx: _AccessGroupTx, team_ids: Sequence[str]) -> None:
|
||||
|
|
@ -425,7 +505,7 @@ async def create_access_group(
|
|||
proxy_logging_obj,
|
||||
)
|
||||
|
||||
return _record_to_response(record)
|
||||
return await _response_for(prisma_client, record)
|
||||
|
||||
|
||||
@router.get(
|
||||
|
|
@ -434,14 +514,13 @@ async def create_access_group(
|
|||
)
|
||||
async def list_access_groups(
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
) -> list[AccessGroupResponse]:
|
||||
) -> Sequence[AccessGroupResponse]:
|
||||
_require_admin_view(user_api_key_dict)
|
||||
prisma_client: Final = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value)
|
||||
|
||||
table: Final = AccessGroupRepository(prisma_client).table
|
||||
records: Final = await table.find_many(order={"created_at": "desc"})
|
||||
attached: Final = await _attached_team_ids_for(TeamRepository(prisma_client).table, records)
|
||||
return [_record_to_response(r, assigned_team_ids=attached[r.access_group_id]) for r in records]
|
||||
return await _responses_for(prisma_client, records)
|
||||
|
||||
|
||||
@router.get(
|
||||
|
|
@ -462,8 +541,7 @@ async def get_access_group(
|
|||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"Access group '{access_group_id}' not found",
|
||||
)
|
||||
attached: Final = await _attached_team_ids_for(TeamRepository(prisma_client).table, (record,))
|
||||
return _record_to_response(record, assigned_team_ids=attached[record.access_group_id])
|
||||
return await _response_for(prisma_client, record)
|
||||
|
||||
|
||||
@router.put(
|
||||
|
|
@ -560,7 +638,7 @@ async def update_access_group(
|
|||
await _patch_key_caches_add_access_group(keys_to_add, access_group_id, user_api_key_cache, proxy_logging_obj)
|
||||
await _patch_key_caches_remove_access_group(keys_to_remove, access_group_id, user_api_key_cache, proxy_logging_obj)
|
||||
|
||||
return _record_to_response(record)
|
||||
return await _response_for(prisma_client, record)
|
||||
|
||||
|
||||
@router.delete(
|
||||
|
|
|
|||
|
|
@ -789,6 +789,26 @@ def _for_teams(team_ids: Sequence[str | None]) -> str:
|
|||
return f" for team {', '.join(named)}" if named else ""
|
||||
|
||||
|
||||
def _validate_model_scope(llm_router: "Router | None", models: Sequence[str]) -> None:
|
||||
"""Reject a scope naming a model no request on this proxy could carry, at start rather
|
||||
than as a job that silently samples nothing. The question is "could any caller ask for
|
||||
this name", not "does it resolve for the job's teams": a user target's traffic can arrive
|
||||
on any team's key, so a team-public name is a legitimate scope for it, and an auto-router
|
||||
is one too (a forward job on router A scoped to router B samples what B serves today).
|
||||
Nothing here is ever dispatched to."""
|
||||
unreachable: Final = tuple(
|
||||
model
|
||||
for model in models
|
||||
if judge_target(llm_router, model).via == "nothing"
|
||||
and (llm_router is None or model not in llm_router.team_public_model_names)
|
||||
)
|
||||
if unreachable:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="models not served by this proxy: " + ", ".join(f"'{model}'" for model in unreachable),
|
||||
)
|
||||
|
||||
|
||||
_JUDGED_ROLES: Final[frozenset[StrategyRouterDependencyRole]] = frozenset({"tier", "default"})
|
||||
|
||||
|
||||
|
|
@ -1080,6 +1100,7 @@ class _LegRow(BaseModel):
|
|||
target_id: str
|
||||
router_name: str
|
||||
router_names: tuple[str, ...] = ()
|
||||
models: tuple[str, ...] = ()
|
||||
direction: ShadowEvalDirection
|
||||
baseline_model: str | None = None
|
||||
judge_model: str
|
||||
|
|
@ -1150,6 +1171,7 @@ def _group_response(
|
|||
for leg in sorted(legs, key=lambda leg: (leg.target_type, leg.target_id))
|
||||
),
|
||||
router_names=first.arm_router_names,
|
||||
models=first.models,
|
||||
direction=first.direction,
|
||||
baseline_model=first.baseline_model,
|
||||
judge_model=first.judge_model,
|
||||
|
|
@ -1322,7 +1344,10 @@ async def start_shadow_eval(
|
|||
A target is a virtual key, a team, or a user. Team and user targets match on the
|
||||
identity every request resolves to at auth time, so they cover JWT-authenticated
|
||||
traffic, which presents no virtual key; a user target samples that user's traffic
|
||||
across all their teams, whether it arrives on a JWT or a key they own.
|
||||
across all their teams, whether it arrives on a JWT or a key they own. models narrows
|
||||
every target to requests for those model groups, so a user plus one model samples that
|
||||
user's traffic on that model across every key they own; it is forward-only, since a
|
||||
reverse job already samples exactly the traffic its own router served.
|
||||
|
||||
A forward job answers whether the targets should adopt router_name: it samples the
|
||||
requests the router did not serve and duplicates them through it. A reverse job
|
||||
|
|
@ -1411,6 +1436,7 @@ async def start_shadow_eval(
|
|||
if data.baseline_model is not None:
|
||||
_validate_plain_model(llm_router, data.baseline_model, "baseline_model", team_ids)
|
||||
_validate_judge_is_not_a_candidate(llm_router, data, team_ids)
|
||||
_validate_model_scope(llm_router, data.models)
|
||||
|
||||
requested_targets: Final[tuple[tuple[ShadowEvalTargetType, str], ...]] = (
|
||||
*(("key", key) for key in data.api_key_ids),
|
||||
|
|
@ -1456,6 +1482,7 @@ async def start_shadow_eval(
|
|||
# a pre-router_names pod samples router_name alone, so it must be a real arm
|
||||
"router_name": data.router_names[0],
|
||||
"router_names": list(data.router_names), # mutable-ok: Prisma payload
|
||||
"models": list(data.models), # mutable-ok: Prisma payload
|
||||
"direction": data.direction,
|
||||
"baseline_model": data.baseline_model,
|
||||
"judge_model": data.judge_model,
|
||||
|
|
@ -1517,6 +1544,7 @@ async def start_shadow_eval(
|
|||
for target_type, target_id in sorted(requested_targets)
|
||||
),
|
||||
router_names=data.router_names,
|
||||
models=data.models,
|
||||
direction=data.direction,
|
||||
baseline_model=data.baseline_model,
|
||||
judge_model=data.judge_model,
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
|
|
@ -32,9 +32,6 @@ class AdmissionControlSettings:
|
|||
queue_timeout_seconds: float
|
||||
|
||||
|
||||
AdmissionControlSettingsGetter: TypeAlias = Callable[[], AdmissionControlSettings | None] # mutable-ok: Callable params
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AdmissionControlStats:
|
||||
admitted: int
|
||||
|
|
@ -66,13 +63,10 @@ class AdmissionControlMetrics:
|
|||
rejected_counter: _Counter
|
||||
|
||||
|
||||
AdmissionControlMetricsFactory: TypeAlias = Callable[[], AdmissionControlMetrics | None] # mutable-ok: Callable params
|
||||
|
||||
|
||||
class AdmissionControlState:
|
||||
"""Per-process admission counters and the in-flight semaphore shared by one worker's requests."""
|
||||
|
||||
def __init__(self, metrics_factory: AdmissionControlMetricsFactory) -> None:
|
||||
def __init__(self, metrics_factory: Callable[[], AdmissionControlMetrics | None]) -> None:
|
||||
self._metrics_factory = metrics_factory
|
||||
self._metrics: AdmissionControlMetrics | None = None
|
||||
self._metrics_init_attempted = False
|
||||
|
|
@ -140,7 +134,7 @@ class AdmissionControlMiddleware:
|
|||
def __init__(
|
||||
self,
|
||||
app: ASGIApp,
|
||||
get_settings: AdmissionControlSettingsGetter,
|
||||
get_settings: Callable[[], AdmissionControlSettings | None],
|
||||
state: AdmissionControlState,
|
||||
) -> None:
|
||||
self.app = app
|
||||
|
|
|
|||
|
|
@ -423,7 +423,7 @@ from litellm.proxy.db.proxy_worker_heartbeat import (
|
|||
PROXY_WORKER_HEARTBEAT_INTERVAL_SECONDS,
|
||||
ProxyWorkerHeartbeat,
|
||||
)
|
||||
from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed
|
||||
from litellm.proxy.db.spend_counter_reseed import END_USER_COUNTER_PREFIX, SpendCounterReseed
|
||||
from litellm.proxy.discovery_endpoints import ui_discovery_endpoints_router
|
||||
from litellm.proxy.fine_tuning_endpoints.endpoints import router as fine_tuning_router
|
||||
from litellm.proxy.fine_tuning_endpoints.endpoints import set_fine_tuning_config
|
||||
|
|
@ -2477,7 +2477,8 @@ async def get_current_spend(
|
|||
authoritative source depends on the counter: primary key/team/user/org
|
||||
counters read the DB row; per-window counters (``window_start`` supplied)
|
||||
read the maintained window-spend row and only aggregate spend logs when
|
||||
that row is missing or stale; end-user/tag counters have no DB row, so the caller's
|
||||
that row is missing or stale; end-user counters read ``LiteLLM_EndUserTable``, the
|
||||
row the budget reset zeroes; tag counters have no DB row, so the caller's
|
||||
``fallback_spend`` (loaded fresh in auth) is authoritative. The DB read is
|
||||
skipped for healthy primary counters (counter at or above recorded spend)
|
||||
and cached in-process for a few seconds, so a persistently stale counter
|
||||
|
|
@ -2511,8 +2512,8 @@ async def get_current_spend(
|
|||
await _repair_stale_spend_counter(counter_key=counter_key, db_spend=authoritative)
|
||||
return authoritative
|
||||
elif fallback_spend > current:
|
||||
# end-user / tag counters have no DB row; fallback_spend is the
|
||||
# authoritative recorded value loaded in auth.
|
||||
# nothing to read (tag counters, an end user without a row or a DB client, a
|
||||
# failed read); fallback_spend is the authoritative recorded value loaded in auth.
|
||||
return fallback_spend
|
||||
|
||||
# Opt-in hard guarantee: when the spend backing this admit decision came
|
||||
|
|
@ -2580,6 +2581,29 @@ async def reseed_spend_counter_from_db(counter_key: str) -> None:
|
|||
await _repair_stale_spend_counter(counter_key=counter_key, db_spend=db_spend)
|
||||
|
||||
|
||||
async def _floor_spend_from_db(
|
||||
counter_key: str,
|
||||
window_entity_type: str | None,
|
||||
window_entity_id: str | None,
|
||||
window_duration: str | None,
|
||||
window_start: datetime | None,
|
||||
) -> float | None:
|
||||
if counter_key.startswith(END_USER_COUNTER_PREFIX):
|
||||
return await SpendCounterReseed.end_user_from_db(prisma_client=prisma_client, counter_key=counter_key)
|
||||
entity_spend: Final = await SpendCounterReseed.from_db(prisma_client=prisma_client, counter_key=counter_key)
|
||||
if entity_spend is not None:
|
||||
return entity_spend
|
||||
if window_entity_type is None or window_entity_id is None or window_start is None:
|
||||
return None
|
||||
return await SpendCounterReseed.window_from_db(
|
||||
prisma_client=prisma_client,
|
||||
entity_type=window_entity_type,
|
||||
entity_id=window_entity_id,
|
||||
window_duration=window_duration,
|
||||
window_start=window_start,
|
||||
)
|
||||
|
||||
|
||||
async def _authoritative_floor_spend(
|
||||
counter_key: str,
|
||||
window_entity_type: str | None = None,
|
||||
|
|
@ -2592,20 +2616,13 @@ async def _authoritative_floor_spend(
|
|||
if cached is not None:
|
||||
return float(cached)
|
||||
|
||||
db_spend = await SpendCounterReseed.from_db(prisma_client=prisma_client, counter_key=counter_key)
|
||||
if (
|
||||
db_spend is None
|
||||
and window_entity_type is not None
|
||||
and window_entity_id is not None
|
||||
and window_start is not None
|
||||
):
|
||||
db_spend = await SpendCounterReseed.window_from_db(
|
||||
prisma_client=prisma_client,
|
||||
entity_type=window_entity_type,
|
||||
entity_id=window_entity_id,
|
||||
window_duration=window_duration,
|
||||
window_start=window_start,
|
||||
)
|
||||
db_spend: Final = await _floor_spend_from_db(
|
||||
counter_key=counter_key,
|
||||
window_entity_type=window_entity_type,
|
||||
window_entity_id=window_entity_id,
|
||||
window_duration=window_duration,
|
||||
window_start=window_start,
|
||||
)
|
||||
if db_spend is None:
|
||||
return None
|
||||
|
||||
|
|
|
|||
|
|
@ -1536,6 +1536,7 @@ model LiteLLM_ShadowEvalJob {
|
|||
target_id String // hashed virtual key, team_id, or user_id whose traffic this leg shadows
|
||||
router_name String // first (often only) auto-router under evaluation; router_names is the full set
|
||||
router_names String[] @default([]) // all routers this job runs as shadow arms; empty on legacy rows, whose set is (router_name)
|
||||
models String[] @default([]) // model groups the sampled traffic is narrowed to; empty samples every model
|
||||
direction String @default("forward") // forward | reverse
|
||||
baseline_model String? // reverse only: the fixed model the router is judged against
|
||||
judge_model String
|
||||
|
|
|
|||
|
|
@ -3365,23 +3365,24 @@ async def view_spend_logs(
|
|||
)
|
||||
sql_query, params = summary_sql_and_params
|
||||
rows: Final[Sequence[_SpendDailySummaryRow]] = await _query_raw(prisma_client, sql_query, *params)
|
||||
if len(rows) == 0:
|
||||
return [] # pyright: ignore[reportUnknownVariableType] # empty summary has no element type
|
||||
|
||||
summary_items: Final = tuple(
|
||||
_daily_summary_item(date.fromisoformat(day), tuple(day_rows))
|
||||
for day, day_rows in groupby(rows, key=lambda row: row["day"])
|
||||
)
|
||||
final_date: Final = date.fromisoformat(rows[-1]["day"])
|
||||
final_date: Final = date.fromisoformat(rows[-1]["day"]) if len(rows) > 0 else None
|
||||
end_date_date: Final = end_date_obj.date()
|
||||
padding: Final[tuple[Mapping[str, object], ...]] = tuple(
|
||||
{
|
||||
"startTime": final_date + timedelta(days=offset),
|
||||
"spend": 0,
|
||||
"users": {},
|
||||
"models": {},
|
||||
}
|
||||
for offset in range(1, (end_date_date - final_date).days + 1)
|
||||
padding: Final[tuple[Mapping[str, object], ...]] = (
|
||||
()
|
||||
if final_date is None
|
||||
else tuple(
|
||||
{
|
||||
"startTime": final_date + timedelta(days=offset),
|
||||
"spend": 0,
|
||||
"users": {},
|
||||
"models": {},
|
||||
}
|
||||
for offset in range(1, (end_date_date - final_date).days + 1)
|
||||
)
|
||||
)
|
||||
return [*summary_items, *padding]
|
||||
|
||||
|
|
|
|||
|
|
@ -37,6 +37,7 @@ from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload, SpendLogsR
|
|||
from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error
|
||||
from litellm.proxy.utils import PrismaClient, hash_token
|
||||
from litellm.types.utils import (
|
||||
PROMPT_CARRYING_GUARDRAIL_FIELDS,
|
||||
CallTypes,
|
||||
CostBreakdown,
|
||||
StandardLoggingGuardrailInformation,
|
||||
|
|
@ -1073,13 +1074,6 @@ def _sanitize_guardrail_information_for_spend_logs(
|
|||
return [_redact_prompt_fields_in_guardrail_entry(entry) for entry in entries if isinstance(entry, dict)]
|
||||
|
||||
|
||||
_PROMPT_CARRYING_GUARDRAIL_FIELDS: Final = (
|
||||
"guardrail_request",
|
||||
"guardrail_response",
|
||||
"match_details",
|
||||
"classification",
|
||||
)
|
||||
|
||||
_NUMERIC_COMPRESSION_STAT_KEYS: Final = (
|
||||
"tokens_before",
|
||||
"tokens_after",
|
||||
|
|
@ -1114,7 +1108,7 @@ def _redact_prompt_fields_in_guardrail_entry(
|
|||
preserved_stats: Final = _numeric_compression_stats_from_guardrail_response(entry.get("guardrail_response"))
|
||||
redacted: Final[StandardLoggingGuardrailInformation] = {
|
||||
**entry,
|
||||
**{key: REDACTED_BY_LITELM_STRING for key in _PROMPT_CARRYING_GUARDRAIL_FIELDS if key in entry},
|
||||
**{key: REDACTED_BY_LITELM_STRING for key in PROMPT_CARRYING_GUARDRAIL_FIELDS if key in entry},
|
||||
}
|
||||
if preserved_stats is None:
|
||||
return redacted
|
||||
|
|
|
|||
|
|
@ -56,7 +56,7 @@ def _row_to_vector_store(row: "_VectorStoreRow") -> LiteLLM_ManagedVectorStore:
|
|||
return LiteLLM_ManagedVectorStore(**row.model_dump())
|
||||
|
||||
|
||||
_LITELLM_PARAMS_MASKER: Final = SensitiveDataMasker()
|
||||
_LITELLM_PARAMS_MASKER: Final = SensitiveDataMasker(extra_sensitive_patterns=frozenset(("connection",)))
|
||||
|
||||
|
||||
_REDACT_LITELLM_PARAMS_MAX_DEPTH: Final = 10
|
||||
|
|
|
|||
|
|
@ -496,10 +496,9 @@ class BaseResponsesAPIStreamingIterator:
|
|||
if logging_response is self.completed_response:
|
||||
return
|
||||
target: Final[object] = getattr(logging_response, "response", None)
|
||||
existing_hidden: Final[object] = getattr(target, "_hidden_params", None)
|
||||
if not isinstance(existing_hidden, Mapping):
|
||||
if not isinstance(target, ResponsesAPIResponse):
|
||||
return
|
||||
existing: Final[Mapping[str, object]] = existing_hidden
|
||||
existing: Final[Mapping[str, object]] = target._hidden_params
|
||||
source_hidden: Final[object] = getattr(
|
||||
getattr(self.completed_response, "response", None), "_hidden_params", None
|
||||
)
|
||||
|
|
@ -510,15 +509,11 @@ class BaseResponsesAPIStreamingIterator:
|
|||
raw_headers: Final[Mapping[str, object]] = raw if isinstance(raw, Mapping) else EMPTY_MAPPING
|
||||
# rebuild by value and let existing keys win: sharing the source dicts would alias what the proxy
|
||||
# splats into the client's HTTP headers, and copying non-header keys would carry response_cost
|
||||
setattr( # noqa: B010 # target is typed object here, so a plain attribute store does not type check
|
||||
target,
|
||||
"_hidden_params",
|
||||
{ # mutable-ok: the cost calculator writes optional_params into _hidden_params
|
||||
"additional_headers": {**headers}, # mutable-ok: fresh copy, logging callbacks may mutate it
|
||||
"headers": {**raw_headers}, # mutable-ok: fresh copy, logging callbacks may mutate it
|
||||
**existing,
|
||||
},
|
||||
)
|
||||
target._hidden_params = { # mutable-ok: the cost calculator writes optional_params into _hidden_params
|
||||
"additional_headers": {**headers}, # mutable-ok: fresh copy, logging callbacks may mutate it
|
||||
"headers": {**raw_headers}, # mutable-ok: fresh copy, logging callbacks may mutate it
|
||||
**existing,
|
||||
}
|
||||
|
||||
def _handle_logging_completed_response(self):
|
||||
"""Base implementation - should be overridden by subclasses"""
|
||||
|
|
|
|||
|
|
@ -247,6 +247,58 @@ unless `modality_routing` is also on.
|
|||
|
||||
`session_affinity_ttl_seconds` is the idle window for both the model pin selected by session affinity and the deployment pin. Every request that reuses a pin refreshes its TTL, so a session actively sending requests stays pinned. After the window passes with no pin reuse, the next request classifies again and creates a fresh pin. Omit the setting to track the default of 3600 seconds.
|
||||
|
||||
### Mid-task stall escalation
|
||||
|
||||
A weak model working an agentic task can get stuck: it keeps calling the same tool with the
|
||||
same arguments, or the same call keeps erroring, when a stronger model would have broken the
|
||||
loop. `stall_escalation_enabled: true` catches this and bumps the request one tier higher, the
|
||||
automatic counterpart to a user typing an escalation keyword:
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: smart-router
|
||||
litellm_params:
|
||||
model: auto_router/complexity_router
|
||||
complexity_router_config:
|
||||
stall_escalation_enabled: true
|
||||
stall_escalation_window: 6
|
||||
stall_escalation_repeat_threshold: 3
|
||||
tiers:
|
||||
SIMPLE: gpt-4o-mini
|
||||
MEDIUM: gpt-4o
|
||||
COMPLEX: claude-sonnet-4
|
||||
REASONING: o1-preview
|
||||
```
|
||||
|
||||
Detection looks at the assistant's own tool calls, not the human's messages. The task counts as
|
||||
stalled when the NEWEST tool call is still part of a stuck pattern: it repeats, or it errored, at
|
||||
least `stall_escalation_repeat_threshold` times across the last `stall_escalation_window` calls.
|
||||
The tier is then bumped one step by the same `_escalate_tier` ladder `escalation_keywords` uses,
|
||||
capped at the highest configured tier. It reads both tool-call shapes: Anthropic Messages
|
||||
`tool_use`/`tool_result` blocks (including `is_error`) and chat-completions `tool_calls`/`tool`
|
||||
messages (which carry no standard error flag, so those calls are judged on repetition alone).
|
||||
|
||||
Anchoring on the newest call is what keeps a recovered task from being escalated on stale
|
||||
evidence. A model that tried the same command three times and then moved on still has those
|
||||
three calls sitting in the window for a few turns, and counting whichever pattern is most common
|
||||
in the window would escalate a request that is already making progress again. Anchoring still
|
||||
leaves room between the matches, so a retry loop broken up by an unrelated lookup counts.
|
||||
|
||||
There is no state to expire or leak: detection reruns on every classified turn from that
|
||||
request's own message list, so the bump lasts only as long as the recent tool calls still look
|
||||
stuck and lifts on its own the moment they don't. This also means it reads the whole
|
||||
conversation rather than only the turns since the newest human ask, so a plain follow-up like
|
||||
"try again" does not discard evidence from before it. Escalation records `stall_escalation` in
|
||||
`routing_decision.signals`; unlike `escalation_keywords`, it does not set the
|
||||
`escalated`/`escalation_keyword` pair, which is reserved for the keyword mechanism specifically.
|
||||
|
||||
`stall_escalation_enabled` cannot be combined with `session_affinity` or
|
||||
`classification_mode: user_turn`: both replay a held routing decision on most turns instead of
|
||||
classifying, so detection would never see the tool calls it needs to look at. It is also
|
||||
rejected together with `tier_definitions`, for the same reason `escalation_keywords` is: both
|
||||
rely on the built-in tier severity order, which a custom tier set does not define. Off by
|
||||
default.
|
||||
|
||||
### Heuristic-first chaining
|
||||
|
||||
`classifier_type: heuristic_first` runs the local scorer on every request and only calls the LLM
|
||||
|
|
|
|||
|
|
@ -40,7 +40,10 @@ from litellm.litellm_core_utils.core_helpers import (
|
|||
get_metadata_variable_name_from_kwargs,
|
||||
)
|
||||
from litellm.litellm_core_utils.internal_call_metadata import forwarded_internal_call_metadata
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import request_contains_image_content
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
as_openai_image_part,
|
||||
request_contains_image_content,
|
||||
)
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import mask_credentials_in_payload
|
||||
from litellm.llms.base_llm.base_utils import type_to_response_format_param
|
||||
from litellm.router_strategy.adaptive_router.classifier import classify_prompt
|
||||
|
|
@ -48,7 +51,11 @@ from litellm.router_strategy.complexity_router.tier_predictor import (
|
|||
TierSuccessPredictor,
|
||||
resolve_tier_artifact,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
ChatCompletionImageObject,
|
||||
ChatCompletionTextObject,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
AUTOROUTER_CLASSIFIER_CALL_ORIGIN,
|
||||
ModelResponse,
|
||||
|
|
@ -76,6 +83,7 @@ from .config import (
|
|||
ComplexityTier,
|
||||
TierDefinition,
|
||||
)
|
||||
from .stall_detector import detect_stalled_task
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from semantic_router.routers import SemanticRouter
|
||||
|
|
@ -434,6 +442,23 @@ def _strip_reminder_blocks(text: str, marker_pairs: tuple[tuple[str, str], ...]
|
|||
return " ".join(kept for a, b in zip(keep_from, keep_to) if (kept := text[a:b].strip()))
|
||||
|
||||
|
||||
def _inline_image_part(part: Mapping[str, object]) -> ChatCompletionImageObject | None:
|
||||
"""One image content part safe to hand the classifier, or None.
|
||||
|
||||
Inline data URIs only. A remote URL is caller-controlled and provider adapters do not uniformly
|
||||
delegate fetching to the provider: gigachat's file handler downloads any non-data URL with
|
||||
`client.get` from the proxy host, so forwarding one would let a key scoped to this router aim a
|
||||
proxy-side request at an internal address, on a call the caller never asked for. The routed
|
||||
model still receives the original URL exactly as before.
|
||||
"""
|
||||
converted: Final = as_openai_image_part(part)
|
||||
if converted is None:
|
||||
return None
|
||||
image_url: Final = converted["image_url"]
|
||||
url: Final = image_url if isinstance(image_url, str) else image_url.get("url", "")
|
||||
return converted if url.startswith("data:") else None
|
||||
|
||||
|
||||
def _human_text(content: object, marker_pairs: tuple[tuple[str, str], ...] = _DEFAULT_REMINDER_MARKERS) -> str:
|
||||
"""Message content as the text a human wrote, with complete reminder blocks removed.
|
||||
|
||||
|
|
@ -1591,6 +1616,10 @@ class ComplexityRouter(CustomLogger):
|
|||
threshold check alone would hand that traffic to the cheapest model without ever consulting
|
||||
the classifier. Scores also go negative when simple indicators fire, so a score threshold
|
||||
would reject exactly the trivial prompts this path exists to serve.
|
||||
|
||||
A turn carrying images the classifier would see is never decided cheaply: the scorer reads
|
||||
text alone, so its confidence describes a request it has only partly seen, and a trivial
|
||||
caption beside a screenshot is exactly the misrouting vision classification exists to stop.
|
||||
"""
|
||||
tier, score, signals, cause = self._score_and_classify(prompt, system_prompt)
|
||||
scored: Final = ClassificationOutcome(tier=tier, score=score, signals=signals, cause=cause)
|
||||
|
|
@ -1598,6 +1627,7 @@ class ComplexityRouter(CustomLogger):
|
|||
decided_cheaply: Final = (
|
||||
threshold is not None
|
||||
and bool(signals)
|
||||
and not self._classifier_image_parts(messages)
|
||||
and self._active_tier_severity(tier) <= self._active_tier_severity(threshold)
|
||||
)
|
||||
if decided_cheaply:
|
||||
|
|
@ -1622,11 +1652,43 @@ class ComplexityRouter(CustomLogger):
|
|||
tier, score, signals, cause = self._score_and_classify(prompt, system_prompt)
|
||||
scored: Final = ClassificationOutcome(tier=tier, score=score, signals=signals, cause=cause)
|
||||
margin: Final = self.config.hybrid_boundary_margin
|
||||
decided: Final = margin is not None and bool(signals) and not self._is_near_tier_boundary(score, margin)
|
||||
decided: Final = (
|
||||
margin is not None
|
||||
and bool(signals)
|
||||
and not self._classifier_image_parts(messages)
|
||||
and not self._is_near_tier_boundary(score, margin)
|
||||
)
|
||||
if decided:
|
||||
return ClassificationOutcome(tier=tier, score=score, signals=signals, cause="hybrid_short_circuit")
|
||||
return await self._llm_classifier_outcome(prompt, system_prompt, request_kwargs, messages, scored=scored)
|
||||
|
||||
def _classifier_image_parts(
|
||||
self, messages: Sequence[Mapping[str, object]] | None
|
||||
) -> tuple[ChatCompletionImageObject, ...]:
|
||||
"""Images from the newest user turn to hand the classifier, capped by max_images.
|
||||
|
||||
Empty unless the operator opted in AND the classifier model is declared vision-capable, so
|
||||
every other deployment keeps today's text-only payload byte for byte. Only the newest user
|
||||
turn is read: earlier turns are context the classifier already gets as quoted text, and an
|
||||
image nested in a tool_result is tool output rather than the ask being classified.
|
||||
Remote-URL images are left out entirely; `_inline_image_part` carries why.
|
||||
"""
|
||||
llm_config: Final = self.config.classifier_llm_config
|
||||
if llm_config is None or not llm_config.vision.enabled or not self.config.uses_llm_classifier or not messages:
|
||||
return ()
|
||||
if not self._model_declares_vision_support(llm_config.model):
|
||||
return ()
|
||||
newest_user_turn: Final = next((msg for msg in reversed(messages) if msg.get("role") == "user"), None)
|
||||
content: Final = newest_user_turn.get("content") if newest_user_turn is not None else None
|
||||
if not isinstance(content, list):
|
||||
return ()
|
||||
return tuple(
|
||||
islice(
|
||||
(part for raw in content if isinstance(raw, Mapping) and (part := _inline_image_part(raw)) is not None),
|
||||
llm_config.vision.max_images,
|
||||
)
|
||||
)
|
||||
|
||||
async def _llm_classifier_outcome(
|
||||
self,
|
||||
prompt: str,
|
||||
|
|
@ -1864,9 +1926,18 @@ class ComplexityRouter(CustomLogger):
|
|||
}
|
||||
turn_off_message_logging: Final = _effective_turn_off_message_logging(request_kwargs)
|
||||
|
||||
image_parts: Final = self._classifier_image_parts(messages)
|
||||
user_content: Final[str | Sequence[ChatCompletionTextObject | ChatCompletionImageObject]] = (
|
||||
[ # mutable-ok: SDK request payload content list is built once
|
||||
{"type": "text", "text": user_payload},
|
||||
*image_parts,
|
||||
]
|
||||
if image_parts
|
||||
else user_payload
|
||||
)
|
||||
messages_for_call: Final[list[AllMessageValues]] = [ # mutable-ok: SDK request payload list is built once
|
||||
{"role": "system", "content": classifier_system_prompt},
|
||||
{"role": "user", "content": user_payload},
|
||||
{"role": "user", "content": user_content},
|
||||
]
|
||||
response_format: Final = classifier_response_format
|
||||
classifier_call_params: Mapping[str, str] = EMPTY_MAPPING
|
||||
|
|
@ -2557,31 +2628,53 @@ class ComplexityRouter(CustomLogger):
|
|||
return pinned_model
|
||||
return self.get_model_for_tier(escalated_tier)
|
||||
|
||||
def _model_accepts_image_input(self, model_name: str) -> bool:
|
||||
"""Whether a routed model or pool entry can serve an image request.
|
||||
def _vision_verdicts(self, model_name: str) -> tuple[bool | None, ...]:
|
||||
"""Declared vision support per deployment serving the name: True, False, or None when
|
||||
nothing declares either way.
|
||||
|
||||
Resolved through the deployments that would actually serve the name; a name with no
|
||||
deployment on the router is served by the SDK directly and is checked against the model
|
||||
cost map itself. Only an explicit supports_vision false excludes, a deployment-level
|
||||
model_info override first and the map otherwise, so unmapped custom names stay routable.
|
||||
cost map itself. A deployment-level model_info override wins over the map.
|
||||
|
||||
One verdict set, two readings, because the two callers fail in opposite directions.
|
||||
Routing a user's image asks whether anything RULES IT OUT, so an undeclared model stays
|
||||
eligible and unmapped custom names keep routing. Handing an image to the classifier asks
|
||||
whether something RULES IT IN: an undeclared model that turns out to be text-only rejects
|
||||
every image request, and that rejection is swallowed by the classifier's own fallback, so
|
||||
the router quietly serves all image traffic from the fallback tier and pays for the failed
|
||||
call each time. An undeclared model instead keeps today's text-only payload, which is a
|
||||
visible no-op the operator fixes by declaring supports_vision on the deployment.
|
||||
"""
|
||||
from litellm.utils import is_vision_explicitly_disabled, supports_vision
|
||||
|
||||
def model_verdict(model: str) -> bool | None:
|
||||
if supports_vision(model):
|
||||
return True
|
||||
return False if is_vision_explicitly_disabled(model) else None
|
||||
|
||||
def deployment_verdict(deployment: Mapping[str, Any]) -> bool | None:
|
||||
declared: Final = (deployment.get("model_info") or EMPTY_MAPPING).get("supports_vision")
|
||||
if declared is not None:
|
||||
return declared is True
|
||||
return model_verdict((deployment.get("litellm_params") or EMPTY_MAPPING).get("model") or model_name)
|
||||
|
||||
deployments: Final = self.litellm_router_instance.get_model_list(model_name=model_name)
|
||||
if not deployments:
|
||||
return (model_verdict(model_name),)
|
||||
return tuple(deployment_verdict(deployment) for deployment in deployments)
|
||||
|
||||
def _model_accepts_image_input(self, model_name: str) -> bool:
|
||||
"""Whether a routed model or pool entry can serve an image request.
|
||||
|
||||
A multi-deployment group must accept on EVERY deployment: the router picks a deployment
|
||||
inside the group after this gate runs, so a mixed group marked eligible could still hand
|
||||
the image to its text-only member and fail with the exact 400 the gate exists to prevent.
|
||||
"""
|
||||
from litellm.utils import is_vision_explicitly_disabled
|
||||
return all(verdict is not False for verdict in self._vision_verdicts(model_name))
|
||||
|
||||
def deployment_accepts(deployment: Mapping[str, Any]) -> bool:
|
||||
declared: Final = (deployment.get("model_info") or EMPTY_MAPPING).get("supports_vision")
|
||||
if declared is not None:
|
||||
return declared is True
|
||||
litellm_model: Final = (deployment.get("litellm_params") or EMPTY_MAPPING).get("model") or model_name
|
||||
return not is_vision_explicitly_disabled(litellm_model)
|
||||
|
||||
deployments: Final = self.litellm_router_instance.get_model_list(model_name=model_name)
|
||||
if not deployments:
|
||||
return not is_vision_explicitly_disabled(model_name)
|
||||
return all(deployment_accepts(deployment) for deployment in deployments)
|
||||
def _model_declares_vision_support(self, model_name: str) -> bool:
|
||||
"""Whether every deployment serving the name is declared vision-capable."""
|
||||
return all(verdict is True for verdict in self._vision_verdicts(model_name))
|
||||
|
||||
def _modality_eligible_models(self) -> frozenset[str]:
|
||||
"""Every configured pool entry, plus default_model, that can serve an image request."""
|
||||
|
|
@ -3373,8 +3466,9 @@ class ComplexityRouter(CustomLogger):
|
|||
has_original_messages: Final = messages is not None and len(messages) > 0
|
||||
|
||||
user_message, system_prompt = _extract_current_ask_and_system_prompt(resolved_messages, self._reminder_markers)
|
||||
classifier_images: Final = self._classifier_image_parts(resolved_messages)
|
||||
|
||||
if user_message is None:
|
||||
if user_message is None and not classifier_images:
|
||||
verbose_router_logger.debug("ComplexityRouter: No user message found, routing to default model")
|
||||
default_model_first: Final = not self.config.plugins and self.config.default_model
|
||||
if default_model_first:
|
||||
|
|
@ -3401,8 +3495,17 @@ class ComplexityRouter(CustomLogger):
|
|||
),
|
||||
)
|
||||
|
||||
ask: Final = user_message or ""
|
||||
newest_ask: Final = _newest_turn_ask(resolved_messages, self._reminder_markers)
|
||||
escalation_keyword: Final = self._matched_escalation_keyword(newest_ask) if newest_ask is not None else None
|
||||
# Resolved here rather than beside the classifier because the keyword-override path below
|
||||
# returns before any classification runs, and a forced tier gets stuck for the same reason
|
||||
# a classified one does.
|
||||
stalled: Final = self.config.stall_escalation_enabled and detect_stalled_task(
|
||||
resolved_messages,
|
||||
window=self.config.stall_escalation_window,
|
||||
repeat_threshold=self.config.stall_escalation_repeat_threshold,
|
||||
)
|
||||
|
||||
plan_mode_sentinel: Final = self._matched_plan_mode_signal(request_kwargs, resolved_messages)
|
||||
plan_floor: Final = self._resolve_plan_mode_floor() if plan_mode_sentinel is not None else None
|
||||
|
|
@ -3430,12 +3533,13 @@ class ComplexityRouter(CustomLogger):
|
|||
),
|
||||
)
|
||||
|
||||
override: Final = await self._resolve_keyword_tier_override(user_message, request_kwargs)
|
||||
override: Final = await self._resolve_keyword_tier_override(ask, request_kwargs)
|
||||
if override is not None:
|
||||
escalated_tier: Final = (
|
||||
keyword_bumped_tier: Final = (
|
||||
self._escalate_tier(override.tier) if escalation_keyword is not None else override.tier
|
||||
)
|
||||
keyword_escalated: Final = escalated_tier != override.tier
|
||||
escalated_tier: Final = self._escalate_tier(keyword_bumped_tier) if stalled else keyword_bumped_tier
|
||||
keyword_escalated: Final = keyword_bumped_tier != override.tier
|
||||
routed_tier: Final = (
|
||||
self._apply_plan_mode_floor(escalated_tier) if plan_floor is not None else escalated_tier
|
||||
)
|
||||
|
|
@ -3463,6 +3567,7 @@ class ComplexityRouter(CustomLogger):
|
|||
conversation_continuing=conversation_continuing,
|
||||
cause=keyword_cause,
|
||||
tier=routed_tier,
|
||||
signals=("stall_escalation",) if stalled else None,
|
||||
matched_keyword=plan_mode_sentinel if keyword_plan_floored else override.matched_keyword,
|
||||
escalation_keyword=escalation_keyword,
|
||||
escalated=keyword_escalated,
|
||||
|
|
@ -3475,9 +3580,7 @@ class ComplexityRouter(CustomLogger):
|
|||
outcome: Final = (
|
||||
ClassificationOutcome(tier=housekeeping_tier, score=None, signals=("housekeeping",), cause="housekeeping")
|
||||
if housekeeping_tier is not None
|
||||
else await self.aclassify(
|
||||
user_message, system_prompt, request_kwargs, resolved_messages, raw_messages=messages
|
||||
)
|
||||
else await self.aclassify(ask, system_prompt, request_kwargs, resolved_messages, raw_messages=messages)
|
||||
)
|
||||
tier, score, signals = outcome.tier, outcome.score, outcome.signals
|
||||
classified_tier: Final = tier
|
||||
|
|
@ -3486,6 +3589,9 @@ class ComplexityRouter(CustomLogger):
|
|||
escalated: Final = tier != classified_tier
|
||||
if escalated:
|
||||
signals = (*signals, "escalation")
|
||||
if stalled:
|
||||
tier = self._escalate_tier(tier)
|
||||
signals = (*signals, "stall_escalation")
|
||||
pre_floor_tier: Final = tier
|
||||
if plan_floor is not None:
|
||||
tier = self._apply_plan_mode_floor(tier)
|
||||
|
|
@ -3544,7 +3650,7 @@ class ComplexityRouter(CustomLogger):
|
|||
# under is not a floor.
|
||||
routed_model = self._soft_floor_pick(
|
||||
tier,
|
||||
user_message,
|
||||
ask,
|
||||
request_kwargs,
|
||||
hard_floor=tier if context_original_tier is not None else plan_floor,
|
||||
hard_ceiling=housekeeping_ceiling,
|
||||
|
|
|
|||
|
|
@ -442,12 +442,47 @@ DEFAULT_TIER_MODELS: Final[dict[str, str]] = {
|
|||
}
|
||||
|
||||
|
||||
class ClassifierVisionConfig(BaseModel):
|
||||
"""Whether the LLM classifier sees the images on the request it is classifying.
|
||||
|
||||
Off by default because images cost far more than the text ask they arrive with, and the
|
||||
classifier runs on every request. A turn whose complexity lives in the image ("what is wrong in
|
||||
this stack trace screenshot") is invisible to a text-only classifier, which is what this buys.
|
||||
"""
|
||||
|
||||
enabled: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Forward image content to the classifier. Requires a classifier model declared "
|
||||
"supports_vision, on the deployment's model_info or in the model cost map; images stay "
|
||||
"stripped otherwise, so a classifier that cannot read them is never sent one. Declare "
|
||||
"model_info.supports_vision on the deployment to enable a model the cost map does not "
|
||||
"describe. Only inline data: URIs are forwarded. A request whose images are http(s) "
|
||||
"URLs still classifies on its text alone, because some providers fetch such a URL from "
|
||||
"the proxy rather than the provider, which would let a caller aim a proxy-side request "
|
||||
"at an address of their choosing."
|
||||
),
|
||||
)
|
||||
max_images: int = Field(
|
||||
default=1,
|
||||
ge=1,
|
||||
description=(
|
||||
"How many images from the newest user turn to forward, in wire order. Bounds the added "
|
||||
"cost of a turn that attaches many images. Images on earlier turns are never forwarded."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class ClassifierLLMConfig(BaseModel):
|
||||
"""Configuration for the LLM-based complexity classifier."""
|
||||
|
||||
model: str = Field(
|
||||
description="Model name (from the router's model_list) to call for classification",
|
||||
)
|
||||
vision: ClassifierVisionConfig = Field(
|
||||
default_factory=ClassifierVisionConfig,
|
||||
description="Whether the classifier sees images on the request, and how many",
|
||||
)
|
||||
reasoning_effort: REASONING_EFFORT | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
|
|
@ -852,6 +887,43 @@ class ComplexityRouterConfig(BaseModel):
|
|||
description="Rules that force a specific tier when their keywords match the prompt",
|
||||
)
|
||||
|
||||
stall_escalation_enabled: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Escalate mid-task to the next-higher configured tier when the assistant's own recent "
|
||||
"tool calls look stuck: the newest tool call repeats, or errors, at least "
|
||||
"stall_escalation_repeat_threshold times across the last stall_escalation_window "
|
||||
"calls. Both tests are anchored on the newest call, so a task that tried the same "
|
||||
"thing a few times and then moved on is not escalated on the strength of those older "
|
||||
"calls alone, while a retry loop broken up by an unrelated lookup still counts. One "
|
||||
"tier at most, on the same ladder escalation_keywords bumps along, and never above "
|
||||
"the highest configured tier. Detection re-runs on every classified turn from the "
|
||||
"tool calls visible in that request, so it needs no state and nothing survives past "
|
||||
"the task. Mutually exclusive with session_affinity and classification_mode="
|
||||
"'user_turn', which both replay a held routing decision instead of classifying most "
|
||||
"turns, so this would never see the tool calls to look at. Off by default."
|
||||
),
|
||||
)
|
||||
stall_escalation_window: int = Field(
|
||||
default=6,
|
||||
gt=0,
|
||||
description=(
|
||||
"How many of the assistant's most recent tool calls stall detection looks at, oldest "
|
||||
"ones dropped as new calls happen. Counted across the whole visible conversation "
|
||||
"rather than reset at the newest human ask, so evidence from before a plain follow-up "
|
||||
"message like 'try again' is still visible on the turn after it."
|
||||
),
|
||||
)
|
||||
stall_escalation_repeat_threshold: int = Field(
|
||||
default=3,
|
||||
ge=2,
|
||||
description=(
|
||||
"How many of the last stall_escalation_window tool calls must repeat the newest call, "
|
||||
"or must have errored alongside it, before the task counts as stalled. Must not "
|
||||
"exceed stall_escalation_window, or the condition could never be reached."
|
||||
),
|
||||
)
|
||||
|
||||
plan_mode_min_tier: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
|
|
@ -1323,6 +1395,7 @@ class ComplexityRouterConfig(BaseModel):
|
|||
("adaptive", self.adaptive),
|
||||
("session_affinity", self.session_affinity),
|
||||
("escalation_keywords", bool(self.escalation_keywords)),
|
||||
("stall_escalation_enabled", self.stall_escalation_enabled),
|
||||
("plugins", bool(self.plugins)),
|
||||
)
|
||||
if enabled
|
||||
|
|
@ -1490,6 +1563,25 @@ class ComplexityRouterConfig(BaseModel):
|
|||
)
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_stall_escalation(self) -> "ComplexityRouterConfig":
|
||||
if not self.stall_escalation_enabled:
|
||||
return self
|
||||
if self.session_affinity or self.classification_mode == "user_turn":
|
||||
raise ValueError(
|
||||
"stall_escalation_enabled cannot be combined with session_affinity or "
|
||||
"classification_mode='user_turn': both replay a held routing decision on most "
|
||||
"turns instead of classifying, so stall detection would never see the tool calls "
|
||||
"of the turns it needs to look at. Disable one or the other."
|
||||
)
|
||||
if self.stall_escalation_repeat_threshold > self.stall_escalation_window:
|
||||
raise ValueError(
|
||||
"stall_escalation_repeat_threshold "
|
||||
f"({self.stall_escalation_repeat_threshold}) cannot exceed stall_escalation_window "
|
||||
f"({self.stall_escalation_window}); the condition could never be reached."
|
||||
)
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_tier_param_placement(self) -> "ComplexityRouterConfig":
|
||||
"""Reject a router setting written into a tier entry's request params.
|
||||
|
|
|
|||
118
litellm/router_strategy/complexity_router/stall_detector.py
Normal file
118
litellm/router_strategy/complexity_router/stall_detector.py
Normal file
|
|
@ -0,0 +1,118 @@
|
|||
"""
|
||||
Mid-task stall detection for the Complexity Router.
|
||||
|
||||
Reads the assistant's own recent tool calls, which every agentic client resends on each
|
||||
turn, and reports whether the task currently looks stuck. No LLM call and no stored state:
|
||||
the same window is rescanned per classified turn, so the verdict follows the conversation
|
||||
rather than latching.
|
||||
|
||||
Tool calls arrive in two shapes and are read in place rather than translated:
|
||||
- Anthropic Messages: assistant `tool_use` content blocks, answered by a user-turn
|
||||
`tool_result` block carrying `is_error`
|
||||
- Chat completions: assistant `tool_calls` entries, answered by a `role: "tool"` message,
|
||||
which has no standard error flag, so those calls are judged on repetition alone
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from itertools import islice
|
||||
from typing import Final, NamedTuple
|
||||
|
||||
_ARGUMENTS_PARSE_FAILED: Final = object()
|
||||
|
||||
|
||||
class _ToolCallEvent(NamedTuple):
|
||||
signature: tuple[str, str]
|
||||
is_error: bool | None
|
||||
"""None where the surface reports no error status, and never counted as an error."""
|
||||
|
||||
|
||||
def _json_arguments(raw: str) -> object:
|
||||
try:
|
||||
return json.loads(raw)
|
||||
except (TypeError, ValueError):
|
||||
return _ARGUMENTS_PARSE_FAILED
|
||||
|
||||
|
||||
def _tool_call_signature(name: str, raw_arguments: object) -> tuple[str, str]:
|
||||
"""Canonicalized so the same call compares equal across both surfaces, which carry
|
||||
arguments as a dict and as a JSON string respectively."""
|
||||
parsed: Final = _json_arguments(raw_arguments) if isinstance(raw_arguments, str) else raw_arguments
|
||||
arguments: Final = raw_arguments if parsed is _ARGUMENTS_PARSE_FAILED else parsed
|
||||
try:
|
||||
return name, json.dumps(arguments, sort_keys=True, default=str)
|
||||
except (TypeError, ValueError):
|
||||
return name, str(arguments)
|
||||
|
||||
|
||||
def _iter_tool_result_error_pairs(messages: Sequence[Mapping[str, object]]) -> Iterator[tuple[str, bool]]:
|
||||
for msg in messages:
|
||||
content = msg.get("content")
|
||||
if msg.get("role") != "user" or not isinstance(content, list):
|
||||
continue
|
||||
for part in content:
|
||||
if isinstance(part, Mapping) and part.get("type") == "tool_result":
|
||||
call_id = part.get("tool_use_id")
|
||||
if isinstance(call_id, str):
|
||||
yield call_id, bool(part.get("is_error", False))
|
||||
|
||||
|
||||
def _iter_tool_call_events_newest_first(messages: Sequence[Mapping[str, object]]) -> Iterator[_ToolCallEvent]:
|
||||
error_by_call_id: Final = dict(_iter_tool_result_error_pairs(messages))
|
||||
for msg in reversed(messages):
|
||||
if msg.get("role") != "assistant":
|
||||
continue
|
||||
content = msg.get("content")
|
||||
if isinstance(content, list):
|
||||
for part in reversed(content):
|
||||
if not (isinstance(part, Mapping) and part.get("type") == "tool_use"):
|
||||
continue
|
||||
name = part.get("name")
|
||||
if isinstance(name, str):
|
||||
call_id = part.get("id")
|
||||
yield _ToolCallEvent(
|
||||
signature=_tool_call_signature(name, part.get("input")),
|
||||
is_error=error_by_call_id.get(call_id) if isinstance(call_id, str) else None,
|
||||
)
|
||||
tool_calls = msg.get("tool_calls")
|
||||
if not isinstance(tool_calls, list):
|
||||
continue
|
||||
for call in reversed(tool_calls):
|
||||
function = call.get("function") if isinstance(call, Mapping) else None
|
||||
name = function.get("name") if isinstance(function, Mapping) else None
|
||||
if isinstance(name, str):
|
||||
yield _ToolCallEvent(
|
||||
signature=_tool_call_signature(name, function.get("arguments") if function else None),
|
||||
is_error=None,
|
||||
)
|
||||
|
||||
|
||||
def detect_stalled_task(
|
||||
messages: Sequence[Mapping[str, object]] | None,
|
||||
*,
|
||||
window: int,
|
||||
repeat_threshold: int,
|
||||
) -> bool:
|
||||
"""Whether the newest tool call is still part of a stuck pattern: it repeats, or it
|
||||
errored, at least repeat_threshold times across the last `window` calls.
|
||||
|
||||
Both tests are anchored on the newest call rather than counting whichever pattern is
|
||||
most common in the window. A task that tried the same thing three times and then moved
|
||||
on has those three calls in the window for a while yet, and counting them alone would
|
||||
escalate a request that already recovered. Anchoring also leaves room between the
|
||||
matches, so a retry loop broken up by an unrelated lookup still reads as stuck.
|
||||
"""
|
||||
if not messages or repeat_threshold <= 0:
|
||||
return False
|
||||
recent: Final = tuple(islice(_iter_tool_call_events_newest_first(messages), window))
|
||||
if len(recent) < repeat_threshold:
|
||||
return False
|
||||
newest: Final = recent[0]
|
||||
repeats: Final = sum(1 for event in recent if event.signature == newest.signature)
|
||||
if repeats >= repeat_threshold:
|
||||
return True
|
||||
if not newest.is_error:
|
||||
return False
|
||||
return sum(1 for event in recent if event.is_error) >= repeat_threshold
|
||||
|
|
@ -116,28 +116,6 @@ async def aattempt(
|
|||
return RustHandled(adapt(value))
|
||||
|
||||
|
||||
def call(operation: Callable[[], ResultT], context: BridgeErrorContext) -> ResultT:
|
||||
exceptions: Final = native_exception_types()
|
||||
if exceptions is None:
|
||||
return operation()
|
||||
upstream: Final = exceptions[1]
|
||||
try:
|
||||
return operation()
|
||||
except upstream as error:
|
||||
_raise_upstream(error, context)
|
||||
|
||||
|
||||
async def acall(operation: Callable[[], Awaitable[ResultT]], context: BridgeErrorContext) -> ResultT:
|
||||
exceptions: Final = native_exception_types()
|
||||
if exceptions is None:
|
||||
return await operation()
|
||||
upstream: Final = exceptions[1]
|
||||
try:
|
||||
return await operation()
|
||||
except upstream as error:
|
||||
_raise_upstream(error, context)
|
||||
|
||||
|
||||
def _decline_reason(error: BaseException) -> str:
|
||||
reason: Final[object] = error.args[0] if error.args else str(error)
|
||||
return reason if isinstance(reason, str) else str(reason)
|
||||
|
|
@ -170,11 +148,3 @@ def _raise_upstream(error: BaseException, context: BridgeErrorContext) -> NoRetu
|
|||
llm_provider=context.provider,
|
||||
model=context.model,
|
||||
) from error
|
||||
|
||||
|
||||
def identity(value: ResultT) -> ResultT:
|
||||
return value
|
||||
|
||||
|
||||
async def async_none() -> None:
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -23,6 +23,13 @@ class AccessGroupUpdateRequest(BaseModel):
|
|||
assigned_key_ids: list[str] | None = None
|
||||
|
||||
|
||||
class AccessGroupResource(BaseModel):
|
||||
"""A resource referenced by an access group. `name` is null when the id no longer resolves or has no alias."""
|
||||
|
||||
id: str
|
||||
name: str | None
|
||||
|
||||
|
||||
class AccessGroupResponse(BaseModel):
|
||||
access_group_id: str
|
||||
access_group_name: str
|
||||
|
|
@ -32,6 +39,10 @@ class AccessGroupResponse(BaseModel):
|
|||
access_agent_ids: list[str]
|
||||
assigned_team_ids: list[str]
|
||||
assigned_key_ids: list[str]
|
||||
access_mcp_servers: tuple[AccessGroupResource, ...]
|
||||
access_agents: tuple[AccessGroupResource, ...]
|
||||
assigned_teams: tuple[AccessGroupResource, ...]
|
||||
assigned_keys: tuple[AccessGroupResource, ...]
|
||||
created_at: datetime
|
||||
created_by: str | None = None
|
||||
updated_at: datetime
|
||||
|
|
|
|||
|
|
@ -292,6 +292,18 @@ class StartShadowEvalRequest(BaseModel):
|
|||
"to across all their teams: JWT requests carrying their subject claim and virtual keys they own"
|
||||
),
|
||||
)
|
||||
models: tuple[str, ...] = Field(
|
||||
default=(),
|
||||
max_length=100,
|
||||
description=(
|
||||
"Model groups to narrow the sampled traffic to, matched on the group the caller "
|
||||
"requested and resolved through model_group_alias, so an alias and its target are one "
|
||||
"name. Empty samples every model the targets use. This ANDs with the targets: a job "
|
||||
"over a user and one model samples that user's requests on that model across every key "
|
||||
"they own, and none of their other traffic. Forward jobs only: a reverse job samples "
|
||||
"exactly the traffic its own router served, which no other model group can name"
|
||||
),
|
||||
)
|
||||
router_name: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
|
|
@ -372,12 +384,20 @@ class StartShadowEvalRequest(BaseModel):
|
|||
def _round_percentage(cls, value: float) -> float:
|
||||
return round(value, 2)
|
||||
|
||||
@field_validator("api_key_ids", "team_ids", "user_ids")
|
||||
@field_validator("api_key_ids", "team_ids", "user_ids", "models")
|
||||
@classmethod
|
||||
def _dedupe_targets(cls, value: tuple[str, ...]) -> tuple[str, ...]:
|
||||
"""A target named twice would collide with itself on the one-active-per-(target, direction) index."""
|
||||
"""A target named twice would collide with itself on the one-active-per-(target, direction)
|
||||
index; a model named twice is one scope entry."""
|
||||
return tuple(dict.fromkeys(value))
|
||||
|
||||
@field_validator("models")
|
||||
@classmethod
|
||||
def _models_are_names(cls, value: tuple[str, ...]) -> tuple[str, ...]:
|
||||
if not all(name.strip() for name in value):
|
||||
raise ValueError("models must be non-empty model group names")
|
||||
return value
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _at_least_one_target_at_most_hundred(self) -> "StartShadowEvalRequest":
|
||||
total: Final = len(self.api_key_ids) + len(self.team_ids) + len(self.user_ids)
|
||||
|
|
@ -387,6 +407,18 @@ class StartShadowEvalRequest(BaseModel):
|
|||
raise ValueError("at most 100 targets per job across api_key_ids, team_ids, and user_ids")
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _model_scope_is_forward_only(self) -> "StartShadowEvalRequest":
|
||||
"""A reverse job admits exactly the requests its own router served, so every one of
|
||||
them names that router and nothing else; any other scope would sample nothing and
|
||||
the router itself is a no-op. Both readings are rejected rather than shipped as a
|
||||
job that silently never samples."""
|
||||
if self.models and self.direction == "reverse":
|
||||
raise ValueError(
|
||||
"models is only meaningful for a forward job; a reverse job samples its own router's traffic"
|
||||
)
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _baseline_model_matches_direction(self) -> "StartShadowEvalRequest":
|
||||
if self.direction == "reverse" and self.baseline_model is None:
|
||||
|
|
@ -599,6 +631,10 @@ class ShadowEvalJobResponse(BaseModel):
|
|||
"traffic and judge every arm against the same real responses"
|
||||
),
|
||||
)
|
||||
models: tuple[str, ...] = Field(
|
||||
default=(),
|
||||
description="Model groups the sampled traffic is narrowed to; empty means every model the targets use",
|
||||
)
|
||||
direction: ShadowEvalDirection = "forward"
|
||||
baseline_model: str | None = None
|
||||
judge_model: str
|
||||
|
|
|
|||
|
|
@ -3080,6 +3080,51 @@ class GuardrailMode(TypedDict, total=False):
|
|||
|
||||
GuardrailStatus = Literal["success", "guardrail_intervened", "guardrail_failed_to_respond", "not_run"]
|
||||
|
||||
# Fields on a guardrail record whose values can quote the caller's prompt: the payload sent to the
|
||||
# guardrail, the provider response that echoes it back, and the two first-party hooks that inline
|
||||
# prompt substrings (``block_code_execution`` and ``litellm_content_filter``). Every other field
|
||||
# reports what the guardrail decided without reproducing the prompt, so redaction replaces these
|
||||
# four and keeps the rest of the record.
|
||||
PROMPT_CARRYING_GUARDRAIL_FIELDS: Final[frozenset[str]] = frozenset(
|
||||
{
|
||||
"guardrail_request",
|
||||
"guardrail_response",
|
||||
"match_details",
|
||||
"classification",
|
||||
}
|
||||
)
|
||||
|
||||
# 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_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(
|
||||
{
|
||||
"guardrail_name",
|
||||
"guardrail_provider",
|
||||
"guardrail_mode",
|
||||
"guardrail_status",
|
||||
"start_time",
|
||||
"end_time",
|
||||
"duration",
|
||||
"masked_entity_count",
|
||||
"guardrail_id",
|
||||
"policy_template",
|
||||
"detection_method",
|
||||
"confidence_score",
|
||||
"patterns_checked",
|
||||
"alert_recipients",
|
||||
"risk_score",
|
||||
"violation_categories",
|
||||
"guardrail_action",
|
||||
"guardrail_usage",
|
||||
"guardrail_cost",
|
||||
"guardrail_cost_by_unit",
|
||||
"guardrail_cost_in_spend",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class StandardLoggingGuardrailInformation(TypedDict, total=False):
|
||||
guardrail_name: str | None
|
||||
|
|
@ -3888,6 +3933,7 @@ class LlmProviders(str, Enum):
|
|||
PG_VECTOR = "pg_vector"
|
||||
S3_VECTORS = "s3_vectors"
|
||||
VALKEY = "valkey"
|
||||
MONGODB = "mongodb"
|
||||
HELICONE = "helicone"
|
||||
HYPERBOLIC = "hyperbolic"
|
||||
RECRAFT = "recraft"
|
||||
|
|
|
|||
|
|
@ -8991,6 +8991,12 @@ class ProviderConfigManager:
|
|||
)
|
||||
|
||||
return ValkeyVectorStoreConfig()
|
||||
elif litellm.LlmProviders.MONGODB == provider:
|
||||
from litellm.llms.mongodb.vector_stores.transformation import (
|
||||
MongoDBVectorStoreConfig,
|
||||
)
|
||||
|
||||
return MongoDBVectorStoreConfig()
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -2880,6 +2880,13 @@
|
|||
"vector_stores_search": true
|
||||
}
|
||||
},
|
||||
"mongodb": {
|
||||
"display_name": "MongoDB Atlas (`mongodb`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/mongodb_vector_stores",
|
||||
"endpoints": {
|
||||
"vector_stores_search": true
|
||||
}
|
||||
},
|
||||
"valkey": {
|
||||
"display_name": "Valkey (`valkey`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/valkey_vector_stores",
|
||||
|
|
|
|||
|
|
@ -112,6 +112,9 @@ utils = [
|
|||
]
|
||||
caching = ["diskcache>=5.6.3,<6.0"]
|
||||
mcp = ["mcp>=1.28.1,<2.0"]
|
||||
# Driver for the MongoDB Atlas vector store; Atlas Vector Search has no HTTP query API.
|
||||
# The floor is 4.9 because that is the release AsyncMongoClient landed in.
|
||||
mongodb = ["pymongo>=4.9,<5.0"]
|
||||
# SAML SSO for the admin UI. python3-saml pulls in xmlsec/lxml, whose wheels
|
||||
# bundle the native libxmlsec1/libxml2 libraries, so no system packages are
|
||||
# required. Kept out of the base `proxy` extra so it stays optional.
|
||||
|
|
|
|||
|
|
@ -78,7 +78,7 @@
|
|||
"limit": 1
|
||||
},
|
||||
"C901": {
|
||||
"limit": 311
|
||||
"limit": 306
|
||||
},
|
||||
"D419": {
|
||||
"limit": 6
|
||||
|
|
|
|||
|
|
@ -1536,6 +1536,7 @@ model LiteLLM_ShadowEvalJob {
|
|||
target_id String // hashed virtual key, team_id, or user_id whose traffic this leg shadows
|
||||
router_name String // first (often only) auto-router under evaluation; router_names is the full set
|
||||
router_names String[] @default([]) // all routers this job runs as shadow arms; empty on legacy rows, whose set is (router_name)
|
||||
models String[] @default([]) // model groups the sampled traffic is narrowed to; empty samples every model
|
||||
direction String @default("forward") // forward | reverse
|
||||
baseline_model String? // reverse only: the fixed model the router is judged against
|
||||
judge_model String
|
||||
|
|
|
|||
|
|
@ -142,6 +142,8 @@ fi
|
|||
|
||||
lint_dashboard() {
|
||||
(
|
||||
trap 'exit 143' TERM
|
||||
trap 'rm -f "${report:-}"' EXIT
|
||||
rc=0
|
||||
prettier_rel=()
|
||||
eslint_rel=()
|
||||
|
|
@ -168,7 +170,6 @@ EOF
|
|||
report=$(mktemp)
|
||||
npx eslint . -f json -o "$report" || true
|
||||
node scripts/check-lint-budgets.mjs "$report" eslint-budgets.json || rc=1
|
||||
rm -f "$report"
|
||||
exit $rc
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -66,6 +66,7 @@ IGNORE_FUNCTIONS = [
|
|||
"_json_safe", # max depth set (_MAX_DEPTH) plus a seen-ids cycle guard for self-referential input.
|
||||
"_redact_agent_params_tree", # max depth set (default 10), same shape as _redact_sensitive_litellm_params.
|
||||
"_restore_redacted_nested_value", # max depth set (default 10), mirrors _redact_agent_params_tree on the write side.
|
||||
"_unqualified", # bounded by the qualifier depth of a static TypedDict annotation (Annotated, Required/NotRequired, ReadOnly around one type, no cycles possible).
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -19,11 +19,13 @@ failures are hard test failures (see `tests/e2e/CLAUDE.md`).
|
|||
| OpenAI | yes | yes | yes | yes | yes (lifecycle + terminal output) | OpenAI Files |
|
||||
| Azure | yes | yes | yes | yes | yes (byte-verbatim) | Azure Files |
|
||||
| Vertex AI | yes | yes | yes | yes | yes (provider-transformed) | GCS (`gcs_bucket_name` / `GCS_BUCKET_NAME` on model) |
|
||||
| Bedrock | yes (unified only) | yes | no (limited upstream) | no | yes (provider-transformed) | S3 (`s3_bucket_name` + `aws_*` + `AWS_BATCH_ROLE_ARN` on model) |
|
||||
| Bedrock | yes (unified only) | yes | yes | yes (unfiltered managed list) | yes (provider-transformed) | S3 (`s3_bucket_name` + `aws_*` + `AWS_BATCH_ROLE_ARN` on model) |
|
||||
|
||||
Bedrock cancel is unreliable upstream and list is unsupported, so both are gated off
|
||||
(`can_cancel=False`, `can_list=False`) when that provider is enabled in the matrix;
|
||||
flipping those gates is tracked in LIT-4774 and deliberately not part of this suite.
|
||||
Bedrock cancel maps to `StopModelInvocationJob` and comes back `cancelling`; the
|
||||
lifecycle asserts it the same way it does for OpenAI (`_CANCEL_ASSERTED_PROVIDERS`).
|
||||
Bedrock has no provider-side list, so list is the proxy's DB-backed managed view: the
|
||||
unified lifecycle lists with the plain `GET /v1/batches` and the batch must appear
|
||||
there. Both were gated off until LIT-5730, after LIT-4774 landed cancel support. A batch that completes inside the 2 s pre-cancel window skips the cancel assertion (a documented vacuous pass for the cancel cell, same as OpenAI); the list assertion runs either way.
|
||||
Bedrock file upload requires a model on the request (`encoded` / `unified` scenarios only);
|
||||
`model_param` and `provider_fallback` are omitted because `POST /bedrock/v1/files` has no
|
||||
model-less passthrough path.
|
||||
|
|
@ -148,6 +150,6 @@ never landed.
|
|||
Unified (managed) batch cost is owned by the hourly `CheckBatchCost` poller, and a
|
||||
terminal DB status short-circuits retrieve for those ids, so the terminal-state cell
|
||||
uses the encoded path; poller timing does not fit an e2e gate and belongs in a
|
||||
DI-stubbed proxy integration test under `tests/test_litellm/proxy/`. Bedrock
|
||||
cancel/list stay gated pending LIT-4774. Gemini (non-Vertex) file content raises
|
||||
`NotImplementedError` upstream and is not a coverage cell.
|
||||
DI-stubbed proxy integration test under `tests/test_litellm/proxy/`. Gemini
|
||||
(non-Vertex) file content raises `NotImplementedError` upstream and is not a
|
||||
coverage cell.
|
||||
|
|
|
|||
|
|
@ -143,8 +143,8 @@ PROVIDERS: tuple[Provider, ...] = (
|
|||
"bedrock",
|
||||
batch_model_name("bedrock-batch"),
|
||||
"bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
|
||||
can_cancel=False,
|
||||
can_list=False,
|
||||
can_cancel=True,
|
||||
can_list=True,
|
||||
),
|
||||
)
|
||||
|
||||
|
|
@ -248,8 +248,9 @@ def coverage_cells_for_lifecycle(cap: Capability) -> tuple[str, ...]:
|
|||
"""Registry cell ids that the parametrized lifecycle test covers for one capability.
|
||||
|
||||
OpenAI has per-scenario cells plus granular create/retrieve/cancel/list/file
|
||||
cells. Other providers have one basic cell each. File-upload cells for the
|
||||
batch-backing path are included when the lifecycle uploads for that provider.
|
||||
cells. Bedrock adds cancel and list cells behind its gates. Other providers
|
||||
have one basic cell each. File-upload cells for the batch-backing path are
|
||||
included when the lifecycle uploads for that provider.
|
||||
"""
|
||||
match cap.provider:
|
||||
case "openai":
|
||||
|
|
@ -279,6 +280,8 @@ def coverage_cells_for_lifecycle(cap: Capability) -> tuple[str, ...]:
|
|||
return (
|
||||
"llm.batches.bedrock.basic.nonstream.works",
|
||||
"llm.files.bedrock.upload.nonstream.works",
|
||||
*(("llm.batches.bedrock.cancel.nonstream.works",) if cap.can_cancel else ()),
|
||||
*(("llm.batches.bedrock.list.nonstream.works",) if cap.can_list else ()),
|
||||
)
|
||||
case _:
|
||||
return ()
|
||||
|
|
|
|||
|
|
@ -77,7 +77,7 @@ BATCH_OP_RETRIES = 5
|
|||
# (connection refused, brief 500s) and the registry only has one basic cell per
|
||||
# provider (shared across scenarios). Create + retrieve already prove routing;
|
||||
# cancel is still deferred for cleanup, just not asserted for these two.
|
||||
_CANCEL_ASSERTED_PROVIDERS = frozenset({"openai"})
|
||||
_CANCEL_ASSERTED_PROVIDERS = frozenset({"openai", "bedrock"})
|
||||
|
||||
|
||||
def _transient_status(status_code: int) -> bool:
|
||||
|
|
@ -293,14 +293,13 @@ def test_batch_lifecycle(
|
|||
f"batch reached {pre_cancel.status!r} before cancel; "
|
||||
"provider likely rejected the input"
|
||||
)
|
||||
if pre_cancel.status == "completed":
|
||||
return
|
||||
cancelled = cancel_batch(client, batch.id, key=key, provider=provider)
|
||||
assert cancelled.id == batch.id
|
||||
assert cancelled.object == "batch"
|
||||
assert cancelled.status in {"cancelling", "cancelled"}, (
|
||||
f"unexpected post-cancel status {cancelled.status!r}"
|
||||
)
|
||||
if pre_cancel.status != "completed":
|
||||
cancelled = cancel_batch(client, batch.id, key=key, provider=provider)
|
||||
assert cancelled.id == batch.id
|
||||
assert cancelled.object == "batch"
|
||||
assert cancelled.status in {"cancelling", "cancelled"}, (
|
||||
f"unexpected post-cancel status {cancelled.status!r}"
|
||||
)
|
||||
|
||||
if cap.can_list:
|
||||
list_result = client.list_batches(key=key, provider=provider)
|
||||
|
|
|
|||
|
|
@ -23,6 +23,8 @@
|
|||
- {id: llm.batches.vertex.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:98", rationale: "Vertex batches"}
|
||||
- {id: llm.batches.bedrock.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:98", rationale: "Bedrock batches (encoded/unified only)"}
|
||||
- {id: llm.batches.bedrock.assume_role.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: assume_role, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Bedrock batch create under STS assume-role credentials"}
|
||||
- {id: llm.batches.bedrock.cancel.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Bedrock batch cancel (StopModelInvocationJob) returns the same id with a cancelling/cancelled status"}
|
||||
- {id: llm.batches.bedrock.list.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "A Bedrock managed batch is present in the GET /v1/batches list envelope"}
|
||||
- {id: llm.batches.hosted_vllm.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: batches, route: hosted_vllm, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "hosted_vllm OpenAI-compatible batch create"}
|
||||
- {id: llm.batches.openai.key_model_access_denied.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Key model restriction 403 on upload/create"}
|
||||
- {id: llm.batches.openai.input_validation.nonstream.works, module: llm, tier: P1, subject_endpoint: batches, route: openai, capability: input_validation, streaming: nonstream, assertions: [works], source: "vendor strategy §9.18 / LIT-4778", rationale: "Missing input_file_id and invalid batch id rejected"}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ spelling of prompt-cache counts (`prompt_tokens_details.cached_tokens`).
|
|||
import json
|
||||
import os
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -20,6 +20,8 @@ import pytest
|
|||
import litellm
|
||||
from litellm.integrations.datadog.datadog_llm_obs import DataDogLLMObsLogger
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import StandardLoggingGuardrailInformation
|
||||
|
||||
TOOL_DEFINITION: dict[str, Any] = {
|
||||
"type": "function",
|
||||
|
|
@ -631,10 +633,133 @@ def test_redaction_drops_every_prompt_carrying_metadata_record(logger: DataDogLL
|
|||
for record in sensitive_metadata:
|
||||
assert record not in redacted["meta"]["metadata"]
|
||||
assert record in unredacted["meta"]["metadata"]
|
||||
assert redacted["meta"]["metadata"]["guardrail_information"] is None
|
||||
assert redacted["meta"]["metadata"]["guardrail_information"] == [
|
||||
{"guardrail_name": "g", "guardrail_request": "REDACTED_BY_LITELM"}
|
||||
] # the record survives; only the field quoting the prompt is replaced
|
||||
assert unredacted["meta"]["metadata"]["guardrail_information"] is not None
|
||||
|
||||
|
||||
_AUDIT_RECORD: Final[StandardLoggingGuardrailInformation] = StandardLoggingGuardrailInformation(
|
||||
guardrail_name="bedrock-pii",
|
||||
guardrail_provider="bedrock",
|
||||
guardrail_mode=GuardrailEventHooks.pre_call,
|
||||
guardrail_status="guardrail_intervened",
|
||||
guardrail_response={"action": "MASK", "match": "alice@acme.com"},
|
||||
match_details=[{"pattern": "email", "match": "alice@acme.com"}],
|
||||
classification="the user asked for alice@acme.com",
|
||||
masked_entity_count={"EMAIL": 2},
|
||||
violation_categories=["pii"],
|
||||
duration=0.01,
|
||||
)
|
||||
|
||||
|
||||
def _payload_with_guardrail_record(guardrail_information: object) -> dict[str, Any]:
|
||||
payload = build_payload()
|
||||
payload["standard_logging_object"]["guardrail_information"] = guardrail_information
|
||||
return payload
|
||||
|
||||
|
||||
def test_redaction_keeps_the_guardrail_audit_record(logger: DataDogLLMObsLogger) -> None:
|
||||
"""Redaction removes the prompt, not the operator's record that a guardrail intervened."""
|
||||
redacted = _span_json(
|
||||
_redacting_logger(turn_off_message_logging=True),
|
||||
_payload_with_guardrail_record([dict(_AUDIT_RECORD)]),
|
||||
)
|
||||
record = redacted["meta"]["metadata"]["guardrail_information"][0]
|
||||
|
||||
for field in ("guardrail_request", "guardrail_response", "match_details", "classification"):
|
||||
assert record.get(field, "REDACTED_BY_LITELM") == "REDACTED_BY_LITELM"
|
||||
assert record["guardrail_name"] == "bedrock-pii"
|
||||
assert record["guardrail_provider"] == "bedrock"
|
||||
assert record["guardrail_mode"] == "pre_call"
|
||||
assert record["guardrail_status"] == "guardrail_intervened"
|
||||
assert record["masked_entity_count"] == {"EMAIL": 2}
|
||||
assert record["violation_categories"] == ["pii"]
|
||||
assert record["duration"] == 0.01
|
||||
assert "alice@acme.com" not in safe_dumps(redacted["meta"]["metadata"])
|
||||
|
||||
|
||||
def test_a_caller_supplied_redaction_header_cannot_blank_the_guardrail_record(
|
||||
logger: DataDogLLMObsLogger,
|
||||
) -> None:
|
||||
"""Any key may redact its own prompts with the header; none may erase what a guardrail caught."""
|
||||
payload = _payload_with_guardrail_record([dict(_AUDIT_RECORD)])
|
||||
payload["litellm_params"] = {"metadata": {"headers": {"x-litellm-enable-message-redaction": "true"}}}
|
||||
|
||||
span = _span_json(logger, payload)
|
||||
record = span["meta"]["metadata"]["guardrail_information"][0]
|
||||
|
||||
assert span["meta"]["input"]["messages"] == [{"role": "user", "content": "redacted-by-litellm"}]
|
||||
assert record["guardrail_status"] == "guardrail_intervened"
|
||||
assert record["masked_entity_count"] == {"EMAIL": 2}
|
||||
assert "alice@acme.com" not in safe_dumps(span["meta"]["metadata"])
|
||||
|
||||
|
||||
def test_a_guardrails_own_extra_field_never_reaches_a_redacted_span(logger: DataDogLLMObsLogger) -> None:
|
||||
"""A guardrail may record whatever it likes; only classified fields survive redaction."""
|
||||
span = _span_json(
|
||||
_redacting_logger(turn_off_message_logging=True),
|
||||
_payload_with_guardrail_record([{**_AUDIT_RECORD, "matched_text": "the caller asked about alice@acme.com"}]),
|
||||
)
|
||||
record = span["meta"]["metadata"]["guardrail_information"][0]
|
||||
|
||||
assert "matched_text" not in record
|
||||
assert record["guardrail_status"] == "guardrail_intervened"
|
||||
assert "alice@acme.com" not in safe_dumps(span["meta"]["metadata"])
|
||||
|
||||
|
||||
def test_a_lone_guardrail_record_survives_redaction(logger: DataDogLLMObsLogger) -> None:
|
||||
"""A guardrail that writes the metadata key itself leaves one record, not a list of them."""
|
||||
span = _span_json(
|
||||
_redacting_logger(turn_off_message_logging=True),
|
||||
_payload_with_guardrail_record(dict(_AUDIT_RECORD)),
|
||||
)
|
||||
metadata = span["meta"]["metadata"]
|
||||
|
||||
assert metadata["guardrail_information"] == [
|
||||
{
|
||||
"guardrail_name": "bedrock-pii",
|
||||
"guardrail_provider": "bedrock",
|
||||
"guardrail_mode": "pre_call",
|
||||
"guardrail_status": "guardrail_intervened",
|
||||
"guardrail_response": "REDACTED_BY_LITELM",
|
||||
"match_details": "REDACTED_BY_LITELM",
|
||||
"classification": "REDACTED_BY_LITELM",
|
||||
"masked_entity_count": {"EMAIL": 2},
|
||||
"violation_categories": ["pii"],
|
||||
"duration": 0.01,
|
||||
}
|
||||
]
|
||||
assert metadata["latency_metrics"]["guardrail_overhead_time_ms"] == 10.0
|
||||
|
||||
|
||||
@pytest.mark.parametrize("guardrail_information", [None, [], 5, "abc", [None, "x"], {}])
|
||||
def test_odd_guardrail_shapes_still_produce_a_span(
|
||||
guardrail_information: object,
|
||||
) -> None:
|
||||
"""The redacted branch replaced an expression that could not fail, so it must not start failing."""
|
||||
span = _span_json(
|
||||
_redacting_logger(turn_off_message_logging=True),
|
||||
_payload_with_guardrail_record(guardrail_information),
|
||||
)
|
||||
|
||||
assert span["meta"]["input"]["messages"] == [{"role": "user", "content": "redacted-by-litellm"}]
|
||||
assert span["meta"]["metadata"]["guardrail_information"] in (None, [], [{}])
|
||||
|
||||
|
||||
def test_a_redacted_span_carries_every_declared_guardrail_field() -> None:
|
||||
"""A field added to the record without a redaction decision would be dropped, so it fails here."""
|
||||
declared = dict.fromkeys(StandardLoggingGuardrailInformation.__annotations__, "alice@acme.com")
|
||||
payload = _payload_with_guardrail_record([{**declared, "duration": 0.01}])
|
||||
|
||||
span = _span_json(_redacting_logger(turn_off_message_logging=True), payload)
|
||||
record = span["meta"]["metadata"]["guardrail_information"][0]
|
||||
|
||||
assert set(record) == set(declared)
|
||||
for field in ("guardrail_request", "guardrail_response", "match_details", "classification"):
|
||||
assert record[field] == "REDACTED_BY_LITELM"
|
||||
|
||||
|
||||
def test_tool_definitions_accept_the_bare_anthropic_shape(logger: DataDogLLMObsLogger) -> None:
|
||||
"""The Anthropic surface declares tools unwrapped, with input_schema instead of parameters."""
|
||||
payload = build(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -24,7 +24,13 @@ from litellm.integrations.shadow_eval_logger import (
|
|||
_unmask_preference,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import SHADOW_EVAL_JUDGE_CALL_ORIGIN, SHADOW_EVAL_ROUTER_CALL_ORIGIN, ModelResponse
|
||||
from litellm.types.utils import (
|
||||
SHADOW_EVAL_JUDGE_CALL_ORIGIN,
|
||||
SHADOW_EVAL_ROUTER_CALL_ORIGIN,
|
||||
ChatCompletionCustomToolCallPayload,
|
||||
ChatCompletionMessageCustomToolCall,
|
||||
ModelResponse,
|
||||
)
|
||||
|
||||
|
||||
def _job(**overrides) -> ActiveShadowEvalJob:
|
||||
|
|
@ -66,6 +72,7 @@ def _job_record(job: ActiveShadowEvalJob, target_type="key", target_id="key-hash
|
|||
target_id=target_id,
|
||||
router_name=job.router_name,
|
||||
router_names=job.router_names,
|
||||
models=sorted(job.models),
|
||||
direction=job.direction,
|
||||
baseline_model=job.baseline_model,
|
||||
shadow_percentage=job.shadow_percentage,
|
||||
|
|
@ -120,6 +127,39 @@ def _router(
|
|||
return router
|
||||
|
||||
|
||||
def _shadow_reply_router(message, finish_reason="stop", routed_model="cheap-model"):
|
||||
"""A router whose shadow arm answers with a caller-supplied message, so a reply that
|
||||
yields no judgeable text can be posed as the two different things it can be: an arm
|
||||
that chose a tool, or an arm that returned nothing."""
|
||||
router = MagicMock()
|
||||
router.model_group_alias = {}
|
||||
router.get_model_list = MagicMock(return_value=[{"litellm_params": {"model": "openai/gpt-4o-mini"}}])
|
||||
|
||||
async def acompletion(**kwargs):
|
||||
if kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN:
|
||||
return {"choices": [{"message": {"content": '{"preference": "A", "confidence": 0.9}'}}]}
|
||||
kwargs["metadata"]["routing_decision"] = {"tier_label": "SIMPLE", "routed_model": routed_model}
|
||||
return {"choices": [{"message": message, "finish_reason": finish_reason}]}
|
||||
|
||||
router.acompletion = MagicMock(side_effect=acompletion)
|
||||
return router
|
||||
|
||||
|
||||
TOOL_CALL_MESSAGE = {
|
||||
"content": None,
|
||||
"tool_calls": [{"id": "c1", "type": "function", "function": {"name": "Read", "arguments": "{}"}}],
|
||||
}
|
||||
|
||||
CUSTOM_TOOL_CALL_MESSAGE = {
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
ChatCompletionMessageCustomToolCall(
|
||||
id="c2", custom=ChatCompletionCustomToolCallPayload(name="exec_sql", input="select 1")
|
||||
)
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def _spend_counter(store=None):
|
||||
"""In-memory stand-in for the proxy's cross-pod spend counter: reads take the max of
|
||||
the counter and the caller's fallback, exactly like get_current_spend does for a key
|
||||
|
|
@ -166,6 +206,7 @@ def _success_kwargs(
|
|||
request_metadata=None,
|
||||
call_type="acompletion",
|
||||
model="claude-opus",
|
||||
model_group="opus-group",
|
||||
response_cost=None,
|
||||
cache_hit=None,
|
||||
):
|
||||
|
|
@ -174,6 +215,7 @@ def _success_kwargs(
|
|||
"id": request_id,
|
||||
"call_type": call_type,
|
||||
"model": model,
|
||||
"model_group": model_group,
|
||||
"metadata": {"user_api_key_hash": api_key_hash},
|
||||
"model_parameters": {"temperature": 0.5, "stream": True},
|
||||
"response_cost": response_cost,
|
||||
|
|
@ -368,7 +410,13 @@ class TestSurfaceNormalization:
|
|||
],
|
||||
ids=["tool-final-chat-turn", "tool-final-responses-turn"],
|
||||
)
|
||||
async def test_unjudgeable_turns_are_skipped_without_consuming_budget(self, response_mutation, kwargs_mutation):
|
||||
async def test_a_tool_final_turn_is_sampled_and_serialized_for_the_judge(
|
||||
self, response_mutation, kwargs_mutation
|
||||
):
|
||||
"""A turn where the real model called a tool used to be dropped before sampling, on
|
||||
every surface. On agentic traffic that is most of the traffic, so a job set to
|
||||
sample 10% was really sampling 10% of the prose-only slice and calling it 10% of
|
||||
the key. The turn is sampled like any other and the call is serialized as text."""
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
|
||||
hook_kwargs = _success_kwargs(**({"call_type": "acompletion"} | kwargs_mutation))
|
||||
|
|
@ -406,6 +454,38 @@ class TestSurfaceNormalization:
|
|||
|
||||
prisma, router = await self._drive(hook_kwargs, response)
|
||||
|
||||
judge_prompt = next(
|
||||
call.kwargs["messages"][-1]["content"]
|
||||
for call in router.acompletion.call_args_list
|
||||
if call.kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN
|
||||
)
|
||||
assert "[tool call] f({})" in judge_prompt
|
||||
prisma.db.litellm_shadowevalattempt.create.assert_called_once()
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"response_mutation,kwargs_mutation",
|
||||
[
|
||||
("chat-no-content", {}),
|
||||
("responses-no-output", {"call_type": "aresponses"}),
|
||||
],
|
||||
ids=["empty-chat-turn", "empty-responses-turn"],
|
||||
)
|
||||
async def test_turns_with_nothing_to_compare_are_skipped_without_consuming_budget(
|
||||
self, response_mutation, kwargs_mutation
|
||||
):
|
||||
"""No prose and no tool call leaves the judge nothing to score, so the turn is
|
||||
still skipped rather than billed."""
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
|
||||
hook_kwargs = _success_kwargs(**({"call_type": "acompletion"} | kwargs_mutation))
|
||||
if response_mutation == "chat-no-content":
|
||||
response = {"choices": [{"message": {"content": ""}}]}
|
||||
else:
|
||||
hook_kwargs["messages"] = "do the thing"
|
||||
response = ResponsesAPIResponse.model_validate(RESPONSES_API_RESPONSE | {"output": []})
|
||||
|
||||
prisma, router = await self._drive(hook_kwargs, response)
|
||||
|
||||
router.acompletion.assert_not_called()
|
||||
prisma.db.litellm_shadowevalattempt.create.assert_not_called()
|
||||
|
||||
|
|
@ -930,6 +1010,79 @@ class TestTargetMatching:
|
|||
assert logger._job_starts == {"key-job": 1, "team-job": 1}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestModelScope:
|
||||
"""A job scoped to model groups samples a target's request only when the group the
|
||||
caller asked for is one of them; an out-of-scope request is not the job's traffic at
|
||||
all, so it records no funnel event, exactly like a direction mismatch."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"requested,sampled",
|
||||
[("sonnet-group", True), ("opus-group", False), ("", False)],
|
||||
ids=["in-scope-group-samples", "other-group-skips", "unknown-group-fails-closed"],
|
||||
)
|
||||
async def test_scope_admits_only_the_named_groups_and_counts_nothing_else(self, requested, sampled):
|
||||
prisma = _prisma()
|
||||
logger = _logger(router=_router(), prisma=prisma, jobs=(_job(models=frozenset({"sonnet-group", "haiku-group"})),))
|
||||
|
||||
await logger.async_log_success_event(_success_kwargs(model_group=requested), RESPONSE, None, None)
|
||||
await _drain(logger)
|
||||
|
||||
assert prisma.db.litellm_shadowevalattempt.create.await_count == (1 if sampled else 0)
|
||||
assert logger._test_funnel == []
|
||||
|
||||
async def test_an_unscoped_job_samples_every_group(self):
|
||||
prisma = _prisma()
|
||||
logger = _logger(router=_router(), prisma=prisma, jobs=(_job(),))
|
||||
|
||||
await logger.async_log_success_event(_success_kwargs(model_group="anything"), RESPONSE, None, None)
|
||||
await _drain(logger)
|
||||
|
||||
prisma.db.litellm_shadowevalattempt.create.assert_awaited_once()
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"scoped_to,requested",
|
||||
[("sonnet-group", "fast"), ("fast", "sonnet-group")],
|
||||
ids=["job-names-the-target-request-uses-the-alias", "job-names-the-alias-request-uses-the-target"],
|
||||
)
|
||||
async def test_an_alias_and_its_target_are_one_group_on_both_sides(self, scoped_to, requested):
|
||||
"""Both the job's scope and the request's group resolve through the router's alias
|
||||
map at match time, so re-pointing an alias follows config rather than freezing at
|
||||
job start."""
|
||||
router = _router()
|
||||
router.model_group_alias = {"fast": "sonnet-group"}
|
||||
prisma = _prisma(jobs=[_job_record(_job(models=frozenset({scoped_to})))])
|
||||
logger = _logger(router=router, prisma=prisma)
|
||||
|
||||
await logger.async_log_success_event(_success_kwargs(model_group=requested), RESPONSE, None, None)
|
||||
await _drain(logger)
|
||||
|
||||
prisma.db.litellm_shadowevalattempt.create.assert_awaited_once()
|
||||
assert prisma.db.litellm_shadowevaljob.find_many.await_count == 1
|
||||
|
||||
async def test_a_repointed_alias_applies_to_the_next_request_without_a_cache_refill(self):
|
||||
router = _router()
|
||||
router.model_group_alias = {"fast": "sonnet-group"}
|
||||
prisma = _prisma(jobs=[_job_record(_job(models=frozenset({"fast"})))])
|
||||
logger = _logger(router=router, prisma=prisma)
|
||||
await logger.async_log_success_event(_success_kwargs(model_group="sonnet-group"), RESPONSE, None, None)
|
||||
await _drain(logger)
|
||||
assert prisma.db.litellm_shadowevalattempt.create.await_count == 1
|
||||
|
||||
router.model_group_alias = {"fast": "haiku-group"}
|
||||
await logger.async_log_success_event(
|
||||
_success_kwargs(request_id="req-2", model_group="sonnet-group"), RESPONSE, None, None
|
||||
)
|
||||
await logger.async_log_success_event(
|
||||
_success_kwargs(request_id="req-3", model_group="haiku-group"), RESPONSE, None, None
|
||||
)
|
||||
await _drain(logger)
|
||||
|
||||
rows = [call.kwargs["data"]["request_id"] for call in prisma.db.litellm_shadowevalattempt.create.call_args_list]
|
||||
assert rows == ["req-1", "req-3"]
|
||||
assert prisma.db.litellm_shadowevaljob.find_many.await_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestActiveJobsCache:
|
||||
async def test_cache_miss_reads_db_once_then_serves_from_cache(self):
|
||||
|
|
@ -1134,6 +1287,206 @@ class TestShadowPipeline:
|
|||
assert row["shadow_cost"] == 0.007
|
||||
assert logger._test_counter["spend:shadow_eval:job-1"] == 0.007
|
||||
|
||||
async def _no_text_error(self, router) -> str:
|
||||
prisma = _prisma()
|
||||
await _logger(router=router, prisma=prisma)._run_shadow_eval(
|
||||
job=_job(),
|
||||
request_id="req-1",
|
||||
messages=({"role": "user", "content": "hi"},),
|
||||
real_text="real answer",
|
||||
real_model="claude-opus",
|
||||
real_cost=0.0,
|
||||
real_classifier_cost=0.0,
|
||||
real_cache_hit=False,
|
||||
control_tier=None,
|
||||
shadow_params={},
|
||||
parent_metadata={},
|
||||
)
|
||||
row = prisma.db.litellm_shadowevalattempt.create.call_args.kwargs["data"]
|
||||
assert row["outcome"] == "error"
|
||||
return row["error"]
|
||||
|
||||
async def _judged_shadow_row(self, router: MagicMock, shadow_params: dict | None = None) -> dict:
|
||||
prisma = _prisma()
|
||||
await _logger(router=router, prisma=prisma)._run_shadow_eval(
|
||||
job=_job(),
|
||||
request_id="req-1",
|
||||
messages=({"role": "user", "content": "hi"},),
|
||||
real_text="real answer",
|
||||
real_model="claude-opus",
|
||||
real_cost=0.0,
|
||||
real_classifier_cost=0.0,
|
||||
real_cache_hit=False,
|
||||
control_tier=None,
|
||||
shadow_params=shadow_params or {},
|
||||
parent_metadata={},
|
||||
)
|
||||
return prisma.db.litellm_shadowevalattempt.create.call_args.kwargs["data"]
|
||||
|
||||
async def test_a_tool_call_shadow_reply_is_judged_rather_than_discarded(self):
|
||||
"""An arm that calls a tool where the real model wrote prose has answered, it just
|
||||
answered by acting. Dropping that turn threw away the comparison the job exists to
|
||||
make, and on agentic traffic it threw away most of them, so the tool call is
|
||||
serialized into text and judged like any other response."""
|
||||
row = await self._judged_shadow_row(_shadow_reply_router(TOOL_CALL_MESSAGE, finish_reason="tool_calls"))
|
||||
|
||||
assert row["outcome"] != "error"
|
||||
assert row["error"] is None
|
||||
assert row["confidence"] == 0.9
|
||||
|
||||
async def test_a_tool_call_reaches_the_judge_as_readable_text(self):
|
||||
"""The judge only ever sees strings, so a tool call has to arrive as its name and
|
||||
arguments. A serialization that dropped either would ask the judge to score a
|
||||
response it cannot tell apart from any other tool call."""
|
||||
router = _shadow_reply_router(TOOL_CALL_MESSAGE, finish_reason="tool_calls")
|
||||
await self._judged_shadow_row(router)
|
||||
|
||||
judge_prompt = next(
|
||||
call.kwargs["messages"][-1]["content"]
|
||||
for call in router.acompletion.call_args_list
|
||||
if call.kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN
|
||||
)
|
||||
|
||||
assert "[tool call] Read({})" in judge_prompt
|
||||
|
||||
async def test_the_judge_sees_what_tools_were_available(self):
|
||||
"""Scoring whether a tool call was the right response needs to know what else the
|
||||
arm could have called instead. Without the tool list, the judge can score the
|
||||
arguments but not whether Read, specifically, was the correct choice."""
|
||||
router = _shadow_reply_router(TOOL_CALL_MESSAGE, finish_reason="tool_calls")
|
||||
tools = [
|
||||
{"type": "function", "function": {"name": "Read", "description": "read a file from disk"}},
|
||||
{"type": "function", "function": {"name": "Bash", "description": "run a shell command"}},
|
||||
]
|
||||
await self._judged_shadow_row(router, shadow_params={"tools": tools})
|
||||
|
||||
judge_prompt = next(
|
||||
call.kwargs["messages"][-1]["content"]
|
||||
for call in router.acompletion.call_args_list
|
||||
if call.kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN
|
||||
)
|
||||
|
||||
assert "Read: read a file from disk" in judge_prompt
|
||||
assert "Bash: run a shell command" in judge_prompt
|
||||
|
||||
async def test_a_custom_tool_definition_is_named_for_the_judge(self):
|
||||
"""A custom tool definition nests name and description under `custom`, not
|
||||
`function`, so reading only `function` renders every one of them as unnamed and
|
||||
tells the judge nothing about what the arm could have called."""
|
||||
from openai.types.chat import ChatCompletionCustomToolParam
|
||||
|
||||
router = _shadow_reply_router(TOOL_CALL_MESSAGE, finish_reason="tool_calls")
|
||||
tools = [
|
||||
ChatCompletionCustomToolParam(
|
||||
type="custom",
|
||||
custom={"name": "exec_sql", "description": "run a read-only sql query"},
|
||||
)
|
||||
]
|
||||
await self._judged_shadow_row(router, shadow_params={"tools": tools})
|
||||
|
||||
judge_prompt = next(
|
||||
call.kwargs["messages"][-1]["content"]
|
||||
for call in router.acompletion.call_args_list
|
||||
if call.kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN
|
||||
)
|
||||
|
||||
assert "exec_sql: run a read-only sql query" in judge_prompt
|
||||
assert "unnamed" not in judge_prompt
|
||||
|
||||
@pytest.mark.parametrize("shadow_params", [{}, {"tools": []}], ids=["omitted", "empty-list"])
|
||||
async def test_no_tool_definitions_section_when_the_turn_offered_no_tools(self, shadow_params):
|
||||
"""Padding every judge prompt with an empty tools section wastes budget on the
|
||||
turns, still the majority, that never offered one, whether tools was left out of
|
||||
the request entirely or sent as an empty list."""
|
||||
router = _shadow_reply_router({"content": "hello"}, finish_reason="stop")
|
||||
await self._judged_shadow_row(router, shadow_params=shadow_params)
|
||||
|
||||
judge_prompt = next(
|
||||
call.kwargs["messages"][-1]["content"]
|
||||
for call in router.acompletion.call_args_list
|
||||
if call.kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN
|
||||
)
|
||||
|
||||
assert "Tools available" not in judge_prompt
|
||||
|
||||
async def test_a_custom_tool_call_serializes_its_name_and_input(self):
|
||||
"""Custom tool calls carry no `function` key: name and arguments live under
|
||||
`custom`, so reading only `function` serializes every one of them as unnamed."""
|
||||
router = _shadow_reply_router(CUSTOM_TOOL_CALL_MESSAGE, finish_reason="tool_calls")
|
||||
await self._judged_shadow_row(router)
|
||||
|
||||
judge_prompt = next(
|
||||
call.kwargs["messages"][-1]["content"]
|
||||
for call in router.acompletion.call_args_list
|
||||
if call.kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN
|
||||
)
|
||||
|
||||
assert "[tool call] exec_sql(select 1)" in judge_prompt
|
||||
|
||||
async def test_the_judge_is_told_a_tool_call_is_not_a_defect(self):
|
||||
"""The judge scores on completeness and clarity. Handed a tool call with no
|
||||
instruction, it marks it down for not reading like an answer, which would bias
|
||||
every verdict against a tool-calling arm on exactly the traffic that calls tools."""
|
||||
router = _shadow_reply_router(TOOL_CALL_MESSAGE, finish_reason="tool_calls")
|
||||
await self._judged_shadow_row(router)
|
||||
|
||||
system_prompt = next(
|
||||
call.kwargs["messages"][0]["content"]
|
||||
for call in router.acompletion.call_args_list
|
||||
if call.kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN
|
||||
)
|
||||
|
||||
assert "tool call" in system_prompt
|
||||
assert "not a defect" in system_prompt
|
||||
|
||||
async def test_prose_written_alongside_a_tool_call_survives_into_the_verdict(self):
|
||||
"""Some providers write a sentence before acting. Serializing only the call would
|
||||
hide half of what the arm actually said from the judge."""
|
||||
router = _shadow_reply_router(
|
||||
{"content": "Let me look that up.", "tool_calls": TOOL_CALL_MESSAGE["tool_calls"]},
|
||||
finish_reason="tool_calls",
|
||||
)
|
||||
await self._judged_shadow_row(router)
|
||||
|
||||
judge_prompt = next(
|
||||
call.kwargs["messages"][-1]["content"]
|
||||
for call in router.acompletion.call_args_list
|
||||
if call.kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN
|
||||
)
|
||||
|
||||
assert "Let me look that up. [tool call] Read({})" in judge_prompt
|
||||
|
||||
async def test_an_empty_shadow_reply_names_the_finish_reason_and_the_routed_model(self):
|
||||
"""A reply that really carried no text is diagnosable only if the row says what
|
||||
the arm was doing when it produced none: a truncated turn and a model that answers
|
||||
with nothing are different faults with different fixes."""
|
||||
error = await self._no_text_error(
|
||||
_shadow_reply_router({"content": ""}, finish_reason="length", routed_model="some-model")
|
||||
)
|
||||
|
||||
assert "empty response" in error
|
||||
assert "finish_reason=length" in error
|
||||
assert "model=some-model" in error
|
||||
|
||||
async def test_no_text_errors_stay_groupable_across_models_and_finish_reasons(self):
|
||||
"""Operators read these rows by grouping on the error text, which is how a job's
|
||||
failures collapse to a handful of causes. Every varying part therefore has to sit
|
||||
behind the first semicolon, or each row becomes its own group and the count that
|
||||
made the problem visible stops existing."""
|
||||
first = await self._no_text_error(
|
||||
_shadow_reply_router({"content": None}, finish_reason="length", routed_model="model-a")
|
||||
)
|
||||
second = await self._no_text_error(
|
||||
_shadow_reply_router(
|
||||
{"content": ""},
|
||||
finish_reason="stop",
|
||||
routed_model="model-b",
|
||||
)
|
||||
)
|
||||
|
||||
assert first != second
|
||||
assert first.split(";")[0] == second.split(";")[0]
|
||||
|
||||
async def test_a_pipeline_error_after_the_shadow_call_keeps_its_billed_cost(self, monkeypatch: pytest.MonkeyPatch):
|
||||
"""An unexpected error between the billed shadow call and the attempt write must
|
||||
still record the shadow cost, or the per-key dollar gate undercounts forever."""
|
||||
|
|
@ -1691,11 +2044,13 @@ class TestSamplingFunnel:
|
|||
prisma.db.litellm_shadowevalattempt.create.assert_not_awaited()
|
||||
|
||||
async def test_an_unjudgeable_sampled_request_counts_unjudgeable(self):
|
||||
"""A tool call still serializes into judgeable text; a turn with neither prose nor
|
||||
a tool call to serialize is the one case left with nothing to compare."""
|
||||
prisma = _prisma()
|
||||
logger = _logger(router=_router(), prisma=prisma, jobs=(_job(),))
|
||||
tool_final = {"choices": [{"message": {"content": None, "tool_calls": [{"type": "function", "function": {}}]}}]}
|
||||
empty = {"choices": [{"message": {"content": None}}]}
|
||||
|
||||
await logger.async_log_success_event(_success_kwargs(), tool_final, None, None)
|
||||
await logger.async_log_success_event(_success_kwargs(), empty, None, None)
|
||||
await _drain(logger)
|
||||
|
||||
assert logger._test_funnel == [("job-1", "unjudgeable")]
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -314,6 +314,36 @@ def test_mask_credentials_in_payload_masks_only_sensitive_string_leaves():
|
|||
assert masked.endswith(plaintext[-4:])
|
||||
|
||||
|
||||
def test_extra_sensitive_patterns_add_to_the_defaults():
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
|
||||
|
||||
masker = SensitiveDataMasker(extra_sensitive_patterns={"connection"})
|
||||
|
||||
assert masker.is_sensitive_key("mongodb_connection_string") is True
|
||||
assert masker.is_sensitive_key("api_key") is True
|
||||
assert masker.is_sensitive_key("aws_secret_access_key") is True
|
||||
assert masker.is_sensitive_key("mongodb_database") is False
|
||||
|
||||
|
||||
def test_extra_sensitive_patterns_do_not_leak_into_other_maskers():
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
|
||||
|
||||
SensitiveDataMasker(extra_sensitive_patterns={"connection"})
|
||||
|
||||
assert SensitiveDataMasker().is_sensitive_key("mongodb_connection_string") is False
|
||||
|
||||
|
||||
def test_the_second_positional_argument_is_still_the_override_set():
|
||||
"""SensitiveDataMasker is public SDK surface, so adding a keyword must not shift what an
|
||||
existing positional call means. Putting extra_sensitive_patterns second would silently turn
|
||||
an override set into an extra sensitive set and start masking the caller's pricing fields."""
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
|
||||
|
||||
masker = SensitiveDataMasker({"token"}, {"session"})
|
||||
|
||||
assert masker.is_sensitive_key("session_token") is False
|
||||
assert masker.is_sensitive_key("auth_token") is True
|
||||
|
||||
def test_redact_credentials_in_payload_leaves_no_fragment_of_the_secret():
|
||||
"""A payload rendered straight to stdout cannot afford the partial reveal
|
||||
mask_credentials_in_payload leaves, so every credential-named value is replaced
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -170,22 +170,27 @@ class TestExtractConverseTexts:
|
|||
texts, _ = _extract_converse_texts(body, skip_system=False, skip_tool=False)
|
||||
assert texts == []
|
||||
|
||||
def test_extracts_tool_config_description_and_schema(self):
|
||||
def test_tool_config_definitions_not_extracted(self):
|
||||
"""Tool definitions are app-authored config, so nothing under
|
||||
toolConfig.tools reaches the guardrail as input content."""
|
||||
body = {
|
||||
"messages": [{"role": "user", "content": [{"text": "hi"}]}],
|
||||
"messages": [
|
||||
{"role": "user", "content": [{"text": "How much lag is there in my data?"}]}
|
||||
],
|
||||
"toolConfig": {
|
||||
"tools": [
|
||||
{
|
||||
"toolSpec": {
|
||||
"name": "lookup",
|
||||
"description": "blocked tool description",
|
||||
"description": "tool description",
|
||||
"inputSchema": {
|
||||
"json": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"q": {
|
||||
"agent_name": {
|
||||
"type": "string",
|
||||
"description": "blocked schema description",
|
||||
"title": "Agent Name",
|
||||
"enum": ["alpha", "beta", "gamma"],
|
||||
}
|
||||
},
|
||||
}
|
||||
|
|
@ -196,20 +201,56 @@ class TestExtractConverseTexts:
|
|||
},
|
||||
}
|
||||
texts, _ = _extract_converse_texts(body, skip_system=False, skip_tool=False)
|
||||
assert "blocked tool description" in texts
|
||||
assert "blocked schema description" in texts
|
||||
assert texts == ["How much lag is there in my data?"]
|
||||
|
||||
def test_tool_config_scanned_even_when_tool_messages_skipped(self):
|
||||
def test_every_tool_definition_excluded_not_just_the_first(self):
|
||||
"""A per-tool scan that only skipped tools[0] would still leak the rest."""
|
||||
body = {
|
||||
"messages": [{"role": "user", "content": [{"text": "hi"}]}],
|
||||
"toolConfig": {
|
||||
"tools": [
|
||||
{"toolSpec": {"name": "fn", "description": "blocked description"}}
|
||||
{"toolSpec": {"name": "first", "description": "first description"}},
|
||||
{"toolSpec": {"name": "second", "description": "second description"}},
|
||||
{"toolSpec": {"name": "third", "description": "third description"}},
|
||||
]
|
||||
},
|
||||
}
|
||||
texts, _ = _extract_converse_texts(body, skip_system=False, skip_tool=False)
|
||||
assert texts == ["hi"]
|
||||
|
||||
def test_tool_config_definitions_not_extracted_when_tool_messages_skipped(self):
|
||||
body = {
|
||||
"messages": [{"role": "user", "content": [{"text": "hi"}]}],
|
||||
"toolConfig": {
|
||||
"tools": [
|
||||
{"toolSpec": {"name": "fn", "description": "tool description"}}
|
||||
]
|
||||
},
|
||||
}
|
||||
texts, _ = _extract_converse_texts(body, skip_system=False, skip_tool=True)
|
||||
assert "blocked description" in texts
|
||||
assert texts == ["hi"]
|
||||
|
||||
def test_tool_use_input_still_extracted_alongside_tool_config(self):
|
||||
"""Only tool DEFINITIONS are excluded; caller content inside a toolUse
|
||||
block is still scanned."""
|
||||
body = {
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"text": "hi"},
|
||||
{"toolUse": {"toolUseId": "t1", "name": "fn", "input": {"q": "user secret"}}},
|
||||
],
|
||||
}
|
||||
],
|
||||
"toolConfig": {
|
||||
"tools": [
|
||||
{"toolSpec": {"name": "fn", "description": "tool description"}}
|
||||
]
|
||||
},
|
||||
}
|
||||
texts, _ = _extract_converse_texts(body, skip_system=False, skip_tool=False)
|
||||
assert texts == ["hi", "user secret"]
|
||||
|
||||
def test_extracts_additional_model_request_fields(self):
|
||||
body = {
|
||||
|
|
@ -437,9 +478,9 @@ class TestBedrockPassthroughGuardrailHandlerInput:
|
|||
assert "blocked content" in sent_texts
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_config_description_scanned_and_masked(self):
|
||||
"""Blocked text hidden in toolConfig.tools[].toolSpec.description is still
|
||||
forwarded to Bedrock, so the guardrail must see it and mask it in place."""
|
||||
async def test_tool_config_definitions_not_sent_and_left_untouched(self):
|
||||
"""Tool definitions never reach the guardrail, and the body forwarded to
|
||||
Bedrock keeps them byte for byte."""
|
||||
handler = BedrockPassthroughGuardrailHandler()
|
||||
data = _converse_data()
|
||||
data["data"]["toolConfig"] = {
|
||||
|
|
@ -453,36 +494,42 @@ class TestBedrockPassthroughGuardrailHandlerInput:
|
|||
}
|
||||
]
|
||||
}
|
||||
guardrail = _make_guardrail(
|
||||
{"texts": ["You are helpful.", "Hello world", "lookup", "[REDACTED]", "object"]}
|
||||
)
|
||||
original_tool_config = copy.deepcopy(data["data"]["toolConfig"])
|
||||
guardrail = _make_guardrail({"texts": ["[REDACTED]", "[REDACTED]"]})
|
||||
|
||||
result = await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
sent_texts = guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"]
|
||||
assert "email john@example.com" in sent_texts
|
||||
tool_spec = result["data"]["toolConfig"]["tools"][0]["toolSpec"]
|
||||
assert tool_spec["description"] == "[REDACTED]"
|
||||
assert sent_texts == ["You are helpful.", "Hello world"]
|
||||
assert result["data"]["toolConfig"] == original_tool_config
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_config_description_blocking_propagates(self):
|
||||
"""A blocking guardrail must reject content hidden in a tool description."""
|
||||
async def test_blocking_guardrail_not_triggered_by_tool_description(self):
|
||||
"""LIT-5797: a request whose only prompt is a benign user message must not
|
||||
be blocked because a denied term appears in a tool definition."""
|
||||
handler = BedrockPassthroughGuardrailHandler()
|
||||
data = _converse_data()
|
||||
data["data"]["toolConfig"] = {
|
||||
"tools": [{"toolSpec": {"name": "fn", "description": "blocked content"}}]
|
||||
}
|
||||
|
||||
async def _block_on_denied_term(**kwargs):
|
||||
texts = kwargs["inputs"]["texts"]
|
||||
if any("blocked content" in text for text in texts):
|
||||
raise GuardrailBlocked("Blocked")
|
||||
return {"texts": texts}
|
||||
|
||||
guardrail = MagicMock()
|
||||
guardrail.guardrail_name = "block-guard"
|
||||
guardrail.skip_system_message_in_guardrail = False
|
||||
guardrail.skip_tool_message_in_guardrail = False
|
||||
guardrail.apply_guardrail = AsyncMock(side_effect=GuardrailBlocked("Blocked"))
|
||||
guardrail.apply_guardrail = AsyncMock(side_effect=_block_on_denied_term)
|
||||
|
||||
with pytest.raises(GuardrailBlocked):
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
result = await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
sent_texts = guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"]
|
||||
assert "blocked content" in sent_texts
|
||||
assert "blocked content" not in sent_texts
|
||||
assert result["data"]["toolConfig"]["tools"][0]["toolSpec"]["description"] == "blocked content"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_additional_model_request_fields_scanned_and_masked(self):
|
||||
|
|
|
|||
|
|
@ -4,11 +4,9 @@ from unittest.mock import MagicMock, patch
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
|
||||
|
||||
from litellm import get_model_info, supports_reasoning, supports_vision
|
||||
from litellm.llms.fireworks_ai.chat.transformation import FireworksAIConfig
|
||||
from litellm.constants import SESSION_ID_GENERATED_METADATA_KEY
|
||||
from litellm.llms.fireworks_ai.chat.transformation import FireworksAIConfig
|
||||
from litellm.llms.fireworks_ai.common_utils import get_fireworks_session_id
|
||||
from litellm.types.utils import (
|
||||
ChatCompletionMessageToolCall,
|
||||
|
|
@ -363,6 +361,27 @@ def test_get_supported_openai_params_parallel_tool_calls():
|
|||
assert "parallel_tool_calls" not in unsupported_params
|
||||
|
||||
|
||||
def test_get_supported_openai_params_short_model_name_resolves_account_prefixed_entry():
|
||||
config = FireworksAIConfig()
|
||||
|
||||
supported_params = config.get_supported_openai_params(
|
||||
"fireworks_ai/deepseek-v4-pro-0813"
|
||||
)
|
||||
|
||||
assert "tool_choice" in supported_params
|
||||
assert "reasoning_effort" in supported_params
|
||||
|
||||
|
||||
def test_get_supported_openai_params_preserves_generic_reasoning_fallback():
|
||||
config = FireworksAIConfig()
|
||||
|
||||
supported_params = config.get_supported_openai_params(
|
||||
"fireworks_ai/accounts/fireworks/models/glm-5p3-flash"
|
||||
)
|
||||
|
||||
assert "reasoning_effort" in supported_params
|
||||
|
||||
|
||||
def test_get_supported_openai_params_parallel_tool_calls_without_tool_choice(
|
||||
monkeypatch,
|
||||
):
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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
|
||||
|
|
@ -1053,6 +1053,8 @@ class TestNumericFormFields:
|
|||
read_only: ReadOnly[int | None]
|
||||
not_required: NotRequired[ReadOnly[int]]
|
||||
required: Required[ReadOnly[Annotated[float, "meta"]]]
|
||||
read_only_not_required: ReadOnly[NotRequired[int]]
|
||||
read_only_required: ReadOnly[Required[float]]
|
||||
|
||||
assert dict(numeric_form_fields(get_type_hints(Schema))) == {
|
||||
"plain": int,
|
||||
|
|
@ -1061,6 +1063,22 @@ class TestNumericFormFields:
|
|||
"read_only": int,
|
||||
"not_required": int,
|
||||
"required": float,
|
||||
"read_only_not_required": int,
|
||||
"read_only_required": float,
|
||||
}
|
||||
|
||||
def test_qualifiers_are_unwrapped_when_get_type_hints_keeps_extras(self):
|
||||
from typing_extensions import Annotated, NotRequired, ReadOnly, Required, TypedDict
|
||||
|
||||
class Schema(TypedDict, total=False):
|
||||
annotated: ReadOnly[Annotated[int, "meta"]]
|
||||
not_required: NotRequired[ReadOnly[int]]
|
||||
required: Required[ReadOnly[Annotated[float, "meta"]]]
|
||||
|
||||
assert dict(numeric_form_fields(get_type_hints(Schema, include_extras=True))) == {
|
||||
"annotated": int,
|
||||
"not_required": int,
|
||||
"required": float,
|
||||
}
|
||||
|
||||
def test_non_scalar_and_bool_fields_are_skipped(self):
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import sys
|
|||
import types
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from datetime import time as dt_time
|
||||
from typing import Any, Dict, List
|
||||
from typing import Any, Dict, Final, List
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import httpx
|
||||
|
|
@ -1495,6 +1495,32 @@ def test_budget_table_reset_invalidates_every_tag_not_just_the_first(reset_budge
|
|||
assert deleted == {"tag:tenant-a", "tag:tenant-b", "tag:tenant-c"}
|
||||
|
||||
|
||||
def test_budget_table_reset_invalidates_enduser_counter_and_cache(reset_budget_job, mock_prisma_client, monkeypatch):
|
||||
"""When an end user's budget resets, its Redis spend counter is zeroed and its management cache is evicted."""
|
||||
counter_cache: Final = _make_counter_invalidation_job(monkeypatch)
|
||||
budget: Final = _budget_row(budget_id="budget-1")
|
||||
mock_prisma_client.data["budget"] = [budget]
|
||||
test_enduser: Final = type(
|
||||
"LiteLLM_EndUserTable",
|
||||
(),
|
||||
{
|
||||
"spend": 20.0,
|
||||
"litellm_budget_table": budget,
|
||||
"budget_id": "budget-1",
|
||||
"user_id": "customer-42",
|
||||
},
|
||||
)
|
||||
mock_prisma_client.data["enduser"] = [test_enduser]
|
||||
|
||||
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
|
||||
|
||||
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:end_user:customer-42", value=0.0, ttl=60)
|
||||
counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:end_user:customer-42", value=0.0, ttl=60)
|
||||
deleted: Final = {call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list}
|
||||
assert "end_user_id:customer-42" in deleted
|
||||
|
||||
|
||||
|
||||
def test_budget_table_reset_commits_even_when_cache_eviction_fails(reset_budget_job, mock_prisma_client, monkeypatch):
|
||||
"""Eviction runs after the commit, so a broken cache cannot undo the write."""
|
||||
counter_cache = _make_counter_invalidation_job(monkeypatch)
|
||||
|
|
@ -3028,6 +3054,38 @@ def test_budget_cascade_carries_enduser_overage_when_rollover_enabled(
|
|||
} in enduser_writes
|
||||
|
||||
|
||||
def test_budget_cascade_carries_default_tier_enduser_counter_when_rollover_enabled(
|
||||
rollover_enabled, reset_budget_job, mock_prisma_client, monkeypatch
|
||||
):
|
||||
"""An end user on the default budget (no budget_id on its row) 5 over the cap
|
||||
keeps a counter of 5 in the next window and loses its cached object."""
|
||||
import litellm
|
||||
|
||||
counter_cache: Final = _make_counter_invalidation_job(monkeypatch)
|
||||
monkeypatch.setattr(litellm, "max_end_user_budget_id", "default-enduser-budget")
|
||||
mock_prisma_client.data["budget"] = [
|
||||
_budget_row(budget_id="default-enduser-budget", budget_duration="1d", max_budget=10.0)
|
||||
]
|
||||
implicit_enduser: Final = type(
|
||||
"EndUserRow",
|
||||
(),
|
||||
{
|
||||
"spend": 15.0,
|
||||
"user_id": "enduser-implicit",
|
||||
"budget_id": None,
|
||||
"model_dump": lambda self=None: {"spend": 15.0, "user_id": "enduser-implicit", "budget_id": None, "blocked": False},
|
||||
},
|
||||
)
|
||||
mock_prisma_client.db.litellm_endusertable.set_find_many_results([implicit_enduser])
|
||||
|
||||
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
|
||||
|
||||
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:end_user:enduser-implicit", value=5.0, ttl=60)
|
||||
counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:end_user:enduser-implicit", value=5.0, ttl=60)
|
||||
deleted: Final = {call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list}
|
||||
assert "end_user_id:enduser-implicit" in deleted
|
||||
|
||||
|
||||
def _replay_spend_writes(writes, spend):
|
||||
"""Apply the queued update_many statements in order, the way the DB
|
||||
transaction executes them, and return the row's final spend."""
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from __future__ import annotations
|
|||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -18,7 +19,7 @@ from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed
|
|||
WINDOW_START = datetime(2026, 8, 1, tzinfo=timezone.utc)
|
||||
|
||||
|
||||
class _FakeWindowSpendTable:
|
||||
class _FakeFindUniqueTable:
|
||||
def __init__(self, row: SimpleNamespace | None, error: Exception | None = None) -> None:
|
||||
self._row = row
|
||||
self._error = error
|
||||
|
|
@ -47,10 +48,13 @@ class _FakePrismaClient:
|
|||
row: SimpleNamespace | None = None,
|
||||
spend_logs_total: float = 0.0,
|
||||
error: Exception | None = None,
|
||||
end_user_row: SimpleNamespace | None = None,
|
||||
end_user_error: Exception | None = None,
|
||||
) -> None:
|
||||
self.db = SimpleNamespace(
|
||||
litellm_budgetwindowspend=_FakeWindowSpendTable(row=row, error=error),
|
||||
litellm_budgetwindowspend=_FakeFindUniqueTable(row=row, error=error),
|
||||
litellm_spendlogs=_FakeSpendLogsTable(total=spend_logs_total),
|
||||
litellm_endusertable=_FakeFindUniqueTable(row=end_user_row, error=end_user_error),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -248,3 +252,65 @@ async def test_coalesced_window_seeds_a_cold_counter_from_the_row():
|
|||
assert result == 4.5
|
||||
assert cache.in_memory_cache.get_cache(key=counter_key) == 4.5
|
||||
assert prisma.db.litellm_spendlogs.call_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_end_user_from_db_reads_the_end_user_row_by_user_id():
|
||||
prisma: Final = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="customer-42", spend=0.0))
|
||||
|
||||
result: Final = await SpendCounterReseed.end_user_from_db(
|
||||
prisma_client=prisma, counter_key="spend:end_user:customer-42"
|
||||
)
|
||||
|
||||
assert result == 0.0
|
||||
assert prisma.db.litellm_endusertable.where_clauses == [{"user_id": "customer-42"}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_end_user_from_db_returns_the_recorded_spend():
|
||||
prisma: Final = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="customer-42", spend=12.5))
|
||||
|
||||
assert (
|
||||
await SpendCounterReseed.end_user_from_db(prisma_client=prisma, counter_key="spend:end_user:customer-42")
|
||||
== 12.5
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("counter_key", ["spend:key:hashed", "spend:team:t1", "spend:tag:t1"])
|
||||
async def test_end_user_from_db_ignores_other_counter_kinds_without_touching_the_db(counter_key):
|
||||
prisma: Final = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="x", spend=5.0))
|
||||
|
||||
assert await SpendCounterReseed.end_user_from_db(prisma_client=prisma, counter_key=counter_key) is None
|
||||
assert prisma.db.litellm_endusertable.where_clauses == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_end_user_from_db_returns_none_without_a_row_a_client_or_on_db_error():
|
||||
assert (
|
||||
await SpendCounterReseed.end_user_from_db(prisma_client=None, counter_key="spend:end_user:customer-42")
|
||||
is None
|
||||
)
|
||||
assert (
|
||||
await SpendCounterReseed.end_user_from_db(
|
||||
prisma_client=_FakePrismaClient(end_user_row=None), counter_key="spend:end_user:customer-42"
|
||||
)
|
||||
is None
|
||||
)
|
||||
assert (
|
||||
await SpendCounterReseed.end_user_from_db(
|
||||
prisma_client=_FakePrismaClient(end_user_error=RuntimeError("db down")),
|
||||
counter_key="spend:end_user:customer-42",
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_from_db_still_never_reads_the_end_user_row():
|
||||
"""A cold end-user counter keeps seeding from the cached end-user object the auth
|
||||
path already loaded; the row is read only as the budget floor."""
|
||||
prisma: Final = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="customer-42", spend=5.0))
|
||||
|
||||
assert await SpendCounterReseed.from_db(prisma_client=prisma, counter_key="spend:end_user:customer-42") is None
|
||||
assert prisma.db.litellm_endusertable.where_clauses == []
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import asyncio
|
||||
import json
|
||||
import time
|
||||
from typing import Final
|
||||
from datetime import datetime, timedelta
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
@ -1190,27 +1191,23 @@ def test_health_liveliness_endpoint(proxy_client):
|
|||
Test that /health/liveliness endpoint returns 200 OK with "I'm alive!" message.
|
||||
This is a critical orchestration endpoint that must be simple and fast.
|
||||
"""
|
||||
# Measure the time taken for the health check call
|
||||
start_time = time.perf_counter()
|
||||
warm_up: Final = proxy_client.get("/health/liveliness")
|
||||
assert warm_up.status_code == 200, f"Expected 200 OK, got {warm_up.status_code}: {warm_up.text}"
|
||||
|
||||
# Make GET request to /health/liveliness
|
||||
response = proxy_client.get("/health/liveliness")
|
||||
def _timed_poll() -> tuple[float, httpx.Response]:
|
||||
start_time: Final = time.perf_counter()
|
||||
response: Final = proxy_client.get("/health/liveliness")
|
||||
return (time.perf_counter() - start_time) * 1000, response
|
||||
|
||||
end_time = time.perf_counter()
|
||||
duration_ms = (end_time - start_time) * 1000
|
||||
polls: Final = tuple(_timed_poll() for _ in range(5))
|
||||
|
||||
# Assert response status
|
||||
assert response.status_code == 200, f"Expected 200 OK, got {response.status_code}: {response.text}"
|
||||
for _, response in polls:
|
||||
assert response.status_code == 200, f"Expected 200 OK, got {response.status_code}: {response.text}"
|
||||
assert response.json() == "I'm alive!", f"Expected 'I'm alive!' message, got: {response.json()}"
|
||||
|
||||
# Assert response content (FastAPI JSON-encodes the string)
|
||||
assert response.json() == "I'm alive!", f"Expected 'I'm alive!' message, got: {response.json()}"
|
||||
|
||||
# Verify response is fast (should be < 100ms for a simple endpoint)
|
||||
# This is critical for orchestration systems that poll frequently
|
||||
assert duration_ms < 100, f"Health check took {duration_ms:.2f}ms, expected < 100ms for a simple endpoint"
|
||||
|
||||
# Log the duration for visibility (useful for CI/CD monitoring)
|
||||
print(f"\n/health/liveliness response time: {duration_ms:.2f}ms")
|
||||
durations_ms: Final = tuple(sorted(duration_ms for duration_ms, _ in polls))
|
||||
median_ms: Final = durations_ms[len(durations_ms) // 2]
|
||||
assert median_ms < 100, f"Median of {len(polls)} health checks took {median_ms:.2f}ms, expected < 100ms"
|
||||
|
||||
|
||||
def test_health_liveness_endpoint(proxy_client):
|
||||
|
|
|
|||
|
|
@ -57,8 +57,20 @@ def _make_access_group_record(
|
|||
return record
|
||||
|
||||
|
||||
def _make_team_record(team_id: str, access_group_ids: list[str] | None = None):
|
||||
return types.SimpleNamespace(team_id=team_id, access_group_ids=access_group_ids or [])
|
||||
def _make_team_record(team_id: str, access_group_ids: list[str] | None = None, team_alias: str | None = None):
|
||||
return types.SimpleNamespace(team_id=team_id, access_group_ids=access_group_ids or [], team_alias=team_alias)
|
||||
|
||||
|
||||
def _make_mcp_server_record(server_id: str, alias: str | None = None, server_name: str | None = None):
|
||||
return types.SimpleNamespace(server_id=server_id, alias=alias, server_name=server_name)
|
||||
|
||||
|
||||
def _make_agent_record(agent_id: str, agent_name: str):
|
||||
return types.SimpleNamespace(agent_id=agent_id, agent_name=agent_name)
|
||||
|
||||
|
||||
def _make_key_record(token: str, key_alias: str | None = None):
|
||||
return types.SimpleNamespace(token=token, key_alias=key_alias)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
|
|
@ -109,6 +121,12 @@ def client_and_mocks(monkeypatch):
|
|||
mock_key_table.find_unique = AsyncMock(return_value=None)
|
||||
mock_key_table.update = AsyncMock(return_value=None)
|
||||
|
||||
mock_mcp_server_table = MagicMock()
|
||||
mock_mcp_server_table.find_many = AsyncMock(return_value=[])
|
||||
|
||||
mock_agents_table = MagicMock()
|
||||
mock_agents_table.find_many = AsyncMock(return_value=[])
|
||||
|
||||
@asynccontextmanager
|
||||
async def mock_tx():
|
||||
tx = types.SimpleNamespace(
|
||||
|
|
@ -122,6 +140,8 @@ def client_and_mocks(monkeypatch):
|
|||
litellm_accessgrouptable=mock_access_group_table,
|
||||
litellm_teamtable=mock_team_table,
|
||||
litellm_verificationtoken=mock_key_table,
|
||||
litellm_mcpservertable=mock_mcp_server_table,
|
||||
litellm_agentstable=mock_agents_table,
|
||||
tx=mock_tx,
|
||||
)
|
||||
mock_prisma.db = mock_db
|
||||
|
|
@ -1447,3 +1467,169 @@ def test_update_access_group_null_assigned_ids_treated_as_empty(client_and_mocks
|
|||
update_call_kwargs = mock_table.update.call_args.kwargs
|
||||
assert update_call_kwargs["data"]["assigned_team_ids"] == []
|
||||
assert update_call_kwargs["data"]["assigned_key_ids"] == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Resolved resource names (LIT-6594)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _mock_resource_tables(mock_prisma, *, mcp_servers=(), agents=(), teams=(), keys=()):
|
||||
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=list(mcp_servers))
|
||||
mock_prisma.db.litellm_agentstable.find_many = AsyncMock(return_value=list(agents))
|
||||
mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=list(teams))
|
||||
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=list(keys))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("base_path", ACCESS_GROUP_PATHS)
|
||||
def test_get_access_group_resolves_resource_names(client_and_mocks, base_path):
|
||||
"""Every id list gets a sibling list of {id, name}; name is null when the id has no alias or no longer resolves."""
|
||||
client, mock_prisma, mock_table, *_ = client_and_mocks
|
||||
mock_table.find_unique = AsyncMock(
|
||||
return_value=_make_access_group_record(
|
||||
access_group_id="ag-123",
|
||||
access_mcp_server_ids=["mcp-a", "mcp-b", "mcp-ghost"],
|
||||
access_agent_ids=["agent-a", "agent-ghost"],
|
||||
assigned_team_ids=["team-a", "team-b"],
|
||||
assigned_key_ids=["key-a", "key-b"],
|
||||
)
|
||||
)
|
||||
_mock_resource_tables(
|
||||
mock_prisma,
|
||||
mcp_servers=[
|
||||
_make_mcp_server_record("mcp-a", alias="GitHub"),
|
||||
_make_mcp_server_record("mcp-b", server_name="jira_tools"),
|
||||
],
|
||||
agents=[_make_agent_record("agent-a", "support-bot")],
|
||||
teams=[
|
||||
_make_team_record("team-a", ["ag-123"], team_alias="Platform"),
|
||||
_make_team_record("team-b", ["ag-123"]),
|
||||
],
|
||||
keys=[_make_key_record("key-a", key_alias="ci-key"), _make_key_record("key-b")],
|
||||
)
|
||||
|
||||
resp = client.get(f"{base_path}/ag-123")
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert body["access_mcp_servers"] == [
|
||||
{"id": "mcp-a", "name": "GitHub"},
|
||||
{"id": "mcp-b", "name": "jira_tools"},
|
||||
{"id": "mcp-ghost", "name": None},
|
||||
]
|
||||
assert body["access_agents"] == [{"id": "agent-a", "name": "support-bot"}, {"id": "agent-ghost", "name": None}]
|
||||
assert body["assigned_teams"] == [{"id": "team-a", "name": "Platform"}, {"id": "team-b", "name": None}]
|
||||
assert body["assigned_keys"] == [{"id": "key-a", "name": "ci-key"}, {"id": "key-b", "name": None}]
|
||||
assert body["access_mcp_server_ids"] == ["mcp-a", "mcp-b", "mcp-ghost"]
|
||||
assert body["assigned_team_ids"] == ["team-a", "team-b"]
|
||||
|
||||
mcp_where = mock_prisma.db.litellm_mcpservertable.find_many.call_args.kwargs["where"]
|
||||
assert sorted(mcp_where["server_id"]["in"]) == ["mcp-a", "mcp-b", "mcp-ghost"]
|
||||
agent_where = mock_prisma.db.litellm_agentstable.find_many.call_args.kwargs["where"]
|
||||
assert sorted(agent_where["agent_id"]["in"]) == ["agent-a", "agent-ghost"]
|
||||
key_where = mock_prisma.db.litellm_verificationtoken.find_many.call_args.kwargs["where"]
|
||||
assert sorted(key_where["token"]["in"]) == ["key-a", "key-b"]
|
||||
|
||||
|
||||
def test_list_access_groups_resolves_names_with_one_query_per_table(client_and_mocks):
|
||||
"""List batches every group's ids into one lookup per table and attributes names back to the right group."""
|
||||
client, mock_prisma, mock_table, *_ = client_and_mocks
|
||||
mock_table.find_many = AsyncMock(
|
||||
return_value=[
|
||||
_make_access_group_record(
|
||||
access_group_id="ag-1", access_mcp_server_ids=["mcp-a"], access_agent_ids=["agent-a"], assigned_key_ids=["key-a"]
|
||||
),
|
||||
_make_access_group_record(
|
||||
access_group_id="ag-2", access_mcp_server_ids=["mcp-b"], access_agent_ids=["agent-b"], assigned_key_ids=["key-b"]
|
||||
),
|
||||
]
|
||||
)
|
||||
_mock_resource_tables(
|
||||
mock_prisma,
|
||||
mcp_servers=[_make_mcp_server_record("mcp-a", alias="A"), _make_mcp_server_record("mcp-b", alias="B")],
|
||||
agents=[_make_agent_record("agent-a", "Agent A"), _make_agent_record("agent-b", "Agent B")],
|
||||
keys=[_make_key_record("key-a", key_alias="Key A"), _make_key_record("key-b", key_alias="Key B")],
|
||||
)
|
||||
|
||||
resp = client.get("/v1/access_group")
|
||||
assert resp.status_code == 200
|
||||
first, second = resp.json()
|
||||
assert first["access_mcp_servers"] == [{"id": "mcp-a", "name": "A"}]
|
||||
assert first["access_agents"] == [{"id": "agent-a", "name": "Agent A"}]
|
||||
assert first["assigned_keys"] == [{"id": "key-a", "name": "Key A"}]
|
||||
assert second["access_mcp_servers"] == [{"id": "mcp-b", "name": "B"}]
|
||||
assert second["access_agents"] == [{"id": "agent-b", "name": "Agent B"}]
|
||||
assert second["assigned_keys"] == [{"id": "key-b", "name": "Key B"}]
|
||||
|
||||
for table, column in (
|
||||
(mock_prisma.db.litellm_mcpservertable, "server_id"),
|
||||
(mock_prisma.db.litellm_agentstable, "agent_id"),
|
||||
(mock_prisma.db.litellm_verificationtoken, "token"),
|
||||
):
|
||||
table.find_many.assert_awaited_once()
|
||||
assert len(table.find_many.call_args.kwargs["where"][column]["in"]) == 2
|
||||
|
||||
|
||||
def test_list_access_groups_skips_lookups_when_nothing_to_resolve(client_and_mocks):
|
||||
"""Groups with no MCP servers, agents or keys must not trigger an empty IN () query per table."""
|
||||
client, mock_prisma, mock_table, *_ = client_and_mocks
|
||||
mock_table.find_many = AsyncMock(
|
||||
return_value=[_make_access_group_record(access_group_id="ag-1"), _make_access_group_record(access_group_id="ag-2")]
|
||||
)
|
||||
|
||||
resp = client.get("/v1/access_group")
|
||||
assert resp.status_code == 200
|
||||
assert all(group["access_mcp_servers"] == [] and group["assigned_keys"] == [] for group in resp.json())
|
||||
|
||||
mock_prisma.db.litellm_mcpservertable.find_many.assert_not_awaited()
|
||||
mock_prisma.db.litellm_agentstable.find_many.assert_not_awaited()
|
||||
mock_prisma.db.litellm_verificationtoken.find_many.assert_not_awaited()
|
||||
|
||||
|
||||
def test_create_access_group_response_carries_resolved_names(client_and_mocks):
|
||||
"""The create response already shows names so the UI never has to refetch to label what it just saved."""
|
||||
client, mock_prisma, *_ = client_and_mocks
|
||||
team_record = _make_team_record("team-1", team_alias="Platform")
|
||||
mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_record)
|
||||
_mock_resource_tables(
|
||||
mock_prisma,
|
||||
mcp_servers=[_make_mcp_server_record("mcp-a", alias="GitHub")],
|
||||
agents=[_make_agent_record("agent-a", "support-bot")],
|
||||
teams=[team_record],
|
||||
)
|
||||
|
||||
resp = client.post(
|
||||
"/v1/access_group",
|
||||
json={
|
||||
"access_group_name": "new-group",
|
||||
"access_mcp_server_ids": ["mcp-a"],
|
||||
"access_agent_ids": ["agent-a"],
|
||||
"assigned_team_ids": ["team-1"],
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 201
|
||||
body = resp.json()
|
||||
assert body["access_mcp_servers"] == [{"id": "mcp-a", "name": "GitHub"}]
|
||||
assert body["access_agents"] == [{"id": "agent-a", "name": "support-bot"}]
|
||||
assert body["assigned_teams"] == [{"id": "team-1", "name": "Platform"}]
|
||||
|
||||
|
||||
def test_update_access_group_response_carries_resolved_names(client_and_mocks):
|
||||
"""The update response reflects the new ids with their names, not the pre-update state."""
|
||||
client, mock_prisma, mock_table, *_ = client_and_mocks
|
||||
mock_table.find_unique = AsyncMock(
|
||||
return_value=_make_access_group_record(access_group_id="ag-update", access_mcp_server_ids=["mcp-old"])
|
||||
)
|
||||
_mock_resource_tables(
|
||||
mock_prisma,
|
||||
mcp_servers=[_make_mcp_server_record("mcp-new", alias="Linear")],
|
||||
agents=[_make_agent_record("agent-a", "support-bot")],
|
||||
)
|
||||
|
||||
resp = client.put(
|
||||
"/v1/access_group/ag-update", json={"access_mcp_server_ids": ["mcp-new"], "access_agent_ids": ["agent-a"]}
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert body["access_mcp_servers"] == [{"id": "mcp-new", "name": "Linear"}]
|
||||
assert body["access_agents"] == [{"id": "agent-a", "name": "support-bot"}]
|
||||
assert body["access_mcp_server_ids"] == ["mcp-new"]
|
||||
|
|
|
|||
|
|
@ -882,6 +882,7 @@ def _leg_record(**overrides: object) -> MagicMock:
|
|||
"target_id": "key-hash",
|
||||
"router_name": "my-router",
|
||||
"router_names": (),
|
||||
"models": (),
|
||||
"direction": "forward",
|
||||
"baseline_model": None,
|
||||
"judge_model": "anthropic/claude-sonnet-5",
|
||||
|
|
@ -1033,6 +1034,7 @@ def _shadow_prisma(
|
|||
"target_id",
|
||||
"router_name",
|
||||
"router_names",
|
||||
"models",
|
||||
"direction",
|
||||
"baseline_model",
|
||||
"judge_model",
|
||||
|
|
@ -1533,6 +1535,87 @@ async def test_start_shadow_eval_forward_leaves_the_baseline_column_empty(monkey
|
|||
assert rows[0]["baseline_model"] is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_start_shadow_eval_writes_the_model_scope_on_every_leg_and_echoes_it(monkeypatch: pytest.MonkeyPatch):
|
||||
"""A model scope is job config, so every leg carries the same copy and both the start
|
||||
response and a later list read report it; an auto-router is a legitimate scope (a
|
||||
forward job on one router may sample what another router serves today)."""
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
|
||||
_configure_anthropic_sdk_judge(monkeypatch)
|
||||
prisma = _shadow_prisma()
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", _shadow_router())
|
||||
|
||||
response = await start_shadow_eval(
|
||||
_start_request(api_key_ids=("key-hash", "key-hash-2"), models=("cheap", "sonnet-router")), ADMIN
|
||||
)
|
||||
|
||||
rows = prisma.db.litellm_shadowevaljob.create_many.call_args.kwargs["data"]
|
||||
assert [row["models"] for row in rows] == [["cheap", "sonnet-router"], ["cheap", "sonnet-router"]]
|
||||
assert response.models == ("cheap", "sonnet-router")
|
||||
|
||||
listed = _shadow_prisma(legs=[_leg_record(models=("cheap",)), _leg_record(id="leg-0", group_id="job-0")])
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", listed)
|
||||
jobs = await list_shadow_eval_jobs(VIEWER, target_type=None, target_id=None, limit=50)
|
||||
assert {job.job_id: job.models for job in jobs} == {"job-1": ("cheap",), "job-0": ()}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_start_shadow_eval_accepts_a_team_public_scope_for_a_user_target(monkeypatch: pytest.MonkeyPatch):
|
||||
"""A user's traffic can arrive on any team's key, so a name only one team can ask for
|
||||
is a legitimate scope for a user target even though it resolves for nobody unscoped."""
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
|
||||
_configure_anthropic_sdk_judge(monkeypatch)
|
||||
prisma = _shadow_prisma(known_users={"dev-alice": "alice@example.com"})
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", _shadow_router())
|
||||
|
||||
response = await start_shadow_eval(
|
||||
_start_request(api_key_ids=(), user_ids=("dev-alice",), models=("house-judge",)), ADMIN
|
||||
)
|
||||
|
||||
assert response.models == ("house-judge",)
|
||||
assert prisma.db.litellm_shadowevaljob.create_many.call_args.kwargs["data"][0]["models"] == ["house-judge"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_start_shadow_eval_rejects_a_model_scope_this_proxy_does_not_serve(monkeypatch: pytest.MonkeyPatch):
|
||||
"""A typo'd model name would otherwise start a job that samples nothing. Only the
|
||||
unresolvable names are reported, so the caller fixes them in one round."""
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
|
||||
_configure_anthropic_sdk_judge(monkeypatch)
|
||||
prisma = _shadow_prisma()
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", _shadow_router())
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await start_shadow_eval(_start_request(models=("cheap", "no-such-model-zzz")), ADMIN)
|
||||
assert exc.value.status_code == 400
|
||||
assert "'no-such-model-zzz'" in exc.value.detail
|
||||
assert "'cheap'" not in exc.value.detail
|
||||
prisma.db.litellm_shadowevaljob.create_many.assert_not_called()
|
||||
|
||||
|
||||
def test_start_request_dedupes_the_model_scope_and_rejects_blank_names():
|
||||
assert _start_request(models=("cheap", "mid", "cheap")).models == ("cheap", "mid")
|
||||
assert _start_request().models == ()
|
||||
with pytest.raises(ValidationError, match="non-empty model group names"):
|
||||
_start_request(models=("cheap", " "))
|
||||
|
||||
|
||||
def test_start_request_rejects_a_model_scope_on_a_reverse_job():
|
||||
"""Reverse admission is the router's own traffic, whose requested group is always the
|
||||
router, so a plain-model scope would sample nothing and the router itself is a no-op."""
|
||||
with pytest.raises(ValidationError, match="only meaningful for a forward job"):
|
||||
_start_request(direction="reverse", baseline_model="cheap", models=("mid",))
|
||||
with pytest.raises(ValidationError, match="only meaningful for a forward job"):
|
||||
_start_request(direction="reverse", baseline_model="cheap", models=("my-router",))
|
||||
assert _start_request(direction="reverse", baseline_model="cheap").models == ()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_start_shadow_eval_rejects_keys_this_proxy_does_not_know(monkeypatch: pytest.MonkeyPatch):
|
||||
"""A typo'd api_key_id would otherwise create a leg no traffic can ever match. Every
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -22,6 +22,7 @@ from __future__ import annotations
|
|||
|
||||
import asyncio
|
||||
from datetime import datetime
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
|
@ -222,16 +223,19 @@ async def test_get_current_spend_floor_caches_db_read(monkeypatch):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_current_spend_floors_end_user_tag_against_fallback(monkeypatch):
|
||||
"""End-user and tag counters have no DB row (from_db returns None). When the
|
||||
counter is stale-low, enforcement falls back to the caller's recorded spend
|
||||
(loaded fresh in auth) instead of trusting the stale counter."""
|
||||
fake_cache = _make_spend_counter_cache(redis_get_value=2.0)
|
||||
@pytest.mark.parametrize("counter_key", ("spend:end_user:e1", "spend:tag:t1"))
|
||||
async def test_get_current_spend_floors_end_user_tag_against_fallback(monkeypatch, counter_key):
|
||||
"""Tag counters have no DB row (from_db returns None), and an end-user counter has
|
||||
none to read without a DB client. When such a counter is stale-low, enforcement
|
||||
falls back to the caller's recorded spend (loaded fresh in auth) instead of
|
||||
trusting the stale counter."""
|
||||
fake_cache: Final = _make_spend_counter_cache(redis_get_value=2.0)
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
monkeypatch.setattr(ps, "prisma_client", None)
|
||||
monkeypatch.setattr(ps.SpendCounterReseed, "from_db", AsyncMock(return_value=None))
|
||||
|
||||
result = await ps.get_current_spend(
|
||||
counter_key="spend:end_user:e1",
|
||||
result: Final = await ps.get_current_spend(
|
||||
counter_key=counter_key,
|
||||
fallback_spend=20.0,
|
||||
max_budget=10.0,
|
||||
)
|
||||
|
|
@ -241,6 +245,72 @@ async def test_get_current_spend_floors_end_user_tag_against_fallback(monkeypatc
|
|||
fake_cache.redis_cache.async_set_max.assert_not_called()
|
||||
|
||||
|
||||
def _make_prisma_with_end_user_row(spend: float | None):
|
||||
prisma: Final = MagicMock()
|
||||
prisma.db.litellm_endusertable.find_unique = AsyncMock(
|
||||
return_value=None if spend is None else MagicMock(spend=spend)
|
||||
)
|
||||
return prisma
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_current_spend_end_user_floor_admits_after_a_reset_on_a_stale_worker(monkeypatch):
|
||||
"""The reset job zeroes LiteLLM_EndUserTable.spend and the shared counter, but it
|
||||
evicts the cached end-user object only on the worker that ran the reset. Every
|
||||
other worker still passes the pre-reset spend as fallback_spend, and that stale
|
||||
copy must not out-vote the reset row."""
|
||||
fake_cache: Final = _make_spend_counter_cache(redis_get_value=0.0)
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
prisma: Final = _make_prisma_with_end_user_row(spend=0.0)
|
||||
monkeypatch.setattr(ps, "prisma_client", prisma)
|
||||
|
||||
result = await ps.get_current_spend(
|
||||
counter_key="spend:end_user:customer-42",
|
||||
fallback_spend=0.000032,
|
||||
max_budget=0.00003,
|
||||
fallback_authoritative=True,
|
||||
)
|
||||
|
||||
assert result == 0.0
|
||||
prisma.db.litellm_endusertable.find_unique.assert_awaited_once_with(where={"user_id": "customer-42"})
|
||||
fake_cache.redis_cache.async_set_max.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_current_spend_end_user_floor_repairs_a_stale_low_counter(monkeypatch):
|
||||
"""After a Redis restart the end-user counter can sit below the recorded spend;
|
||||
the row wins and the shared counter is raised so other workers stop admitting on
|
||||
the stale value."""
|
||||
fake_cache: Final = _make_spend_counter_cache(redis_get_value=2.0)
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
monkeypatch.setattr(ps, "prisma_client", _make_prisma_with_end_user_row(spend=12.0))
|
||||
|
||||
result: Final = await ps.get_current_spend(
|
||||
counter_key="spend:end_user:customer-42",
|
||||
fallback_spend=12.0,
|
||||
max_budget=10.0,
|
||||
)
|
||||
|
||||
assert result == 12.0
|
||||
fake_cache.redis_cache.async_set_max.assert_awaited_once_with(key="spend:end_user:customer-42", value=12.0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_current_spend_end_user_without_a_row_keeps_the_cached_spend(monkeypatch):
|
||||
fake_cache: Final = _make_spend_counter_cache(redis_get_value=0.0)
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
monkeypatch.setattr(ps, "prisma_client", _make_prisma_with_end_user_row(spend=None))
|
||||
|
||||
result: Final = await ps.get_current_spend(
|
||||
counter_key="spend:end_user:customer-42",
|
||||
fallback_spend=20.0,
|
||||
max_budget=10.0,
|
||||
)
|
||||
|
||||
assert result == 20.0
|
||||
fake_cache.redis_cache.async_set_max.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_current_spend_floors_window_against_spend_logs(monkeypatch):
|
||||
"""Per-window counters have no DB row but aggregate from spend logs. A
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import json
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
|
@ -2463,6 +2464,41 @@ class TestRedactSensitiveLitellmParams:
|
|||
for k, v in params.items():
|
||||
assert out[k] == v, f"{k} should be preserved verbatim"
|
||||
|
||||
def test_redacts_wire_protocol_connection_strings(self):
|
||||
"""
|
||||
A MongoDB vector store's whole credential is its connection string:
|
||||
``mongodb+srv://<user>:<password>@<cluster>`` embeds the database
|
||||
password, and none of the default api_key/secret/token patterns match
|
||||
the key name, so an unextended masker returns it verbatim to every
|
||||
caller of /vector_store/list and /vector_store/info.
|
||||
"""
|
||||
from litellm.constants import REDACTED_BY_LITELM_STRING
|
||||
from litellm.proxy.vector_store_endpoints.management_endpoints import (
|
||||
_redact_sensitive_litellm_params,
|
||||
)
|
||||
|
||||
password = "hunter2-not-for-callers"
|
||||
params = {
|
||||
"mongodb_connection_string": f"mongodb+srv://dbuser:{password}@cluster0.mongodb.net",
|
||||
"mongodb_database": "sample_mflix",
|
||||
"mongodb_collection": "embedded_movies",
|
||||
"mongodb_embedding_field": "plot_embedding",
|
||||
"mongodb_text_field": "plot",
|
||||
"litellm_embedding_model": "openai/text-embedding-ada-002",
|
||||
}
|
||||
out = _redact_sensitive_litellm_params(params)
|
||||
|
||||
assert out["mongodb_connection_string"] == REDACTED_BY_LITELM_STRING
|
||||
assert password not in json.dumps(out)
|
||||
for k in (
|
||||
"mongodb_database",
|
||||
"mongodb_collection",
|
||||
"mongodb_embedding_field",
|
||||
"mongodb_text_field",
|
||||
"litellm_embedding_model",
|
||||
):
|
||||
assert out[k] == params[k], f"{k} is not a credential and must survive redaction"
|
||||
|
||||
def test_handles_none_and_empty(self):
|
||||
from litellm.proxy.vector_store_endpoints.management_endpoints import (
|
||||
_redact_sensitive_litellm_params,
|
||||
|
|
|
|||
|
|
@ -2897,9 +2897,7 @@ class TestRouterPreRoutingAliasOverrides:
|
|||
import time
|
||||
|
||||
monkeypatch.setenv("GITHUB_COPILOT_TOKEN_DIR", str(tmp_path))
|
||||
(tmp_path / "api-key.json").write_text(
|
||||
json.dumps({"token": "tid=test", "expires_at": int(time.time()) + 3600})
|
||||
)
|
||||
(tmp_path / "api-key.json").write_text(json.dumps({"token": "tid=test", "expires_at": int(time.time()) + 3600}))
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
|
|
@ -2924,7 +2922,9 @@ class TestRouterPreRoutingAliasOverrides:
|
|||
copilot_resolutions: List = []
|
||||
|
||||
def _guarded(*args, **kwargs):
|
||||
target = str(kwargs.get("model") or (args[0] if args else "")) + str(kwargs.get("custom_llm_provider") or "")
|
||||
target = str(kwargs.get("model") or (args[0] if args else "")) + str(
|
||||
kwargs.get("custom_llm_provider") or ""
|
||||
)
|
||||
if "github_copilot" in target:
|
||||
copilot_resolutions.append(target)
|
||||
raise RuntimeError("routing must not resolve an authenticating provider")
|
||||
|
|
@ -6133,6 +6133,150 @@ class TestEscalationKeywords:
|
|||
assert result.model == "o1-b" # unchanged: no random hop to o1-a / o1-c
|
||||
|
||||
|
||||
def _stalled_tool_history(repeats: int = 3) -> List[Dict]:
|
||||
"""`repeats` identical bash tool calls in a row, the automatic counterpart to a user
|
||||
typing an escalation keyword: the assistant, not the human, is the one stuck."""
|
||||
return [
|
||||
turn
|
||||
for i in range(repeats)
|
||||
for turn in (
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [{"type": "tool_use", "id": f"call-{i}", "name": "bash", "input": {"cmd": "pytest"}}],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"type": "tool_result", "tool_use_id": f"call-{i}", "is_error": True, "content": "fail"}],
|
||||
},
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
class TestStallEscalation:
|
||||
"""Mid-task auto-escalation when the assistant's own recent tool calls look stuck: the
|
||||
automatic counterpart to escalation_keywords, gated by stall_escalation_enabled and off
|
||||
by default."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_repeated_tool_calls_escalate_the_classified_tier(self, mock_router_instance, basic_config):
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={**basic_config, "stall_escalation_enabled": True},
|
||||
)
|
||||
messages = [*_stalled_tool_history(), {"role": "user", "content": "Hello there!"}]
|
||||
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
|
||||
assert result.model == "gpt-4o" # SIMPLE bumped to MEDIUM
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_varied_tool_calls_do_not_escalate(self, mock_router_instance, basic_config):
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={**basic_config, "stall_escalation_enabled": True},
|
||||
)
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [{"type": "tool_use", "id": "c1", "name": "bash", "input": {"cmd": "ls"}}],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"type": "tool_result", "tool_use_id": "c1", "is_error": False, "content": "ok"}],
|
||||
},
|
||||
{"role": "user", "content": "Hello there!"},
|
||||
]
|
||||
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
|
||||
assert result.model == "gpt-4o-mini" # not escalated
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disabled_by_default_ignores_repeated_tool_calls(self, mock_router_instance, basic_config):
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config=basic_config,
|
||||
)
|
||||
messages = [*_stalled_tool_history(), {"role": "user", "content": "Hello there!"}]
|
||||
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
|
||||
assert result.model == "gpt-4o-mini" # stall_escalation_enabled defaults False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_signals_record_stall_escalation(self, mock_router_instance, basic_config):
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={**basic_config, "stall_escalation_enabled": True},
|
||||
)
|
||||
messages = [*_stalled_tool_history(), {"role": "user", "content": "Hello there!"}]
|
||||
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
|
||||
assert "stall_escalation" in result.routing_decision["signals"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stall_escalation_caps_at_highest_tier(self, mock_router_instance, basic_config):
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={**basic_config, "stall_escalation_enabled": True},
|
||||
)
|
||||
messages = [
|
||||
*_stalled_tool_history(),
|
||||
{"role": "user", "content": "Let's think step by step and reason through this carefully."},
|
||||
]
|
||||
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
|
||||
assert result.model == "o1-preview" # already REASONING, stays there
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stall_escalation_stacks_with_keyword_escalation(self, mock_router_instance, basic_config):
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={**basic_config, "stall_escalation_enabled": True},
|
||||
)
|
||||
messages = [*_stalled_tool_history(), {"role": "user", "content": "LITELLM ESCALATE Hello there!"}]
|
||||
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
|
||||
assert result.model == "claude-sonnet-4-20250514" # SIMPLE -> MEDIUM (keyword) -> COMPLEX (stall)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_keyword_forced_tier_still_escalates_when_stalled(self, mock_router_instance, basic_config):
|
||||
"""A keyword rule forces its tier and returns before any classification runs, so
|
||||
without its own bump the one path that can pin a weak model to a whole conversation
|
||||
would be the one path a stall could never lift."""
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={
|
||||
**basic_config,
|
||||
"stall_escalation_enabled": True,
|
||||
"keyword_tier_rules": [{"keywords": ["billing"], "tier": "SIMPLE"}],
|
||||
},
|
||||
)
|
||||
healthy = await router.async_pre_routing_hook(
|
||||
model="test-model", request_kwargs={}, messages=[{"role": "user", "content": "a billing question"}]
|
||||
)
|
||||
assert healthy.model == "gpt-4o-mini" # forced SIMPLE, nothing stuck
|
||||
|
||||
stalled = await router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=[*_stalled_tool_history(), {"role": "user", "content": "a billing question"}],
|
||||
)
|
||||
assert stalled.model == "gpt-4o" # forced SIMPLE bumped to MEDIUM
|
||||
assert "stall_escalation" in stalled.routing_decision["signals"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_evidence_survives_a_new_human_ask(self, mock_router_instance, basic_config):
|
||||
"""A plain follow-up like 'try again' must not erase the stall evidence that came
|
||||
before it: escalation still fires on the turn carrying that follow-up."""
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={**basic_config, "stall_escalation_enabled": True},
|
||||
)
|
||||
messages = [*_stalled_tool_history(), {"role": "user", "content": "try again"}]
|
||||
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
|
||||
assert result.model == "gpt-4o" # SIMPLE ("try again" carries no signal) bumped to MEDIUM
|
||||
|
||||
|
||||
class TestRoutingDecisionContents:
|
||||
"""Every routing path must return a PreRoutingHookResponse carrying a routing_decision
|
||||
that names the mechanism that actually decided, with the facts of that path only."""
|
||||
|
|
@ -8027,7 +8171,6 @@ class TestClientHousekeepingCalls:
|
|||
assert result is not None
|
||||
assert result.model == "claude-sonnet-4-20250514"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_classifier_plugin_still_decides_its_own_routers(self, mock_router_instance):
|
||||
"""A plugin is where an operator encodes policy the tier ladder cannot express.
|
||||
|
|
@ -8062,9 +8205,7 @@ class TestClientHousekeepingCalls:
|
|||
assert result.model == "o1-preview"
|
||||
assert result.routing_decision["cause"] == "classifier_plugin"
|
||||
|
||||
def _adaptive_router(
|
||||
self, tier_distance_penalty: float, plan_mode_min_tier: str | None = None
|
||||
) -> ComplexityRouter:
|
||||
def _adaptive_router(self, tier_distance_penalty: float, plan_mode_min_tier: str | None = None) -> ComplexityRouter:
|
||||
adaptive_instance = MagicMock()
|
||||
adaptive_instance.model_list = [
|
||||
{
|
||||
|
|
@ -8101,9 +8242,7 @@ class TestClientHousekeepingCalls:
|
|||
return router
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_bandit_cannot_route_a_housekeeping_call_above_the_cheapest_tier(
|
||||
self, mock_router_instance
|
||||
):
|
||||
async def test_the_bandit_cannot_route_a_housekeeping_call_above_the_cheapest_tier(self, mock_router_instance):
|
||||
"""The tier here is what the request IS, not how hard it is, so the bandit has nothing to win.
|
||||
|
||||
Without a ceiling the tier distance penalty is the only thing holding the tier, so a
|
||||
|
|
@ -8136,7 +8275,6 @@ class TestClientHousekeepingCalls:
|
|||
assert result is not None
|
||||
assert result.model == "premium"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_housekeeping_call_never_becomes_the_session_pin(self, mock_router_instance):
|
||||
"""Pinning this is the most expensive mistake of the transient causes.
|
||||
|
|
@ -8178,9 +8316,7 @@ class TestClientHousekeepingCalls:
|
|||
assert work_turn.routing_decision["cause"] == "llm_classifier"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_decision_records_which_sentinel_matched(
|
||||
self, mock_router_instance, llm_classifier_config
|
||||
):
|
||||
async def test_the_decision_records_which_sentinel_matched(self, mock_router_instance, llm_classifier_config):
|
||||
"""The cause's contract says the sentinel rides in matched_keyword, so it has to be there.
|
||||
|
||||
Without it an operator reading the logs can see that a call was treated as housekeeping but
|
||||
|
|
@ -8201,7 +8337,6 @@ class TestClientHousekeepingCalls:
|
|||
"Write the title in the predominant language of the session"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_plan_mode_floor_raises_a_housekeeping_call_under_adaptive(self, mock_router_instance):
|
||||
"""Floor and ceiling must not contradict each other on the same request.
|
||||
|
|
@ -9428,6 +9563,7 @@ class TestTierDefinitions:
|
|||
({"adaptive": True}, "severity order"),
|
||||
({"session_affinity": True}, "severity order"),
|
||||
({"escalation_keywords": ["GO UP"]}, "severity order"),
|
||||
({"stall_escalation_enabled": True}, "severity order"),
|
||||
(
|
||||
{"classifier_llm_config": {"model": "haiku-classifier", "system_prompt": "grade it"}},
|
||||
"system_prompt",
|
||||
|
|
@ -10735,9 +10871,7 @@ class TestHeuristicFirst:
|
|||
|
||||
# Scores 0.175 with one signal, so it sits 0.025 from simple_medium: the pair of tiers either side of
|
||||
# that boundary are different model pools, and a hair's difference in score picks the other one.
|
||||
NEAR_BOUNDARY_PROMPT = (
|
||||
"design a distributed cache with consistent hashing, then explain the failure modes step by step"
|
||||
)
|
||||
NEAR_BOUNDARY_PROMPT = "design a distributed cache with consistent hashing, then explain the failure modes step by step"
|
||||
|
||||
# Scores 0.075 with signals, the far side of any margin under 0.075: the scorer is decided here.
|
||||
CLEAR_OF_BOUNDARY_PROMPT = "explain step by step how consistent hashing rebalances keys"
|
||||
|
|
@ -11179,6 +11313,7 @@ class TestContextWindowEscalation:
|
|||
litellm_router_instance=_windowed_router(_SMALL, _BIG),
|
||||
complexity_router_config=_tier_config(session_affinity=True),
|
||||
)
|
||||
|
||||
def session_kwargs() -> dict[str, object]:
|
||||
return {"metadata": {"session_id": "s-1", "user_api_key_hash": "k-1"}}
|
||||
|
||||
|
|
@ -11203,6 +11338,7 @@ class TestContextWindowEscalation:
|
|||
litellm_router_instance=_windowed_router(_SMALL, _BIG),
|
||||
complexity_router_config=_tier_config(session_affinity=True),
|
||||
)
|
||||
|
||||
def session_kwargs() -> dict[str, object]:
|
||||
return {"metadata": {"session_id": "s-2", "user_api_key_hash": "k-2"}}
|
||||
|
||||
|
|
@ -11280,7 +11416,9 @@ class TestContextWindowEscalation:
|
|||
copilot_resolutions: List = []
|
||||
|
||||
def _guarded(*args, **kwargs):
|
||||
target = str(kwargs.get("model") or (args[0] if args else "")) + str(kwargs.get("custom_llm_provider") or "")
|
||||
target = str(kwargs.get("model") or (args[0] if args else "")) + str(
|
||||
kwargs.get("custom_llm_provider") or ""
|
||||
)
|
||||
if "github_copilot" in target:
|
||||
copilot_resolutions.append(target)
|
||||
raise RuntimeError("the gate must not resolve an authenticating provider")
|
||||
|
|
@ -12291,3 +12429,277 @@ class TestTierHealthFailover:
|
|||
for _ in range(20)
|
||||
]
|
||||
assert {r.model for r in results} == {"live-c"}
|
||||
|
||||
|
||||
ANTHROPIC_IMG_PART = {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "aGk="}}
|
||||
RESPONSES_IMG_PART = {"type": "input_image", "image_url": "data:image/png;base64,aGk="}
|
||||
|
||||
|
||||
class TestClassifierVision:
|
||||
"""classifier_llm_config.vision: what the LLM classifier is shown for an image-bearing turn."""
|
||||
|
||||
TIERS = {"SIMPLE": "t-simple", "MEDIUM": "t-medium", "COMPLEX": "t-complex", "REASONING": "t-reasoning"}
|
||||
|
||||
@staticmethod
|
||||
def _router(mock_router_instance, *, vision, classifier_declares_vision=True, classifier_type="llm", **extra):
|
||||
def get_model_list(model_name=None):
|
||||
if model_name != "clf":
|
||||
return [{"model_name": model_name, "litellm_params": {"model": "openai/gpt-4o"}}]
|
||||
declared = classifier_declares_vision
|
||||
return [
|
||||
{
|
||||
"model_name": "clf",
|
||||
"litellm_params": {"model": "openai/unmapped-classifier"},
|
||||
"model_info": {} if declared is None else {"supports_vision": declared},
|
||||
}
|
||||
]
|
||||
|
||||
mock_router_instance.get_model_list = get_model_list
|
||||
classifier_llm_config = {"model": "clf", "circuit_breaker_enabled": False}
|
||||
return ComplexityRouter(
|
||||
model_name="vision-classifier-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={
|
||||
"classifier_type": classifier_type,
|
||||
"classifier_llm_config": (
|
||||
classifier_llm_config if vision is None else {**classifier_llm_config, "vision": vision}
|
||||
),
|
||||
"tiers": dict(TestClassifierVision.TIERS),
|
||||
**extra,
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _classifier_user_content(mock_router_instance):
|
||||
return mock_router_instance.acompletion.call_args.kwargs["messages"][-1]["content"]
|
||||
|
||||
@staticmethod
|
||||
def _turn(*parts):
|
||||
return [{"role": "user", "content": list(parts)}]
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _classifier_answers_complex(self, mock_router_instance):
|
||||
mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "COMPLEX"}'))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"vision, classifier_declares_vision",
|
||||
[
|
||||
(None, True),
|
||||
({"enabled": False}, True),
|
||||
({"enabled": True}, False),
|
||||
({"enabled": True}, None),
|
||||
],
|
||||
ids=["vision_unset", "vision_disabled", "classifier_declared_text_only", "classifier_undeclared"],
|
||||
)
|
||||
async def test_payload_stays_text_only(self, mock_router_instance, vision, classifier_declares_vision):
|
||||
"""Off, or a classifier not declared vision-capable, keeps the plain-string payload.
|
||||
|
||||
The undeclared case is the polarity. A text-only classifier handed an image rejects the
|
||||
call, the rejection is swallowed by the classifier's own fallback, and every image request
|
||||
then serves from the fallback tier while still paying for the failed call. Staying text-only
|
||||
is instead a visible no-op the operator fixes by declaring supports_vision.
|
||||
"""
|
||||
router = self._router(
|
||||
mock_router_instance, vision=vision, classifier_declares_vision=classifier_declares_vision
|
||||
)
|
||||
await router.async_pre_routing_hook(
|
||||
model="m", request_kwargs={}, messages=self._turn({"type": "text", "text": "what is this"}, IMG_PART)
|
||||
)
|
||||
content = self._classifier_user_content(mock_router_instance)
|
||||
assert isinstance(content, str)
|
||||
assert "what is this" in content
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deployment_model_info_enables_a_classifier_the_cost_map_does_not_describe(
|
||||
self, mock_router_instance
|
||||
):
|
||||
"""The escape hatch for an unmapped classifier name, and the reason undeclared can stay off.
|
||||
|
||||
`_router` gives every deployment an `openai/unmapped-*` litellm_params model, so nothing in
|
||||
the cost map declares it and the verdict comes only from model_info.
|
||||
"""
|
||||
router = self._router(mock_router_instance, vision={"enabled": True}, classifier_declares_vision=True)
|
||||
await router.async_pre_routing_hook(
|
||||
model="m", request_kwargs={}, messages=self._turn({"type": "text", "text": "what is this"}, IMG_PART)
|
||||
)
|
||||
assert [b["type"] for b in self._classifier_user_content(mock_router_instance)] == ["text", "image_url"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"part",
|
||||
[IMG_PART, ANTHROPIC_IMG_PART, RESPONSES_IMG_PART],
|
||||
ids=["chat_completions", "anthropic_messages", "responses"],
|
||||
)
|
||||
async def test_image_reaches_the_classifier_in_chat_completions_dialect(self, mock_router_instance, part):
|
||||
"""Every surface's dialect arrives as a chat-completions image_url on the classifier call.
|
||||
|
||||
/v1/messages hands the hook an Anthropic image block untranslated, so forwarding verbatim
|
||||
would send the classifier a content part its own request dialect has no meaning for.
|
||||
"""
|
||||
router = self._router(mock_router_instance, vision={"enabled": True})
|
||||
await router.async_pre_routing_hook(
|
||||
model="m", request_kwargs={}, messages=self._turn({"type": "text", "text": "what is this"}, part)
|
||||
)
|
||||
content = self._classifier_user_content(mock_router_instance)
|
||||
assert [block["type"] for block in content] == ["text", "image_url"]
|
||||
assert content[1]["image_url"] == {"url": "data:image/png;base64,aGk="}
|
||||
assert "what is this" in content[0]["text"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"part",
|
||||
[
|
||||
{"type": "image_url", "image_url": {"url": "http://169.254.169.254/latest/meta-data/"}},
|
||||
{"type": "image_url", "image_url": {"url": "https://example.internal/secret.png"}},
|
||||
{"type": "input_image", "image_url": "https://example.internal/secret.png"},
|
||||
{"type": "image", "source": {"type": "url", "url": "https://example.internal/secret.png"}},
|
||||
],
|
||||
ids=["metadata_service", "chat_completions", "responses", "anthropic"],
|
||||
)
|
||||
async def test_remote_url_images_are_never_forwarded(self, mock_router_instance, part):
|
||||
"""A caller-supplied URL must not reach an internal call the caller did not ask for.
|
||||
|
||||
Provider adapters do not uniformly delegate fetching: gigachat downloads any non-data URL
|
||||
from the proxy host, so forwarding one would turn a router-scoped key into a proxy-side GET
|
||||
at an address of the caller's choosing.
|
||||
"""
|
||||
router = self._router(mock_router_instance, vision={"enabled": True})
|
||||
await router.async_pre_routing_hook(
|
||||
model="m", request_kwargs={}, messages=self._turn({"type": "text", "text": "what is this"}, part)
|
||||
)
|
||||
assert isinstance(self._classifier_user_content(mock_router_instance), str)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_remote_url_image_only_turn_does_not_reach_the_classifier(self, mock_router_instance):
|
||||
"""With nothing forwardable left, the turn stays unclassifiable rather than sending the URL."""
|
||||
router = self._router(mock_router_instance, vision={"enabled": True})
|
||||
response = await router.async_pre_routing_hook(
|
||||
model="m",
|
||||
request_kwargs={},
|
||||
messages=self._turn({"type": "image_url", "image_url": {"url": "https://example.internal/x.png"}}),
|
||||
)
|
||||
assert response.routing_decision["cause"] == "default_fallback"
|
||||
mock_router_instance.acompletion.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_image_only_turn_is_classified_instead_of_falling_back(self, mock_router_instance):
|
||||
"""A turn carrying only an image reaches the classifier rather than the default model.
|
||||
|
||||
It flattens to empty text, so before this it never reached the classifier at all and was
|
||||
routed as default_fallback on text the request never contained.
|
||||
"""
|
||||
router = self._router(mock_router_instance, vision={"enabled": True})
|
||||
response = await router.async_pre_routing_hook(
|
||||
model="m", request_kwargs={}, messages=self._turn(IMG_PART)
|
||||
)
|
||||
assert response.routing_decision["cause"] == "llm_classifier"
|
||||
assert response.model == "t-complex"
|
||||
assert [block["type"] for block in self._classifier_user_content(mock_router_instance)] == [
|
||||
"text",
|
||||
"image_url",
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_image_only_turn_still_falls_back_when_vision_is_off(self, mock_router_instance):
|
||||
router = self._router(mock_router_instance, vision={"enabled": False})
|
||||
response = await router.async_pre_routing_hook(
|
||||
model="m", request_kwargs={}, messages=self._turn(IMG_PART)
|
||||
)
|
||||
assert response.routing_decision["cause"] == "default_fallback"
|
||||
mock_router_instance.acompletion.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("max_images, expected", [(1, 1), (2, 2), (5, 3)])
|
||||
async def test_max_images_caps_what_is_forwarded(self, mock_router_instance, max_images, expected):
|
||||
router = self._router(mock_router_instance, vision={"enabled": True, "max_images": max_images})
|
||||
images = [dict(IMG_PART, image_url={"url": f"data:image/png;base64,{n}"}) for n in ("a", "b", "c")]
|
||||
await router.async_pre_routing_hook(
|
||||
model="m", request_kwargs={}, messages=self._turn({"type": "text", "text": "look"}, *images)
|
||||
)
|
||||
content = self._classifier_user_content(mock_router_instance)
|
||||
forwarded = [block for block in content if block["type"] == "image_url"]
|
||||
assert len(forwarded) == expected
|
||||
assert [block["image_url"]["url"] for block in forwarded] == [
|
||||
f"data:image/png;base64,{n}" for n in ("a", "b", "c")[:expected]
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_earlier_turn_images_are_not_forwarded(self, mock_router_instance):
|
||||
"""Only the newest user turn's images ride along, so history cannot inflate every call.
|
||||
|
||||
The two turns carry different images on purpose: identical ones would pass this assertion
|
||||
whichever turn the helper read.
|
||||
"""
|
||||
older = dict(IMG_PART, image_url={"url": "data:image/png;base64,OLDER"})
|
||||
newer = dict(IMG_PART, image_url={"url": "data:image/png;base64,NEWER"})
|
||||
router = self._router(mock_router_instance, vision={"enabled": True, "max_images": 5})
|
||||
await router.async_pre_routing_hook(
|
||||
model="m",
|
||||
request_kwargs={},
|
||||
messages=[
|
||||
{"role": "user", "content": [{"type": "text", "text": "first"}, older]},
|
||||
{"role": "assistant", "content": "ok"},
|
||||
{"role": "user", "content": [{"type": "text", "text": "second"}, newer]},
|
||||
],
|
||||
)
|
||||
content = self._classifier_user_content(mock_router_instance)
|
||||
forwarded = [block for block in content if block["type"] == "image_url"]
|
||||
assert [block["image_url"]["url"] for block in forwarded] == ["data:image/png;base64,NEWER"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_logged_request_body_matches_what_was_sent(self, mock_router_instance):
|
||||
"""proxy_server_request is the logged copy of the classifier call and must not drift."""
|
||||
router = self._router(mock_router_instance, vision={"enabled": True})
|
||||
await router.async_pre_routing_hook(
|
||||
model="m", request_kwargs={}, messages=self._turn({"type": "text", "text": "what is this"}, IMG_PART)
|
||||
)
|
||||
call_kwargs = mock_router_instance.acompletion.call_args.kwargs
|
||||
assert call_kwargs["proxy_server_request"]["body"]["messages"] == call_kwargs["messages"]
|
||||
|
||||
SHORT_CIRCUIT_ARMS = [
|
||||
("heuristic_first", {"heuristic_first_max_tier": "SIMPLE"}, "heuristic_first_short_circuit"),
|
||||
("hybrid", {"hybrid_boundary_margin": 0.05}, "hybrid_short_circuit"),
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"classifier_type, extra, short_circuit_cause", SHORT_CIRCUIT_ARMS, ids=["heuristic_first", "hybrid"]
|
||||
)
|
||||
async def test_local_scorer_cannot_short_circuit_a_turn_it_cannot_see(
|
||||
self, mock_router_instance, classifier_type, extra, short_circuit_cause
|
||||
):
|
||||
"""The scorer reads text alone, so its confidence is not a verdict on an image turn.
|
||||
|
||||
Both arms are tuned so the scorer WOULD short-circuit on this exact text, which is what
|
||||
makes the image the only variable; a margin loose enough to leave the score undecided
|
||||
would pass whether or not the guard exists.
|
||||
"""
|
||||
router = self._router(
|
||||
mock_router_instance, vision={"enabled": True}, classifier_type=classifier_type, **extra
|
||||
)
|
||||
response = await router.async_pre_routing_hook(
|
||||
model="m", request_kwargs={}, messages=self._turn({"type": "text", "text": "what is this"}, IMG_PART)
|
||||
)
|
||||
assert response.routing_decision["cause"] == "llm_classifier"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"classifier_type, extra, short_circuit_cause", SHORT_CIRCUIT_ARMS, ids=["heuristic_first", "hybrid"]
|
||||
)
|
||||
async def test_local_scorer_still_short_circuits_without_images(
|
||||
self, mock_router_instance, classifier_type, extra, short_circuit_cause
|
||||
):
|
||||
"""The negative class: same router, same text, no image, and the scorer still decides."""
|
||||
router = self._router(
|
||||
mock_router_instance, vision={"enabled": True}, classifier_type=classifier_type, **extra
|
||||
)
|
||||
response = await router.async_pre_routing_hook(
|
||||
model="m", request_kwargs={}, messages=[{"role": "user", "content": "what is this"}]
|
||||
)
|
||||
assert response.routing_decision["cause"] == short_circuit_cause
|
||||
mock_router_instance.acompletion.assert_not_awaited()
|
||||
|
||||
def test_max_images_must_be_positive(self):
|
||||
with pytest.raises(ValidationError):
|
||||
ClassifierLLMConfig(model="clf", vision={"enabled": True, "max_images": 0})
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
154
tests/test_litellm/router_strategy/test_stall_detector.py
Normal file
154
tests/test_litellm/router_strategy/test_stall_detector.py
Normal file
|
|
@ -0,0 +1,154 @@
|
|||
"""
|
||||
Tests for mid-task stall detection: repeated identical tool calls or repeated tool
|
||||
errors, read from both Anthropic Messages and chat-completions tool-call shapes.
|
||||
"""
|
||||
|
||||
from litellm.router_strategy.complexity_router.stall_detector import detect_stalled_task
|
||||
|
||||
|
||||
def _anthropic_call(call_id: str, name: str, arguments: dict, *, is_error: bool) -> list[dict]:
|
||||
return [
|
||||
{"role": "assistant", "content": [{"type": "tool_use", "id": call_id, "name": name, "input": arguments}]},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"type": "tool_result", "tool_use_id": call_id, "is_error": is_error, "content": "result"}],
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def _chat_completions_call(call_id: str, name: str, arguments_json: str) -> list[dict]:
|
||||
return [
|
||||
{
|
||||
"role": "assistant",
|
||||
"tool_calls": [
|
||||
{"id": call_id, "type": "function", "function": {"name": name, "arguments": arguments_json}}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": call_id, "content": "result"},
|
||||
]
|
||||
|
||||
|
||||
class TestDetectStalledTask:
|
||||
def test_repeated_identical_anthropic_calls_are_stalled(self):
|
||||
messages = [
|
||||
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=False),
|
||||
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=False),
|
||||
*_anthropic_call("t3", "bash", {"cmd": "pytest"}, is_error=False),
|
||||
]
|
||||
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is True
|
||||
|
||||
def test_repeated_errors_are_stalled_even_with_varied_arguments(self):
|
||||
messages = [
|
||||
*_anthropic_call("t1", "bash", {"cmd": "pytest tests/a.py"}, is_error=True),
|
||||
*_anthropic_call("t2", "bash", {"cmd": "pytest tests/b.py"}, is_error=True),
|
||||
*_anthropic_call("t3", "bash", {"cmd": "pytest tests/c.py"}, is_error=True),
|
||||
]
|
||||
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is True
|
||||
|
||||
def test_varied_successful_calls_are_not_stalled(self):
|
||||
messages = [
|
||||
*_anthropic_call("t1", "bash", {"cmd": "ls"}, is_error=False),
|
||||
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=False),
|
||||
*_anthropic_call("t3", "grep", {"pattern": "x"}, is_error=False),
|
||||
]
|
||||
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is False
|
||||
|
||||
def test_chat_completions_repeats_are_stalled(self):
|
||||
messages = [
|
||||
*_chat_completions_call("c1", "bash", '{"cmd": "pytest"}'),
|
||||
*_chat_completions_call("c2", "bash", '{"cmd": "pytest"}'),
|
||||
*_chat_completions_call("c3", "bash", '{"cmd": "pytest"}'),
|
||||
]
|
||||
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is True
|
||||
|
||||
def test_chat_completions_has_no_structured_error_signal(self):
|
||||
"""A chat-completions tool message carries no standard error flag, so varied calls
|
||||
whose content happens to read like failures still aren't flagged on error alone."""
|
||||
messages = [
|
||||
*_chat_completions_call("c1", "bash", '{"cmd": "a"}'),
|
||||
*_chat_completions_call("c2", "bash", '{"cmd": "b"}'),
|
||||
*_chat_completions_call("c3", "bash", '{"cmd": "c"}'),
|
||||
]
|
||||
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is False
|
||||
|
||||
def test_dict_and_json_string_arguments_compare_equal_across_surfaces(self):
|
||||
messages = [
|
||||
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=False),
|
||||
*_chat_completions_call("c2", "bash", '{"cmd": "pytest"}'),
|
||||
*_anthropic_call("t3", "bash", {"cmd": "pytest"}, is_error=False),
|
||||
]
|
||||
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is True
|
||||
|
||||
def test_below_repeat_threshold_is_not_stalled(self):
|
||||
messages = [
|
||||
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=False),
|
||||
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=False),
|
||||
]
|
||||
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is False
|
||||
|
||||
def test_evidence_older_than_the_window_does_not_count(self):
|
||||
"""Only the most recent `window` tool calls are considered, so a stall the model
|
||||
already recovered from does not keep re-triggering forever."""
|
||||
messages = [
|
||||
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=False),
|
||||
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=False),
|
||||
*_anthropic_call("t3", "bash", {"cmd": "pytest"}, is_error=False),
|
||||
*_anthropic_call("t4", "grep", {"pattern": "a"}, is_error=False),
|
||||
*_anthropic_call("t5", "grep", {"pattern": "b"}, is_error=False),
|
||||
]
|
||||
assert detect_stalled_task(messages, window=2, repeat_threshold=2) is False
|
||||
|
||||
def test_evidence_survives_a_new_human_ask(self):
|
||||
"""A follow-up like 'try again' must not erase evidence from before it: detection
|
||||
reads the whole message list, not just the turns since the newest human ask."""
|
||||
messages = [
|
||||
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=False),
|
||||
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=False),
|
||||
*_anthropic_call("t3", "bash", {"cmd": "pytest"}, is_error=False),
|
||||
{"role": "user", "content": [{"type": "text", "text": "try again"}]},
|
||||
]
|
||||
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is True
|
||||
|
||||
def test_a_recovered_task_is_not_stalled_while_its_old_failures_sit_in_the_window(self):
|
||||
"""The three identical failures stay in the window for a few turns after the model
|
||||
breaks out of them, and counting them on their own would escalate a request that is
|
||||
already making progress again."""
|
||||
messages = [
|
||||
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=True),
|
||||
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=True),
|
||||
*_anthropic_call("t3", "bash", {"cmd": "pytest"}, is_error=True),
|
||||
*_anthropic_call("t4", "read_file", {"path": "conftest.py"}, is_error=False),
|
||||
*_anthropic_call("t5", "edit_file", {"path": "conftest.py"}, is_error=False),
|
||||
]
|
||||
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is False
|
||||
|
||||
def test_a_retry_loop_broken_up_by_an_unrelated_call_still_counts(self):
|
||||
"""Anchoring on the newest call must not require the repeats to be adjacent: a model
|
||||
re-running the same failing command around a lookup in between is still stuck."""
|
||||
messages = [
|
||||
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=True),
|
||||
*_anthropic_call("t2", "read_file", {"path": "conftest.py"}, is_error=False),
|
||||
*_anthropic_call("t3", "bash", {"cmd": "pytest"}, is_error=True),
|
||||
*_anthropic_call("t4", "bash", {"cmd": "pytest"}, is_error=True),
|
||||
]
|
||||
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is True
|
||||
|
||||
def test_errors_only_count_while_the_newest_call_is_still_failing(self):
|
||||
messages = [
|
||||
*_anthropic_call("t1", "bash", {"cmd": "pytest a"}, is_error=True),
|
||||
*_anthropic_call("t2", "bash", {"cmd": "pytest b"}, is_error=True),
|
||||
*_anthropic_call("t3", "bash", {"cmd": "pytest c"}, is_error=True),
|
||||
*_anthropic_call("t4", "bash", {"cmd": "pytest d"}, is_error=False),
|
||||
]
|
||||
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is False
|
||||
|
||||
def test_no_messages_is_not_stalled(self):
|
||||
assert detect_stalled_task(None, window=6, repeat_threshold=3) is False
|
||||
assert detect_stalled_task([], window=6, repeat_threshold=3) is False
|
||||
|
||||
def test_zero_threshold_never_flags_stalled(self):
|
||||
messages = [
|
||||
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=True),
|
||||
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=True),
|
||||
]
|
||||
assert detect_stalled_task(messages, window=6, repeat_threshold=0) is False
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 22325
|
||||
"limit": 22324
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26747
|
||||
"limit": 26746
|
||||
},
|
||||
"LIT003": {
|
||||
"limit": 261
|
||||
|
|
|
|||
|
|
@ -1208,11 +1208,6 @@
|
|||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/app/(dashboard)/vector-stores/_components/VectorStoreForm.tsx": {
|
||||
"no-nested-ternary": {
|
||||
"count": 2
|
||||
}
|
||||
},
|
||||
"src/app/(dashboard)/vector-stores/_components/index.tsx": {
|
||||
"local/filename-pascal-case": {
|
||||
"count": 1
|
||||
|
|
|
|||
6
ui/litellm-dashboard/public/assets/logos/mongodb.svg
Normal file
6
ui/litellm-dashboard/public/assets/logos/mongodb.svg
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
<?xml version="1.0" encoding="UTF-8"?>
|
||||
<svg width="64" height="64" viewBox="0 0 64 64" xmlns="http://www.w3.org/2000/svg" role="img" aria-label="MongoDB">
|
||||
<path fill="#00684A" fill-rule="evenodd" d="M 33.1 2.4 C 33.1 2.4 36.6 8.9 44.4 15.1 C 52.2 21.3 55.1 28.5 54.2 37.1 C 53.3 45.7 47.6 52.7 40.1 55.6 C 38.5 56.2 37.3 57.4 36.7 59 L 34.7 64 L 31.4 64 L 30.3 59.4 C 29.9 57.6 28.7 56.1 27 55.3 C 19.4 51.8 14 44.6 13.4 36 C 12.7 25.8 18.1 19.6 24.6 14 C 30.2 9.2 33.1 2.4 33.1 2.4 Z"/>
|
||||
<path fill="#00ED64" d="M 33.1 2.4 C 33.1 2.4 30.2 9.2 24.6 14 C 18.1 19.6 12.7 25.8 13.4 36 C 14 44.6 19.4 51.8 27 55.3 C 28.7 56.1 29.9 57.6 30.3 59.4 L 31.4 64 L 32.9 64 Z"/>
|
||||
<path fill="#B8C4C2" d="M 32.4 46.9 L 31.9 46.1 C 31.5 40.4 31.4 34.6 31.6 28.9 C 31.7 26.1 31.8 20.1 32.9 17.1 C 32.6 20.6 32.7 43.1 32.8 45.4 C 32.7 45.9 32.6 46.4 32.4 46.9 Z"/>
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 866 B |
|
|
@ -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", () => {
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue