From eb7eeb54199968aaf9eaee3d89c65f2e47752cd9 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 15:57:51 -0700 Subject: [PATCH 01/10] test(straiker): deterministic integration audit of the v3 platform relay (#42781) * test(straiker): deterministic integration audit of the v3 platform relay 34 cells against a real two-worker proxy, Postgres and Redis with a scripted provider upstream and a local Straiker sink: v3 allow, block, deny, replay and killswitch verdicts on chat completions, messages, responses and completions across the OpenAI and Anthropic SDKs and raw httpx, pre_call, post_call and logging_only modes, header and identity precedence, credential redaction, sink outages, malformed verdicts, unauthenticated and unknown-model requests, management endpoints, the unchanged v1 path, and a mixed burst through a sink outage, a worker kill and a proxy restart Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(straiker): bind the spend-row pattern inside the outage burst poll Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(straiker): kill a real uvicorn worker and prove detect runs before the unknown-model error Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(straiker): assert the v1 webhook ran on the v1 block cell and check every non-streaming burst spend row Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yucheng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_straiker_v3_platform.py | 1090 +++++++++++++++++ 1 file changed, 1090 insertions(+) create mode 100644 tests/integration/observability/test_straiker_v3_platform.py diff --git a/tests/integration/observability/test_straiker_v3_platform.py b/tests/integration/observability/test_straiker_v3_platform.py new file mode 100644 index 00000000000..e44abf4e066 --- /dev/null +++ b/tests/integration/observability/test_straiker_v3_platform.py @@ -0,0 +1,1090 @@ +"""Straiker guardrail on both platform APIs, driven through a real proxy. + +The Straiker platform is the only double: an owned HTTP sink that speaks the v1 webhook and the v3 +detect wire protocols and records every request. The provider is a second owned sink. The proxy, +its guardrail registry, Postgres and Redis run for real with two workers. +""" + +from __future__ import annotations + +import hashlib +import itertools +import json +import os +import signal +import socket +import threading +import uuid +from collections.abc import Callable, Iterator +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass, field +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from pathlib import Path +from typing import Final + +import anthropic +import httpx +import openai +import psutil +import pytest +import yaml +from integration._support.client import Gateway, eventually, gateway_from_environment, object_value +from integration._support.database import read_rows +from integration._support.process import OwnedProxy, owned_proxy, owned_proxy_process +from integration._support.wire import Reply, Request, wire_server + +V3_KEY: Final = "sk_agt_synthetic_integration_key" +V1_KEY: Final = "synthetic-v1-collection-key" +V3_PATH: Final = "/api/v3/detect" +V1_PATH: Final = "/api/v1/detect/webhook" +BLOCK_MARK: Final = "SYNTHETIC-INJECTION" +KILL_MARK: Final = "SYNTHETIC-KILLSWITCH" +DENY_MARK: Final = "SYNTHETIC-DENY" +SINK_500_MARK: Final = "SYNTHETIC-SINK-500" +SINK_401_MARK: Final = "SYNTHETIC-SINK-401" +SINK_GARBAGE_MARK: Final = "SYNTHETIC-SINK-GARBAGE" +LOG_BLOCK_MARK: Final = "SYNTHETIC-LOG-ONLY-BLOCK" +OPEN_500_MARK: Final = "SYNTHETIC-OPEN-500" +V1_500_MARK: Final = "SYNTHETIC-V1-500" +V1_BLOCK_MARK: Final = "SYNTHETIC-V1-BLOCK" +AUDIT_AGENT: Final = "audit-agent" +POST_AGENT: Final = "post-agent" +LOG_AGENT: Final = "log-agent" +OPEN_AGENT: Final = "open-agent" +BLOCK_MESSAGE: Final = "Straiker blocked this turn: prompt-injection" +DENY_MESSAGE: Final = "Straiker denied this turn" + + +@dataclass(frozen=True, slots=True) +class Seen: + target: str + headers: dict[str, str] + body: dict[str, object] + + +@dataclass(slots=True) +class Sink: + """Owned Straiker platform double on a fixed port so a test can stop and restart it.""" + + port: int + seen: list[Seen] = field(default_factory=list) + lock: threading.Lock = field(default_factory=threading.Lock) + server: ThreadingHTTPServer | None = None + thread: threading.Thread | None = None + + @property + def url(self) -> str: + return f"http://127.0.0.1:{self.port}" + + def start(self) -> None: + sink: Final = self + + class Handler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def do_POST(self) -> None: + raw: Final = self.rfile.read(int(self.headers.get("content-length", "0"))) + body: Final = json.loads(raw) + seen: Final = Seen(self.path, {k.lower(): v for k, v in self.headers.items()}, body) + with sink.lock: + sink.seen.append(seen) + status, payload = _verdict(seen, raw.decode()) + self.send_response(status) + self.send_header("content-type", "application/json") + self.send_header("content-length", str(len(payload))) + self.send_header("connection", "close") + self.end_headers() + self.wfile.write(payload) + + def log_message(self, format: str, *args: object) -> None: + pass + + class Server(ThreadingHTTPServer): + allow_reuse_address = True + daemon_threads = True + + self.server = Server(("127.0.0.1", self.port), Handler) + self.thread = threading.Thread(target=self.server.serve_forever, daemon=True) + self.thread.start() + + def stop(self) -> None: + assert self.server is not None and self.thread is not None + self.server.shutdown() + self.server.server_close() + self.thread.join(timeout=5) + self.server = None + self.thread = None + + def drain(self) -> tuple[Seen, ...]: + with self.lock: + taken: Final = tuple(self.seen) + self.seen.clear() + return taken + + def for_marker(self, marker: str) -> tuple[Seen, ...]: + with self.lock: + return tuple(s for s in self.seen if marker in json.dumps(s.body)) + + +def _verdict(seen: Seen, text: str) -> tuple[int, bytes]: + agent: Final = seen.headers.get("x-s6r-agent") + if ( + SINK_500_MARK in text + or (OPEN_500_MARK in text and agent == OPEN_AGENT) + or (V1_500_MARK in text and seen.target == V1_PATH) + ): + return 500, b'{"error":"synthetic outage"}' + if SINK_401_MARK in text: + return 401, b'{"error":"synthetic bad key"}' + if SINK_GARBAGE_MARK in text: + return 200, b"not json" + if seen.target == V1_PATH: + if BLOCK_MARK in text or V1_BLOCK_MARK in text: + return 200, json.dumps({"action": "BLOCKED", "blocked_reason": BLOCK_MESSAGE}).encode() + return 200, json.dumps({"action": "NONE"}).encode() + assert seen.target == V3_PATH, seen.target + turn: Final = "turn-" + hashlib.sha256(text.encode()).hexdigest()[:12] + if BLOCK_MARK in text or (LOG_BLOCK_MARK in text and agent == LOG_AGENT): + return 200, json.dumps( + { + "hookSpecificOutput": {"permissionDecision": "block"}, + "straiker": { + "action": "block", + "blocked_by": ["prompt-injection"], + "block_message": BLOCK_MESSAGE, + "turn_id": turn, + }, + } + ).encode() + if DENY_MARK in text: + return 200, json.dumps({"action": "deny", "deny_reason": DENY_MESSAGE, "turn_id": turn}).encode() + if KILL_MARK in text: + return 200, json.dumps( + {"straiker": {"action": "block", "block_message": BLOCK_MESSAGE, "turn_id": turn}} + ).encode() + return 200, json.dumps( + {"hookSpecificOutput": {"permissionDecision": "allow"}, "straiker": {"action": "allow", "turn_id": turn}} + ).encode() + + +def _marker_in(body: bytes) -> str: + text: Final = body.decode() + start: Final = text.find("mark-") + return text[start : start + 37] if start >= 0 else "mark-" + uuid.uuid4().hex + + +def _chat_body(marker: str, answer: str) -> bytes: + return json.dumps( + { + "id": "chatcmpl-" + marker, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": answer}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + } + ).encode() + + +def _chat_chunks(marker: str, answer: str) -> tuple[bytes, ...]: + def chunk(delta: dict[str, object], finish: str | None) -> bytes: + return ( + "data: " + + json.dumps( + { + "id": "chatcmpl-" + marker, + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "delta": delta, "finish_reason": finish}], + } + ) + + "\n\n" + ).encode() + + return ( + chunk({"role": "assistant", "content": answer[:3]}, None), + chunk({"content": answer[3:]}, "stop"), + b"data: [DONE]\n\n", + ) + + +def _messages_body(marker: str, answer: str) -> bytes: + return json.dumps( + { + "id": "msg_" + marker, + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5-20250929", + "content": [{"type": "text", "text": answer}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 10, "output_tokens": 5}, + } + ).encode() + + +def _messages_chunks(marker: str, answer: str) -> tuple[bytes, ...]: + def event(name: str, payload: dict[str, object]) -> bytes: + return f"event: {name}\ndata: {json.dumps(payload)}\n\n".encode() + + return ( + event( + "message_start", + { + "type": "message_start", + "message": { + "id": "msg_" + marker, + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5-20250929", + "content": [], + "stop_reason": None, + "stop_sequence": None, + "usage": {"input_tokens": 10, "output_tokens": 1}, + }, + }, + ), + event( + "content_block_start", + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + ), + event( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": answer}}, + ), + event("content_block_stop", {"type": "content_block_stop", "index": 0}), + event( + "message_delta", + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn", "stop_sequence": None}, + "usage": {"output_tokens": 5}, + }, + ), + event("message_stop", {"type": "message_stop"}), + ) + + +def _responses_body(marker: str, answer: str) -> bytes: + return json.dumps( + { + "id": "resp_" + marker, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "type": "message", + "id": "msgo_" + marker, + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": answer, "annotations": []}], + } + ], + "usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15}, + } + ).encode() + + +def _completion_body(marker: str, answer: str) -> bytes: + return json.dumps( + { + "id": "cmpl-" + marker, + "object": "text_completion", + "created": 1, + "model": "gpt-3.5-turbo-instruct", + "choices": [{"index": 0, "text": answer, "finish_reason": "stop", "logprobs": None}], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + } + ).encode() + + +_PROVIDER_CALLS: Final = itertools.count(1) + + +def _provider(request: Request) -> Reply: + if not request.body: + return Reply(status=404, body=b'{"error":"synthetic provider: no body"}') + marker: Final = _marker_in(request.body) + ident: Final = f"{marker}-{next(_PROVIDER_CALLS)}" + body: Final = json.loads(request.body) + answer: Final = "synthetic answer " + marker + (" " + BLOCK_MARK if "ANSWER-BLOCK" in request.body.decode() else "") + streaming: Final = bool(body.get("stream")) + if request.target.endswith("/v1/messages"): + return ( + Reply(chunks=_messages_chunks(ident, answer), content_type="text/event-stream") + if streaming + else Reply(body=_messages_body(ident, answer)) + ) + if request.target.endswith("/v1/responses"): + return Reply(body=_responses_body(ident, answer)) + if request.target.endswith("/v1/completions"): + return Reply(body=_completion_body(ident, answer)) + assert request.target.endswith("/v1/chat/completions"), request.target + return ( + Reply(chunks=_chat_chunks(ident, answer), content_type="text/event-stream") + if streaming + else Reply(body=_chat_body(ident, answer)) + ) + + +def _guardrail(name: str, key: str, url: str, mode: str, default_on: bool, **params: object) -> dict[str, object]: + return { + "guardrail_name": name, + "litellm_params": { + "guardrail": "straiker", + "mode": mode, + "default_on": default_on, + "api_key": key, + "api_base": url, + "max_retries": 0, + **params, + }, + } + + +def _rig_config(sink_url: str, root: Path) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["cache"] = False + config["guardrails"] = [ + _guardrail("straiker-v3", V3_KEY, sink_url, "pre_call", True, agent_ref=AUDIT_AGENT), + _guardrail("straiker-v3-post", V3_KEY, sink_url, "post_call", False, agent_ref=POST_AGENT), + _guardrail("straiker-v3-log", V3_KEY, sink_url, "logging_only", False, agent_ref=LOG_AGENT), + _guardrail("straiker-v3-open", V3_KEY, sink_url, "pre_call", False, fail_on_error=False, agent_ref=OPEN_AGENT), + _guardrail( + "straiker-v3-hint", + V3_KEY, + sink_url, + "pre_call", + False, + client="named-client", + format_hint="anthropic.messages", + ), + _guardrail("straiker-v3-as-v1", V3_KEY, sink_url, "pre_call", False, api_version="v1"), + _guardrail("straiker-v1", V1_KEY, sink_url, "pre_call", False), + _guardrail("straiker-v1-post", V1_KEY, sink_url, "post_call", False), + ] + path: Final = root / "straiker.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +@dataclass(frozen=True, slots=True) +class Rig: + proxy: Gateway + owned: OwnedProxy + sink: Sink + provider_url: str + provider_drain: Callable[[], tuple[Request, ...]] + chat_model: str + anthropic_model: str + completion_model: str + + def marker(self) -> str: + return "mark-" + uuid.uuid4().hex + + def _base(self) -> str: + return str(self.proxy.client.base_url).rstrip("/") + + def openai(self, key: str | None = None) -> openai.OpenAI: + return openai.OpenAI(base_url=self._base() + "/v1", api_key=key or self.proxy.key, max_retries=0) + + def async_openai(self, key: str | None = None) -> openai.AsyncOpenAI: + return openai.AsyncOpenAI(base_url=self._base() + "/v1", api_key=key or self.proxy.key, max_retries=0) + + def anthropic(self) -> anthropic.Anthropic: + return anthropic.Anthropic(base_url=self._base(), api_key=self.proxy.key, max_retries=0) + + def async_anthropic(self) -> anthropic.AsyncAnthropic: + return anthropic.AsyncAnthropic(base_url=self._base(), api_key=self.proxy.key, max_retries=0) + + def sink_calls(self, marker: str) -> tuple[Seen, ...]: + return self.sink.for_marker(marker) + + def provider_calls(self, marker: str, requests: tuple[Request, ...]) -> tuple[Request, ...]: + return tuple(r for r in requests if marker.encode() in r.body) + + def spend_row(self, request_id: str) -> dict[str, object]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id, model, call_type, metadata FROM "LiteLLM_SpendLogs" WHERE request_id = %s', + (request_id,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + return rows[0] + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Rig]: + root: Final = tmp_path_factory.mktemp("straiker") + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + port: Final = reserve.getsockname()[1] + sink: Final = Sink(port) + sink.start() + with gateway_from_environment() as gateway, wire_server(_provider) as provider: + config: Final = _rig_config(sink.url, root) + with ( + owned_proxy_process(gateway, root, {}, config=config, workers=2) as owned, + owned.gateway.scenario() as scenario, + ): + chat: Final = scenario.model( + model="openai/gpt-4o-mini", api_base=provider.url + "/v1", api_key="synthetic-openai-key" + ) + claude: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", api_base=provider.url, api_key="synthetic-anthropic-key" + ) + completion: Final = scenario.model( + model="text-completion-openai/gpt-3.5-turbo-instruct", + api_base=provider.url + "/v1", + api_key="synthetic-openai-key", + ) + yield Rig(owned.gateway, owned, sink, provider.url, provider.drain, chat, claude, completion) + if sink.server is not None: + sink.stop() + + +def _messages(text: str, system: str | None = None) -> list[dict[str, object]]: + return ([{"role": "system", "content": system}] if system else []) + [{"role": "user", "content": text}] + + +def _chat( + rig: Rig, text: str, *, key: str | None = None, headers: dict[str, str] | None = None, **extra: object +) -> httpx.Response: + return rig.proxy.client.post( + "/v1/chat/completions", + json={"model": rig.chat_model, "messages": _messages(text), **extra}, + headers={"Authorization": f"Bearer {key or rig.proxy.key}", **(headers or {})}, + ) + + +def _v3_request_calls(rig: Rig, marker: str, agent: str | None = AUDIT_AGENT) -> tuple[Seen, ...]: + return tuple( + s + for s in rig.sink_calls(marker) + if s.target == V3_PATH and "straiker_phase" not in s.body and s.headers.get("x-s6r-agent") == agent + ) + + +def _v3_response_calls(rig: Rig, marker: str, agent: str | None = POST_AGENT) -> tuple[Seen, ...]: + return tuple( + s + for s in rig.sink_calls(marker) + if s.target == V3_PATH + and s.body.get("straiker_phase") == "response-sync" + and s.headers.get("x-s6r-agent") == agent + ) + + +def _v1_calls(rig: Rig, marker: str, key: str) -> tuple[Seen, ...]: + return tuple( + s for s in rig.sink_calls(marker) if s.target == V1_PATH and s.headers.get("authorization") == "Bearer " + key + ) + + +# H1: default_on v3 pre_call, OpenAI SDK sync, non-streaming +def test_v3_pre_call_allow_relays_provider_body_and_key_identity(rig: Rig) -> None: + marker: Final = rig.marker() + with rig.proxy.scenario() as scenario: + key: Final = scenario.key(key_alias="alias-" + marker, metadata={"user_api_key_user_email": "n/a"}) + response: Final = rig.openai(key).chat.completions.create( + model=rig.chat_model, + messages=[{"role": "user", "content": "hello " + marker}], + temperature=0.2, + user="end-" + marker, + ) + assert response.id.startswith("chatcmpl-" + marker), response.id + assert response.choices[0].message.content == "synthetic answer " + marker + calls: Final = _v3_request_calls(rig, marker) + assert len(calls) == 1, calls + sent: Final = calls[0] + assert sent.headers["authorization"] == "Bearer " + V3_KEY + assert "x-straiker-webhook-format" not in sent.headers + assert sent.headers["x-s6r-agent"] == "audit-agent" + assert sent.body["messages"] == [{"role": "user", "content": "hello " + marker}] + assert sent.body["temperature"] == 0.2 + assert sent.body["model"] == rig.chat_model + assert "api_key" not in sent.body and "synthetic-openai-key" not in json.dumps(sent.body) + assert object_value(sent.body["metadata"])["user_api_key_alias"] == "alias-" + marker + assert sent.body.get("session_id", "").startswith("litellm-") + upstream: Final = rig.provider_calls(marker, rig.provider_drain()) + assert len(upstream) == 1 and upstream[0].target == "/v1/chat/completions" + row: Final = rig.spend_row(response.id) + assert row["model"] == "openai/gpt-4o-mini", row + + +# H2: v3 block verdict on the request phase blocks with the platform's message +def test_v3_block_verdict_returns_400_with_block_message_and_no_provider_call(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = _chat(rig, f"{BLOCK_MARK} {marker}") + assert response.status_code == 400, response.text + assert response.json()["error"]["message"] == BLOCK_MESSAGE, response.text + assert len(_v3_request_calls(rig, marker)) == 1 + assert rig.provider_calls(marker, rig.provider_drain()) == () + + +# H3: a resend of a blocked conversation is blocked by the process that saw the block, without a second detect call +def test_v3_blocked_conversation_replays_block_without_asking_again(rig: Rig) -> None: + marker: Final = rig.marker() + session: Final = {"x-claude-code-session-id": "session-" + marker} + first: Final = _chat(rig, f"{BLOCK_MARK} {marker}", headers=session) + assert first.status_code == 400, first.text + baseline: Final = len(_v3_request_calls(rig, marker)) + assert baseline == 1 + outcomes: Final = tuple(_chat(rig, f"{BLOCK_MARK} {marker}", headers=session) for _ in range(6)) + assert all(r.status_code == 400 and r.json()["error"]["message"] == BLOCK_MESSAGE for r in outcomes), [ + r.text for r in outcomes + ] + later: Final = len(_v3_request_calls(rig, marker)) + # Two workers: only the worker that saw the block replays from memory, the other asks Straiker once + assert baseline <= later <= 2, later + grown: Final = rig.proxy.client.post( + "/v1/chat/completions", + json={ + "model": rig.chat_model, + "messages": _messages(f"{BLOCK_MARK} {marker}") + + [{"role": "assistant", "content": "x"}, {"role": "user", "content": "more"}], + }, + headers={"Authorization": f"Bearer {rig.proxy.key}", **session}, + ) + assert grown.status_code == 400, grown.text + assert rig.provider_calls(marker, rig.provider_drain()) == () + + +# H4: a kill-switch block (no blocked_by) blocks but is not remembered, so Straiker is asked every time +def test_v3_killswitch_block_is_not_remembered(rig: Rig) -> None: + marker: Final = rig.marker() + session: Final = {"x-claude-code-session-id": "session-" + marker} + outcomes: Final = tuple(_chat(rig, f"{KILL_MARK} {marker}", headers=session) for _ in range(3)) + assert all(r.status_code == 400 and r.json()["error"]["message"] == BLOCK_MESSAGE for r in outcomes) + assert len(_v3_request_calls(rig, marker)) == 3 + + +# H4b: a deny decision on the flat envelope also blocks, with the deny_reason +def test_v3_flat_deny_decision_blocks_with_deny_reason(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = _chat(rig, f"{DENY_MARK} {marker}") + assert response.status_code == 400, response.text + assert response.json()["error"]["message"] == DENY_MESSAGE + + +# H5: post_call non-streaming, selected per request, async OpenAI SDK +@pytest.mark.asyncio +async def test_v3_post_call_sends_response_phase_with_answer(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = await rig.async_openai().chat.completions.create( + model=rig.chat_model, + messages=[{"role": "user", "content": "post " + marker}], + extra_body={"guardrails": ["straiker-v3-post"]}, + ) + assert response.id.startswith("chatcmpl-" + marker), response.id + calls: Final = eventually(lambda: _v3_response_calls(rig, marker), lambda c: len(c) == 1) + phase: Final = calls[0].body + assert phase["model"] == "gpt-4o-mini", "the deployment's model, not the alias" + assert object_value(phase["request"])["messages"] == [{"role": "user", "content": "post " + marker}] + assert json.loads(str(phase["sse"]))["id"].startswith("chatcmpl-" + marker) + assert json.loads(str(phase["sse"]))["choices"][0]["message"]["content"] == "synthetic answer " + marker + assert len(_v3_request_calls(rig, marker)) == 1, "the default_on pre_call route still runs beside the selected one" + assert rig.spend_row(response.id)["request_id"] == response.id + + +# H5b: post_call block replaces the answer with the block message as a 200 +def test_v3_post_call_block_replaces_answer(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = _chat(rig, "ANSWER-BLOCK " + marker, guardrails=["straiker-v3-post"]) + assert response.status_code == 200, response.text + assert response.json()["choices"][0]["message"]["content"] == BLOCK_MESSAGE, response.text + assert len(_v3_response_calls(rig, marker)) == 1 + + +# H6: post_call streaming, OpenAI SDK sync; the stream is consumed to the end before the phase is sent +def test_v3_post_call_streaming_sends_assembled_answer(rig: Rig) -> None: + marker: Final = rig.marker() + stream: Final = rig.openai().chat.completions.create( + model=rig.chat_model, + messages=[{"role": "user", "content": "stream " + marker}], + stream=True, + extra_body={"guardrails": ["straiker-v3-post"]}, + ) + chunks: Final = list(stream) + assert chunks and all(c.id.startswith("chatcmpl-" + marker) for c in chunks) + text: Final = "".join(c.choices[0].delta.content or "" for c in chunks if c.choices) + assert text == "synthetic answer " + marker + calls: Final = eventually(lambda: _v3_response_calls(rig, marker), lambda c: len(c) == 1) + sse: Final = json.loads(str(calls[0].body["sse"])) + assert "synthetic answer " + marker in json.dumps(sse) + assert object_value(calls[0].body["request"])["stream"] is True + + +# H7: Anthropic Messages sync, pre_call, session header and recognised client +def test_v3_anthropic_messages_relays_system_and_routing_headers(rig: Rig) -> None: + marker: Final = rig.marker() + client: Final = rig.anthropic().with_options( + default_headers={"x-claude-code-session-id": "cc-" + marker, "User-Agent": "claude-cli/2.0.0 (external, cli)"} + ) + response: Final = client.messages.create( + model=rig.anthropic_model, + max_tokens=16, + system="synthetic system " + marker, + messages=[{"role": "user", "content": "anthropic " + marker}], + ) + assert response.id.startswith("msg_" + marker), response.id + assert response.content[0].text == "synthetic answer " + marker + calls: Final = _v3_request_calls(rig, marker) + assert len(calls) == 1, calls + sent: Final = calls[0] + assert sent.headers["x-claude-code-session-id"] == "cc-" + marker + assert sent.headers["x-s6r-client"] == "claude" + assert sent.headers["x-s6r-agent"] == "audit-agent", "YAML agent_ref wins over the User-Agent derived agent" + assert sent.body["session_id"] == "cc-" + marker + assert sent.body["system"] == "synthetic system " + marker + assert sent.body["messages"] == [{"role": "user", "content": "anthropic " + marker}] + assert sent.body["max_tokens"] == 16 + upstream: Final = rig.provider_calls(marker, rig.provider_drain()) + assert len(upstream) == 1 and upstream[0].target == "/v1/messages" + assert rig.spend_row(response.id)["call_type"] == "anthropic_messages" + + +# H8: Anthropic Messages streaming, async SDK, post_call: the answer is scored in Messages shape +@pytest.mark.asyncio +async def test_v3_anthropic_streaming_post_call_scores_messages_shaped_answer(rig: Rig) -> None: + marker: Final = rig.marker() + client: Final = rig.async_anthropic() + async with client.messages.stream( + model=rig.anthropic_model, + max_tokens=16, + messages=[{"role": "user", "content": "astream " + marker}], + extra_body={"guardrails": ["straiker-v3-post"]}, + ) as stream: + final: Final = await stream.get_final_message() + assert final.id.startswith("msg_" + marker), final.id + assert final.content[0].text == "synthetic answer " + marker + calls: Final = eventually(lambda: _v3_response_calls(rig, marker), lambda c: len(c) == 1) + sse: Final = json.loads(str(calls[0].body["sse"])) + assert sse.get("type") == "message", sse + assert sse["content"][0]["text"] == "synthetic answer " + marker + + +# H9: Responses API, raw httpx, pre_call relays `input`, `instructions`, and the answer on post_call +def test_v3_responses_api_relays_input_and_answer(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = rig.proxy.client.post( + "/v1/responses", + json={ + "model": rig.chat_model, + "input": "responses " + marker, + "instructions": "be brief", + "guardrails": ["straiker-v3", "straiker-v3-post"], + }, + headers={"Authorization": f"Bearer {rig.proxy.key}"}, + ) + assert response.status_code == 200, response.text + assert response.json()["id"].startswith("resp_"), response.text + assert "synthetic answer " + marker in response.text + pre: Final = _v3_request_calls(rig, marker) + assert len(pre) == 1 and pre[0].body["input"] == "responses " + marker and pre[0].body["instructions"] == "be brief" + post: Final = eventually(lambda: _v3_response_calls(rig, marker), lambda c: len(c) == 1) + assert "synthetic answer " + marker in str(post[0].body["sse"]) + upstream: Final = rig.provider_calls(marker, rig.provider_drain()) + assert len(upstream) == 1 and upstream[0].target == "/v1/responses" + + +# H10/H11: Completions API prompt becomes messages on the request phase; the answer is sent as a chat completion +def test_v3_completions_prompt_is_relayed_as_messages_and_answer_as_chat(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = rig.proxy.client.post( + "/v1/completions", + json={ + "model": rig.completion_model, + "prompt": "complete " + marker, + "guardrails": ["straiker-v3", "straiker-v3-post"], + }, + headers={"Authorization": f"Bearer {rig.proxy.key}"}, + ) + assert response.status_code == 200, response.text + assert response.json()["id"].startswith("cmpl-" + marker), response.json()["id"] + pre: Final = _v3_request_calls(rig, marker) + assert len(pre) == 1, pre + assert pre[0].body["messages"] == [{"role": "user", "content": "complete " + marker}] + assert "prompt" not in pre[0].body + post: Final = eventually(lambda: _v3_response_calls(rig, marker), lambda c: len(c) == 1) + sse: Final = json.loads(str(post[0].body["sse"])) + assert sse["object"] == "chat.completion", sse + assert sse["choices"][0]["message"]["content"] == "synthetic answer " + marker + + +# H12: tool and MCP server credentials are redacted one level deep; a schema property named headers is kept +def test_v3_redacts_tool_credentials_but_keeps_schema_properties(rig: Rig) -> None: + marker: Final = rig.marker() + tools: Final = [ + { + "type": "function", + "authorization": "Bearer synthetic-tool-secret", + "function": { + "name": "lookup", + "parameters": {"type": "object", "properties": {"headers": {"type": "string"}}}, + }, + } + ] + response: Final = _chat( + rig, + "tools " + marker, + tools=tools, + mcp_servers=[{"url": "http://mcp", "authorization_token": "synthetic-mcp-secret"}], + ) + assert response.status_code == 200, response.text + sent: Final = _v3_request_calls(rig, marker)[0].body + assert sent["tools"][0]["authorization"] == "[redacted]" # pyright: ignore[reportIndexIssue] # sink body is loose JSON + assert sent["tools"][0]["function"]["parameters"]["properties"]["headers"] == {"type": "string"} # pyright: ignore[reportIndexIssue] # sink body is loose JSON + assert sent["mcp_servers"][0]["authorization_token"] == "[redacted]" # pyright: ignore[reportIndexIssue] # sink body is loose JSON + assert "synthetic-tool-secret" not in json.dumps(sent) and "synthetic-mcp-secret" not in json.dumps(sent) + + +# U1: a v1 collection key still speaks the v1 webhook with the litellm envelope +def test_v1_key_keeps_webhook_envelope_and_format_header(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = _chat(rig, "v1 " + marker, guardrails=["straiker-v1"]) + assert response.status_code == 200, response.text + calls: Final = _v1_calls(rig, marker, V1_KEY) + assert len(calls) == 1, rig.sink_calls(marker) + assert calls[0].headers["x-straiker-webhook-format"] == "litellm" + assert calls[0].body["schema_version"] and object_value(calls[0].body["event"])["type"] + assert "v1 " + marker in json.dumps(object_value(calls[0].body["request"])) + assert len(_v3_request_calls(rig, marker)) == 1, "the default_on v3 route runs beside it" + + +# U2: v1 block verdict still blocks +def test_v1_block_verdict_still_blocks(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = _chat(rig, f"{V1_BLOCK_MARK} {marker}", guardrails=["straiker-v1"]) + assert response.status_code == 400, response.text + assert response.json()["error"]["message"] == BLOCK_MESSAGE + assert len(_v1_calls(rig, marker, V1_KEY)) == 1, rig.sink_calls(marker) + assert rig.provider_calls(marker, rig.provider_drain()) == () + + +# U3: v1 post_call still receives the response envelope +def test_v1_post_call_sends_response_envelope(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = _chat(rig, "v1post " + marker, guardrails=["straiker-v1-post"]) + assert response.status_code == 200, response.text + calls: Final = eventually(lambda: _v1_calls(rig, marker, V1_KEY), lambda c: len(c) == 1) + assert "synthetic answer " + marker in json.dumps(calls[0].body.get("response")) + + +# E: explicit api_version v1 with a v3-shaped key follows the configuration, not the key +def test_explicit_api_version_v1_overrides_key_prefix(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = _chat(rig, "explicit " + marker, guardrails=["straiker-v3-as-v1"]) + assert response.status_code == 200, response.text + calls: Final = _v1_calls(rig, marker, V3_KEY) + assert len(calls) == 1, rig.sink_calls(marker) + assert calls[0].headers["x-straiker-webhook-format"] == "litellm" + + +# E: configured client and format_hint ride as headers; request header for agent fills in when YAML has none +def test_v3_client_and_format_hint_headers_and_request_agent_header(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = _chat( + rig, "hint " + marker, guardrails=["straiker-v3-hint"], headers={"x-s6r-agent": "caller-agent"} + ) + assert response.status_code == 200, response.text + hinted: Final = tuple(s for s in _v3_request_calls(rig, marker, agent="caller-agent")) + assert len(hinted) == 1, rig.sink_calls(marker) + sent: Final = hinted[0] + assert sent.headers["x-s6r-client"] == "named-client" + assert sent.headers["x-s6r-format"] == "anthropic.messages" + assert sent.headers["x-s6r-agent"] == "caller-agent" + + +# E: identity precedence: the key's user email wins over an end user in the body +def test_v3_user_prefers_key_email_over_body_user(rig: Rig) -> None: + marker: Final = rig.marker() + with rig.proxy.scenario() as scenario: + user: Final = scenario.user(user_email=f"{marker}@example.test") + key: Final = scenario.key(user_id=user) + response: Final = _chat(rig, "identity " + marker, key=key, user="body-user-" + marker) + assert response.status_code == 200, response.text + sent: Final = _v3_request_calls(rig, marker)[0].body + meta: Final = object_value(sent["original"]) + assert object_value(object_value(object_value(meta["processed"])["Meta"]))["user"] == f"{marker}@example.test" + assert object_value(sent["metadata"])["user_api_key_user_email"] == f"{marker}@example.test" + assert sent["user"] == "body-user-" + marker + + +# E: logging_only observes the turn but never blocks +def test_v3_logging_only_observes_block_verdict_without_blocking(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = _chat(rig, f"{LOG_BLOCK_MARK} {marker}") + assert response.status_code == 200, response.text + assert response.json()["id"].startswith("chatcmpl-" + marker), response.json()["id"] + calls: Final = eventually(lambda: _v3_request_calls(rig, marker, agent=LOG_AGENT), lambda c: len(c) >= 1) + assert calls[0].headers["x-s6r-agent"] == LOG_AGENT + assert len(rig.provider_calls(marker, rig.provider_drain())) == 1 + row: Final = rig.spend_row(response.json()["id"]) + assert row["request_id"] == response.json()["id"] + + +# E: the same identical allowed request three times yields three detect calls and three spend rows +def test_v3_repeated_allowed_request_is_scored_and_logged_each_time(rig: Rig) -> None: + marker: Final = rig.marker() + responses: Final = tuple(_chat(rig, "repeat " + marker) for _ in range(3)) + assert all(r.status_code == 200 for r in responses), [r.text for r in responses] + ids: Final = {r.json()["id"] for r in responses} + assert len(ids) == 3 and all(i.startswith("chatcmpl-" + marker) for i in ids), ids + assert len(_v3_request_calls(rig, marker)) == 3 + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id LIKE %s', ("chatcmpl-" + marker + "%",) + ), + lambda values: len(values) == 3, + seconds=70, + ) + assert {str(r["request_id"]) for r in rows} == ids + + +# S1: platform answers 500: fail closed with the reason in the body, no provider call +def test_v3_sink_500_fails_closed_with_reason(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = _chat(rig, f"{SINK_500_MARK} {marker}") + assert response.status_code == 400, response.text + assert "Straiker detection unavailable" in response.json()["error"]["message"], response.text + assert rig.provider_calls(marker, rig.provider_drain()) == () + + +# S1a: the v1 webhook route fails the same way when the platform answers 500 +def test_v1_sink_500_fails_closed_with_reason(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = _chat(rig, f"{V1_500_MARK} {marker}", guardrails=["straiker-v1"]) + assert response.status_code == 400, response.text + assert "Straiker detection unavailable" in response.json()["error"]["message"], response.text + assert len(_v1_calls(rig, marker, V1_KEY)) == 1 + assert rig.provider_calls(marker, rig.provider_drain()) == () + + +# S1b: fail_on_error false lets the request through on a 500 +def test_v3_fail_open_guardrail_passes_on_sink_500(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = _chat(rig, f"{OPEN_500_MARK} {marker}", guardrails=["straiker-v3-open"]) + assert response.status_code == 200, response.text + assert response.json()["id"].startswith("chatcmpl-" + marker), response.json()["id"] + assert len(_v3_request_calls(rig, marker, agent=OPEN_AGENT)) == 1 + + +# S2: platform rejects the key: 401 is not retried and fails closed +def test_v3_sink_401_fails_closed_once(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = _chat(rig, f"{SINK_401_MARK} {marker}") + assert response.status_code == 400, response.text + assert "401" in response.json()["error"]["message"], response.text + assert len(_v3_request_calls(rig, marker)) == 1 + + +# S3: platform answers non JSON: fail closed, caller sees the parse failure +def test_v3_sink_garbage_fails_closed(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = _chat(rig, f"{SINK_GARBAGE_MARK} {marker}") + assert response.status_code == 400, response.text + assert "Straiker detection unavailable" in response.json()["error"]["message"] + + +# S4: unauthenticated request never reaches the platform +def test_unauthenticated_request_does_not_reach_platform(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = _chat(rig, "anon " + marker, key="sk-not-a-real-key") + assert response.status_code == 401, response.text + assert rig.sink_calls(marker) == () + + +# S5: unknown model: the guardrail still runs, then the router error reaches the caller +def test_unknown_model_error_reaches_caller_after_detect(rig: Rig) -> None: + marker: Final = rig.marker() + response: Final = rig.proxy.client.post( + "/v1/chat/completions", + json={"model": "no-such-model-" + marker, "messages": _messages("unknown " + marker)}, + headers={"Authorization": f"Bearer {rig.proxy.key}"}, + ) + assert response.status_code in (400, 401, 404), response.text + assert "no-such-model-" + marker in response.text + assert len(_v3_request_calls(rig, marker)) == 1, rig.sink_calls(marker) + assert rig.provider_calls(marker, rig.provider_drain()) == () + + +# S6: odd shapes in the routing header and a 5 KB prompt are relayed verbatim, not crashed on +def test_v3_oversized_prompt_and_odd_header_values_are_relayed(rig: Rig) -> None: + marker: Final = rig.marker() + big: Final = "x" * 5000 + " " + marker + response: Final = _chat(rig, big, headers={"x-claude-code-session-id": "", "x-s6r-agent": "1"}) + assert response.status_code == 200, response.text + sent: Final = _v3_request_calls(rig, marker)[0] + assert sent.body["messages"] == [{"role": "user", "content": big}] + assert sent.headers["x-s6r-agent"] == "audit-agent" + assert "x-claude-code-session-id" not in sent.headers + assert sent.body.get("session_id", "").startswith("litellm-") + + +# S7: a guardrail with a malformed format_hint is rejected at /guardrails/apply_guardrail time, not at boot +def test_malformed_format_hint_config_is_rejected_by_guardrail_management(rig: Rig) -> None: + response: Final = rig.proxy.client.post( + "/guardrails", + json={ + "guardrail": { + "guardrail_name": "straiker-bad-" + uuid.uuid4().hex, + "litellm_params": { + "guardrail": "straiker", + "mode": "pre_call", + "api_key": V3_KEY, + "api_base": rig.sink.url, + "format_hint": "bogus", + }, + } + }, + headers={"Authorization": f"Bearer {rig.proxy.key}"}, + ) + assert response.status_code in (400, 422, 500), response.text + assert "format_hint" in response.text or "bogus" in response.text, response.text + healthy: Final = _chat(rig, "still-fine " + uuid.uuid4().hex) + assert healthy.status_code == 200, healthy.text + + +# S8: /key/health reports the key without touching the platform +def test_key_health_does_not_call_platform(rig: Rig) -> None: + marker: Final = rig.marker() + with rig.proxy.scenario() as scenario: + key: Final = scenario.key(key_alias="health-" + marker) + response: Final = rig.proxy.client.post("/key/health", headers={"Authorization": f"Bearer {key}"}) + assert response.status_code == 200, response.text + assert response.json()["key"] == "healthy" + assert rig.sink_calls(marker) == () + + +# C1: 30 request mixed burst while the platform sink is down mid burst, then recovers; every allowed id lands once +def test_burst_with_platform_outage_recovers_without_duplicate_spend(rig: Rig) -> None: + burst: Final = 30 + markers: Final = tuple(rig.marker() for _ in range(burst)) + down: Final = threading.Event() + up: Final = threading.Event() + + def call(index: int) -> tuple[int, int, str]: + if index == 8: + rig.sink.stop() + down.set() + if index == 20: + assert down.wait(10) + rig.sink.start() + up.set() + marker: Final = markers[index] + if index % 3 == 0: + response: Final = rig.proxy.client.post( + "/v1/messages", + json={ + "model": rig.anthropic_model, + "max_tokens": 8, + "messages": [{"role": "user", "content": "burst " + marker}], + }, + headers={"Authorization": f"Bearer {rig.proxy.key}"}, + ) + return index, response.status_code, response.text + streaming: Final = index % 2 == 1 + response = _chat(rig, "burst " + marker, stream=streaming) + return index, response.status_code, response.text + + with ThreadPoolExecutor(max_workers=6) as pool: + results: Final = sorted(pool.map(call, range(burst))) + assert up.is_set() + statuses: Final = {index: status for index, status, _ in results} + assert all(status in (200, 400) for status in statuses.values()), results + failed: Final = tuple(index for index, status, text in results if status == 400) + assert failed, "the outage must be visible to at least one caller" + assert all("Straiker detection unavailable" in text for index, status, text in results if status == 400), results + for index, status, text in results: + if status != 200 or (index % 3 != 0 and index % 2 == 1): + continue + marker = markers[index] + expected: Final = ("msg_" if index % 3 == 0 else "chatcmpl-") + marker + "%" + rows: Final = eventually( + lambda like=expected: read_rows( + 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id LIKE %s', (like,) + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert len(rows) == 1, rows + provider_seen: Final = rig.provider_drain() + for index, status, _ in results: + if status == 400: + assert rig.provider_calls(markers[index], provider_seen) == (), ( + "a failed-closed turn must not reach the provider" + ) + recovered: Final = _chat(rig, "after-outage " + rig.marker()) + assert recovered.status_code == 200, recovered.text + + +# C2: one proxy worker is killed during a burst; the other keeps serving and detect still runs for each call +def _uvicorn_workers(parent: psutil.Process, *, exclude: int = 0) -> tuple[psutil.Process, ...]: + return tuple( + c for c in parent.children() if c.is_running() and c.pid != exclude and "spawn_main" in " ".join(c.cmdline()) + ) + + +def test_burst_survives_one_worker_kill(rig: Rig) -> None: + parent: Final = psutil.Process(rig.owned.process.pid) + workers: Final = eventually(lambda: _uvicorn_workers(parent), lambda c: len(c) >= 2) + victim: Final = workers[0].pid + markers: Final = tuple(rig.marker() for _ in range(24)) + + def fresh_chat(text: str) -> tuple[int, str]: + with httpx.Client(base_url=rig._base(), timeout=15, trust_env=False) as fresh: + try: + response: Final = fresh.post( + "/v1/chat/completions", + json={"model": rig.chat_model, "messages": _messages(text)}, + headers={"Authorization": f"Bearer {rig.proxy.key}"}, + ) + except httpx.TransportError as error: + return 0, repr(error) + return response.status_code, response.text + + def call(index: int) -> tuple[int, str]: + if index == 6: + os.kill(victim, signal.SIGKILL) + return fresh_chat("kill " + markers[index]) + + with ThreadPoolExecutor(max_workers=4) as pool: + results: Final = tuple(pool.map(call, range(24))) + ok: Final = tuple(i for i, (status, _) in enumerate(results) if status == 200) + assert len(ok) >= 20, results + for index in ok: + assert len(_v3_request_calls(rig, markers[index])) >= 1, markers[index] + eventually(lambda: _uvicorn_workers(parent, exclude=victim), lambda c: len(c) >= 2) + after: Final = fresh_chat("after-kill " + rig.marker()) + assert after[0] == 200, after + + +# C3: proxy restart between a blocked turn and its replay: the memory is per process and empties, so Straiker is asked again +def test_proxy_restart_forgets_blocked_turns_and_asks_platform_again(tmp_path: Path, rig: Rig) -> None: + with gateway_from_environment() as gateway: + config: Final = _rig_config(rig.sink.url, tmp_path) + marker: Final = rig.marker() + session: Final = {"x-claude-code-session-id": "restart-" + marker} + body: Final = {"model": rig.chat_model, "messages": _messages(f"{BLOCK_MARK} {marker}")} + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as first: + blocked: Final = first.client.post( + "/v1/chat/completions", json=body, headers={"Authorization": f"Bearer {first.key}", **session} + ) + assert blocked.status_code == 400, blocked.text + replayed: Final = first.client.post( + "/v1/chat/completions", json=body, headers={"Authorization": f"Bearer {first.key}", **session} + ) + assert replayed.status_code == 400, replayed.text + assert len(_v3_request_calls(rig, marker)) == 1, "one worker replays from memory" + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as second: + again: Final = second.client.post( + "/v1/chat/completions", json=body, headers={"Authorization": f"Bearer {second.key}", **session} + ) + assert again.status_code == 400, again.text + assert len(_v3_request_calls(rig, marker)) == 2, "a restarted process has no memory and asks once more" From 6c8afb221fd6a8cfc1cac6d1cce0f0d928ac2336 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 19:04:00 -0500 Subject: [PATCH 02/10] fix(otel): root post-response service spans in their own trace linked to the request (#42826) * fix(otel): root post-response service spans in their own trace linked to the request Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(otel): trim service span context docstring Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/integrations/otel/README.md | 19 ++++- litellm/integrations/otel/logger.py | 11 ++- litellm/integrations/otel/plumbing/context.py | 23 +++++ .../integrations/otel/test_otel_v2_logger.py | 85 +++++++++++++++++++ 4 files changed, 133 insertions(+), 5 deletions(-) diff --git a/litellm/integrations/otel/README.md b/litellm/integrations/otel/README.md index 023caf06d12..d8dfabe23d6 100644 --- a/litellm/integrations/otel/README.md +++ b/litellm/integrations/otel/README.md @@ -63,7 +63,24 @@ Spans are named `"{service} {call_type}"` (e.g. `"redis set"`) so repeated calls to one service stay distinguishable. Like every other span they parent to the **ambient** context, falling back to the threaded `litellm_parent_otel_span` only when ambient has no live span; a background job with neither starts its own root -trace. Caller-supplied `event_metadata` is **sanitized** before it reaches a span +trace. + +**Post-response work is its own trace.** Spend tracking, the response cache write +and the spend-counter increment all run after the response is on the wire, so they +add nothing to the request's latency. Parenting them under the (already ended) +server span stretched the request trace past the request itself, which is what a +viewer shows as trace duration. `context.resolve_service_span_context` compares +the call's end time with the resolved parent's end time: a call that finished +after its parent ended starts a **new root trace** carrying a **span link** back +to the request span (the `FollowsFrom` relationship of OpenTracing; the default +`:link` propagation style of the OTel Ruby ActiveJob and Sidekiq +instrumentations). Identity Baggage still rides along, so the detached span keeps +its team / key / user attributes. Only an SDK span that has really ended detaches: +a sampled-out or remote `NonRecordingSpan` is never recording but is still the +right parent. A call that ended before the server span did stays a child even when +its `asyncio.create_task`-dispatched hook runs after the response. + +Caller-supplied `event_metadata` is **sanitized** before it reaches a span (primitives only, no live objects, no secrets/headers, bounded) — see `payloads.sanitize_event_metadata`. diff --git a/litellm/integrations/otel/logger.py b/litellm/integrations/otel/logger.py index bab0d7ec092..0466e00a959 100644 --- a/litellm/integrations/otel/logger.py +++ b/litellm/integrations/otel/logger.py @@ -56,8 +56,8 @@ from litellm.integrations.otel.plumbing.context import ( request_root_http_route, request_root_span, resolve_mcp_span_context, - resolve_parent_context, resolve_request_span_context, + resolve_service_span_context, set_request_baggage, set_request_root_span, ) @@ -671,14 +671,17 @@ class OpenTelemetryV2(CustomLogger): # rides along and the call nests under whatever request phase is active — # e.g. a DB lookup under the live ``auth`` span), falling back to the # server span the proxy threaded as ``parent_otel_span``. A background - # service call has neither, so it starts its own root trace. - parent_context: Final = resolve_parent_context(threaded=parent_otel_span) + # service call has neither, so it starts its own root trace, as does one + # that finished after the request span ended (linked back to it). + end_time_ns: Final = to_ns(end_time) + parent_context, links = resolve_service_span_context(threaded=parent_otel_span, end_time_ns=end_time_ns) return self._emitter.emit( role, data, parent_context=parent_context, start_time_ns=to_ns(start_time), - end_time_ns=to_ns(end_time), + end_time_ns=end_time_ns, + links=links, ) # ====================================================================== # diff --git a/litellm/integrations/otel/plumbing/context.py b/litellm/integrations/otel/plumbing/context.py index 19243d64c64..9de5c1ac1cb 100644 --- a/litellm/integrations/otel/plumbing/context.py +++ b/litellm/integrations/otel/plumbing/context.py @@ -9,6 +9,7 @@ from opentelemetry import baggage from opentelemetry.context import Context, get_current from opentelemetry.sdk.trace import ReadableSpan from opentelemetry.trace import ( + INVALID_SPAN, Link, NonRecordingSpan, Span, @@ -225,6 +226,28 @@ def resolve_parent_context(threaded: Span | None = None) -> Context: return ctx +def resolve_service_span_context( + threaded: Span | None = None, end_time_ns: int | None = None +) -> tuple[Context, tuple[Link, ...]]: + """Parent context + links for a service/DB span that ended at ``end_time_ns``. + + A call that finished after its parent ended (post-response spend tracking) + starts its own root trace with a span link back to the parent instead of + stretching the parent's trace. Baggage stays on the returned context. + """ + ctx: Final = resolve_parent_context(threaded) + parent: Final = get_current_span(ctx) + if not _ended_before(parent, end_time_ns): + return ctx, () + return set_span_in_context(INVALID_SPAN, ctx), (Link(parent.get_span_context()),) + + +def _ended_before(span: Span, end_time_ns: int | None) -> bool: + if not isinstance(span, ReadableSpan) or span.end_time is None: + return False + return end_time_ns is None or end_time_ns > span.end_time + + def resolve_request_span_context() -> Context: """The parent context for a request-level span (the LLM call, a guardrail). diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py b/tests/test_litellm/integrations/otel/test_otel_v2_logger.py index 287f15a7183..00c1343f72e 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_logger.py @@ -2001,6 +2001,91 @@ def test_service_span_prefers_ambient_context_over_threaded_parent(): assert by_name["redis get"].parent.span_id == ambient.get_span_context().span_id +_REQUEST_END = 1_000.0 + + +def _ended_request_span(logger): + """A PROXY_REQUEST span whose response already went out at ``_REQUEST_END``.""" + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) + server.end(end_time=to_ns(_REQUEST_END)) + return server + + +@pytest.mark.parametrize("parent_source", ["ambient", "threaded"]) +def test_service_call_that_outlives_the_request_roots_its_own_trace_linked_to_the_request(parent_source): + """Post-response work (spend tracking, the cache write, the spend-counter + increment) finishes after the server span ended, so it did not add to the + request's latency. Nesting it under the request would stretch the request + trace past the response, so it starts its own trace and keeps the request + reachable through a span link, whether the request span is the ambient + context or the threaded ``parent_otel_span``.""" + logger, exporter = _logger() + server = _ended_request_span(logger) + hook = logger.async_service_success_hook( + payload=_ServicePayload("batch_write_to_db", "_PROXY_track_cost_callback"), + parent_otel_span=server if parent_source == "threaded" else None, + start_time=_REQUEST_END + 0.1, + end_time=_REQUEST_END + 0.5, + ) + if parent_source == "ambient": + with trace.use_span(server, end_on_exit=False): + asyncio.run(hook) + else: + asyncio.run(hook) + by_name = {s.name: s for s in exporter.get_finished_spans()} + span = by_name["batch_write_to_db _PROXY_track_cost_callback"] + request_ctx = server.get_span_context() + assert span.parent is None + assert span.context.trace_id != request_ctx.trace_id + assert [(link.context.trace_id, link.context.span_id) for link in span.links] == [ + (request_ctx.trace_id, request_ctx.span_id) + ] + + +def test_service_call_that_finished_before_the_response_stays_in_the_request_trace(): + """The hook is dispatched with ``asyncio.create_task`` and can run after the + response went out even though the call itself completed during the request. + Its own end time decides: a call that ended before the request span did is + request latency and stays a child of the request.""" + logger, exporter = _logger() + server = _ended_request_span(logger) + asyncio.run( + logger.async_service_success_hook( + payload=_ServicePayload("postgres", "get_data"), + parent_otel_span=server, + start_time=_REQUEST_END - 0.5, + end_time=_REQUEST_END - 0.1, + ) + ) + span = {s.name: s for s in exporter.get_finished_spans()}["postgres get_data"] + assert span.parent.span_id == server.get_span_context().span_id + assert span.context.trace_id == server.get_span_context().trace_id + assert list(span.links) == [] + + +def test_service_call_under_a_remote_parent_is_never_detached(): + """A propagated parent is a ``NonRecordingSpan`` with no end time of its own. + Not recording is not the same as ended, so the call stays its child.""" + from opentelemetry.trace import NonRecordingSpan, SpanContext, TraceFlags + + logger, exporter = _logger() + remote = NonRecordingSpan( + SpanContext(trace_id=0xABC, span_id=0x123, is_remote=True, trace_flags=TraceFlags(TraceFlags.SAMPLED)) + ) + asyncio.run( + logger.async_service_success_hook( + payload=_ServicePayload("redis", "get"), + parent_otel_span=remote, + start_time=_REQUEST_END + 0.1, + end_time=_REQUEST_END + 0.5, + ) + ) + span = {s.name: s for s in exporter.get_finished_spans()}["redis get"] + assert span.parent.span_id == 0x123 + assert span.context.trace_id == 0xABC + assert list(span.links) == [] + + # --------------------------------------------------------------------------- # # Proxy SERVER span lifecycle # --------------------------------------------------------------------------- # From 5c738cc3cd609b59c7028434114a5eba1aebcce8 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 17:17:52 -0700 Subject: [PATCH 03/10] Revert "docs: simplify pull request template into plain English questions (#42813)" (#42828) This reverts commit deab92408ee6e29f714cbc18d6f659f667260c35. Co-authored-by: kerry --- .github/pull_request_template.md | 161 ++++++++++++++++++++++++++++--- AGENTS.md | 2 +- 2 files changed, 149 insertions(+), 14 deletions(-) diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index 12dff300afe..7a9883df356 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -1,27 +1,162 @@ - + -## What's the problem? +## TLDR -## What's the solution? + - +Problem this solves: -## How does it fix it? +- +- ... - +How it solves it: -## How does the product experience change? +- +- ... - +## User Flow -## What caveats are there, if any? + +Example: + +Before: a developer whose app streams chat completions gets no token counts back, so their cost dashboard reads zero + +1. They send POST https://litellm-domain/v1/chat/completions with `"stream": true` and no `stream_options` +2. The last SSE chunk arrives with `"usage": null`, so their app records 0 prompt and 0 completion tokens +3. They open https://litellm-domain/ui/?page=logs and see the request logged at $0 spend + +After: the same request comes back with real token counts, so the dashboard shows real spend + +1. The proxy admin sets `always_include_stream_usage: true` and restarts the proxy +2. The developer sends the same POST https://litellm-domain/v1/chat/completions with `"stream": true` and no `stream_options` +3. The last SSE chunk now carries a `usage` object with real prompt and completion token counts +4. https://litellm-domain/ui/?page=logs shows that request at non-zero spend +--> + +## Relevant issues + + + +## Affected release + + ## Linear ticket - + -## How did you test this? +## Pre-Submission checklist + +**Please complete all items before asking a LiteLLM maintainer to review your PR** + +- [ ] I have added meaningful tests +- [ ] The handful of test files covering my change pass locally, e.g. `uv run pytest tests/test_litellm/.py -v`. Leave the suites (`make test-unit-*`, `make test-unit`) to CI: it finishes in ~15 minutes where a laptop takes an hour or more +- [ ] My PR passes all required CI/CD checks (e.g., lint, schema.d.ts sync check, etc.) +- [ ] My PR's scope is as isolated as possible; it only solves 1 specific problem +- [ ] I have received a Greptile **Confidence Score of at least 4/5** before requesting a maintainer review (Greptile reviews automatically once the PR is opened; only comment `@greptileai` to re-request a review after pushing changes) + +## Delays in PR merge? + +If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slack (#pr-review)](https://join.slack.com/t/litellmossslack/shared_invite/zt-3o7nkuyfr-p_kbNJj8taRfXGgQI1~YyA). + +## Screenshots / Proof of Fix + + + +## Type + + + + +🆕 New Feature +🐛 Bug Fix +🧹 Refactoring +📖 Documentation +🚄 Infrastructure +✅ Test + +## Caveats (if any) + + + +## QA runbook + + + +## Final Attestation + +- [ ] The tests check the right things, including the edge cases, and regressions in the respective real-world customer use-cases are not possible after this PR - diff --git a/AGENTS.md b/AGENTS.md index 9e0543753b0..cade08bdd02 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -39,7 +39,7 @@ Same applies for filing bug reports and feature requests, with .github/ISSUE_TEM If you're resolving a linear ticket, in the "## Linear ticket" section of the PR, say "Resolves LIT-1234", replacing "LIT-1234" with the actual ticket id that you're resolving. If you don't have the ticket id, don't make one up or search for it. Just leave the section blank -Never use `pytest` commands or the like as the answer to "How did you test this?". We prefer curl'ing a live proxy instance running on localhost:4000 (I like to run it with `python litellm/proxy/proxy_cli.py --config litellm/proxy/dev_config.yaml --detailed_debug --reload --use_v2_migration_resolver 2>&1 | tee litellm.log`; the Admin UI dev server is `npm run dev` in `ui/litellm-dashboard`, served on port 3000) and showing both the command run and the output. Also, it should hit real LLM provider APIs, not mocks, and cost real $$$ because that is the most realistic test. The proof of fix should be exactly what the end user / customer would see / do. The run logs in PR #27703 is a prime example of how to do it (not a huge fan of using a python test script that future me and the team will have no visibility into; I prefer just curl commands or a short list of bash commands (e.g., using `for`)). If it's a UI thing, or the main use case runs through a headful agentic coding tool like Claude Code or Codex, drive that surface yourself and embed your own before and after screenshots of it in the PR (the Admin UI page, or what the coding tool shows), next to an ordered list of the URLs to go to (e.g., http://localhost:4000/ui/?page=logs), where to click, and what fields to fill out so a reviewer can reproduce it +Never use `pytest` commands or the like as "Screenshots / Proof of Fix". We prefer curl'ing a live proxy instance running on localhost:4000 (I like to run it with `python litellm/proxy/proxy_cli.py --config litellm/proxy/dev_config.yaml --detailed_debug --reload --use_v2_migration_resolver 2>&1 | tee litellm.log`; the Admin UI dev server is `npm run dev` in `ui/litellm-dashboard`, served on port 3000) and showing both the command run and the output. Also, it should hit real LLM provider APIs, not mocks, and cost real $$$ because that is the most realistic test. The proof of fix should be exactly what the end user / customer would see / do. The run logs in PR #27703 is a prime example of how to do it (not a huge fan of using a python test script that future me and the team will have no visibility into; I prefer just curl commands or a short list of bash commands (e.g., using `for`)). If it's a UI thing, or the main use case runs through a headful agentic coding tool like Claude Code or Codex, drive that surface yourself and embed your own before and after screenshots of it in the PR (the Admin UI page, or what the coding tool shows), next to an ordered list of the URLs to go to (e.g., http://localhost:4000/ui/?page=logs), where to click, and what fields to fill out so a reviewer can reproduce it If you ever write any human-facing text (pull requests, issues, commit messages, discussion posts, github comments, release notes, docs, etc.), always follow these guidelines to sound less AI-y: - don't use emojis From bab86555e72a51ac8495b8a88ab0503c2904314b Mon Sep 17 00:00:00 2001 From: tin-berri Date: Wed, 23 Sep 2026 17:32:22 -0700 Subject: [PATCH 04/10] fix(ui): hide LiteAdmin while Playground is open (#42755) --- .../(dashboard)/hooks/useDisableLiteAdmin.ts | 34 ++++++++ .../src/app/(dashboard)/layout.test.tsx | 33 +++++++- .../src/app/(dashboard)/layout.tsx | 5 +- .../Navbar/UserDropdown/UserDropdown.tsx | 25 +++++- .../SidebarAccountMenu/SidebarAccountMenu.tsx | 26 +++++- .../liteadmin/LiteAdmin.integration.test.tsx | 81 ++++++++++++++++++- .../src/components/liteadmin/LiteAdmin.tsx | 4 +- 7 files changed, 201 insertions(+), 7 deletions(-) create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableLiteAdmin.ts diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableLiteAdmin.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableLiteAdmin.ts new file mode 100644 index 00000000000..cbc3a5a81f6 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableLiteAdmin.ts @@ -0,0 +1,34 @@ +import { useSyncExternalStore } from "react"; +import { getProxyBaseUrl } from "@/components/networking"; +import { + LOCAL_STORAGE_EVENT, + emitLocalStorageChange, + getLocalStorageItem, + removeLocalStorageItem, + setLocalStorageItem, +} from "@/utils/localStorageUtils"; + +function subscribe(callback: () => void) { + window.addEventListener("storage", callback); + window.addEventListener(LOCAL_STORAGE_EVENT, callback); + return () => { + window.removeEventListener("storage", callback); + window.removeEventListener(LOCAL_STORAGE_EVENT, callback); + }; +} + +export function useDisableLiteAdmin(userId: string | null) { + const key = userId ? `disableLiteAdmin:${JSON.stringify([getProxyBaseUrl(), userId])}` : null; + const disabled = useSyncExternalStore( + subscribe, + () => key !== null && getLocalStorageItem(key) === "true", + () => false, + ); + const setDisabled = (value: boolean) => { + if (key === null) return; + if (value) setLocalStorageItem(key, "true"); + else removeLocalStorageItem(key); + emitLocalStorageChange(key); + }; + return [disabled, setDisabled] as const; +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/layout.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/layout.test.tsx index d854befa197..3b52a2eac33 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/layout.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/layout.test.tsx @@ -1,5 +1,6 @@ import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; import { render, screen, waitFor } from "@testing-library/react"; +import { usePathname } from "next/navigation"; import { AuthProvider } from "@/contexts/AuthContext"; import Layout from "./layout"; @@ -10,7 +11,11 @@ let searchParamsValue = new URLSearchParams(); vi.mock("next/navigation", () => ({ useRouter: vi.fn(() => ({ push: vi.fn(), replace: replaceMock })), useSearchParams: vi.fn(() => searchParamsValue), - usePathname: vi.fn(() => "/ui/guardrails"), + usePathname: vi.fn(), +})); + +vi.mock("@/components/liteadmin/LiteAdmin", () => ({ + default: () => , })); vi.mock("@/components/DashboardHeader", () => ({ @@ -79,8 +84,34 @@ describe("(dashboard) Layout", () => { vi.clearAllMocks(); pendingUiConfig = createDeferred(); searchParamsValue = new URLSearchParams(); + vi.mocked(usePathname).mockReturnValue("/ui/guardrails"); }); + it.each(["/ui/playground", "/ui/playground/"])( + "hides LiteAdmin on %s and restores it after leaving Playground", + async (pathname) => { + const dashboard = () => ( + + +
+ + + ); + const { rerender } = render(dashboard()); + pendingUiConfig.resolve(); + expect(await screen.findByRole("button", { name: "LiteAdmin" })).toBeInTheDocument(); + + vi.mocked(usePathname).mockReturnValue(pathname); + rerender(dashboard()); + expect(screen.queryByRole("button", { name: "LiteAdmin" })).not.toBeInTheDocument(); + expect(screen.getByTestId("page-content")).toBeInTheDocument(); + + vi.mocked(usePathname).mockReturnValue("/ui/api-keys"); + rerender(dashboard()); + expect(screen.getByRole("button", { name: "LiteAdmin" })).toBeInTheDocument(); + }, + ); + it("does not mount route content until getUiConfig has resolved", async () => { render( diff --git a/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx b/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx index 03612cee3c6..406a323fbfb 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx @@ -15,7 +15,7 @@ import { LicenseExpiryBanner } from "@/components/LicenseExpiryBanner"; import { UserBanner } from "@/components/UserBanner"; import LiteAdmin from "@/components/liteadmin/LiteAdmin"; import { UpgradeBanner } from "@/components/UpgradeBanner"; -import { uiHref } from "@/utils/uiHref"; +import { routeSegmentForPathname, uiHref } from "@/utils/uiHref"; import { PluginModeProvider, usePluginMode } from "@/contexts/PluginModeContext"; import { createApiClient } from "@/lib/http/client"; import { getProxyBaseUrl } from "@/components/networking"; @@ -103,6 +103,7 @@ function DashboardShell({ children }: { children: React.ReactNode }) { const { accessToken } = useAuth(); const [sidebarCollapsed, setSidebarCollapsed] = useState(false); const { mode } = usePluginMode(); + const isPlayground = routeSegmentForPathname(usePathname()) === "playground"; const isGateway = mode === "ai-gateway"; @@ -142,7 +143,7 @@ function DashboardShell({ children }: { children: React.ReactNode }) {
{children}
- + {!isPlayground && }
); diff --git a/ui/litellm-dashboard/src/components/Navbar/UserDropdown/UserDropdown.tsx b/ui/litellm-dashboard/src/components/Navbar/UserDropdown/UserDropdown.tsx index 9c02defc778..ab9a723e57a 100644 --- a/ui/litellm-dashboard/src/components/Navbar/UserDropdown/UserDropdown.tsx +++ b/ui/litellm-dashboard/src/components/Navbar/UserDropdown/UserDropdown.tsx @@ -1,6 +1,7 @@ import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import { useDisableBlogPosts } from "@/app/(dashboard)/hooks/useDisableBlogPosts"; import { useDisableBouncingIcon } from "@/app/(dashboard)/hooks/useDisableBouncingIcon"; +import { useDisableLiteAdmin } from "@/app/(dashboard)/hooks/useDisableLiteAdmin"; import { useDisableShowPrompts } from "@/app/(dashboard)/hooks/useDisableShowPrompts"; import { emitLocalStorageChange, @@ -10,6 +11,7 @@ import { } from "@/utils/localStorageUtils"; import { navAccountDisplayName } from "@/components/Navbar/navDisplayName"; import { uiHref } from "@/utils/uiHref"; +import { isProxyAdminRole } from "@/utils/roles"; import { ChevronDown, ChevronsUpDown, Crown, KeyRound, LogOut, Mail, ShieldCheck, User } from "lucide-react"; import { useRouter } from "next/navigation"; import { Avatar, AvatarFallback } from "@/components/ui/avatar"; @@ -65,12 +67,22 @@ interface UserDropdownProps { } const UserDropdown: React.FC = ({ onLogout, variant = "navbar", collapsed = false }) => { - const { userId, userEmail, userRoleLabel: userRole, premiumUser, loginMethod } = useAuthorized(); + const { + userId, + userEmail, + userRole: role, + userRoleLabel: userRole, + isViewOnly, + premiumUser, + loginMethod, + } = useAuthorized(); const router = useRouter(); const [open, setOpen] = useState(false); const disableShowPrompts = useDisableShowPrompts(); const disableBlogPosts = useDisableBlogPosts(); const disableBouncingIcon = useDisableBouncingIcon(); + const [disableLiteAdmin, setDisableLiteAdmin] = useDisableLiteAdmin(userId); + const canUseLiteAdmin = userId && !isViewOnly && isProxyAdminRole(role); const [disableShowNewBadge, setDisableShowNewBadge] = useState(false); useEffect(() => { @@ -192,6 +204,17 @@ const UserDropdown: React.FC = ({ onLogout, variant = "navbar aria-label="Toggle hide bouncing icon" /> + {canUseLiteAdmin && ( +
+ Hide LiteAdmin + +
+ )} ); diff --git a/ui/litellm-dashboard/src/components/SidebarAccountMenu/SidebarAccountMenu.tsx b/ui/litellm-dashboard/src/components/SidebarAccountMenu/SidebarAccountMenu.tsx index e697cc6f34a..e580b8610d0 100644 --- a/ui/litellm-dashboard/src/components/SidebarAccountMenu/SidebarAccountMenu.tsx +++ b/ui/litellm-dashboard/src/components/SidebarAccountMenu/SidebarAccountMenu.tsx @@ -2,6 +2,7 @@ import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import { useHealthReadinessDetails } from "@/app/(dashboard)/hooks/healthReadiness/useHealthReadinessDetails"; import { useDisableBlogPosts } from "@/app/(dashboard)/hooks/useDisableBlogPosts"; import { useDisableBouncingIcon } from "@/app/(dashboard)/hooks/useDisableBouncingIcon"; +import { useDisableLiteAdmin } from "@/app/(dashboard)/hooks/useDisableLiteAdmin"; import { useDisableShowNewBadge } from "@/app/(dashboard)/hooks/useDisableShowNewBadge"; import { useDisableShowPrompts } from "@/app/(dashboard)/hooks/useDisableShowPrompts"; import { emitLocalStorageChange, removeLocalStorageItem, setLocalStorageItem } from "@/utils/localStorageUtils"; @@ -15,6 +16,7 @@ import { Separator } from "@/components/ui/separator"; import { Switch } from "@/components/ui/switch"; import { cn } from "@/lib/cva.config"; import { uiHref } from "@/utils/uiHref"; +import { isProxyAdminRole } from "@/utils/roles"; import { ChevronsUpDown, Crown, IdCard, KeyRound, LogOut, Mail, ShieldCheck } from "lucide-react"; import { useRouter } from "next/navigation"; import React from "react"; @@ -83,7 +85,16 @@ interface SidebarAccountMenuProps { } const SidebarAccountMenu: React.FC = ({ onLogout, collapsed = false }) => { - const { userId, userEmail, userRoleLabel: userRole, premiumUser, accessToken, loginMethod } = useAuthorized(); + const { + userId, + userEmail, + userRole: role, + userRoleLabel: userRole, + isViewOnly, + premiumUser, + accessToken, + loginMethod, + } = useAuthorized(); const router = useRouter(); const [open, setOpen] = React.useState(false); const { data: healthData } = useHealthReadinessDetails(accessToken); @@ -92,6 +103,8 @@ const SidebarAccountMenu: React.FC = ({ onLogout, colla const disableBlogPosts = useDisableBlogPosts(); const disableBouncingIcon = useDisableBouncingIcon(); const disableShowNewBadge = useDisableShowNewBadge(); + const [disableLiteAdmin, setDisableLiteAdmin] = useDisableLiteAdmin(userId); + const canUseLiteAdmin = userId && !isViewOnly && isProxyAdminRole(role); const setFlag = (key: string, checked: boolean) => { if (checked) { @@ -235,6 +248,17 @@ const SidebarAccountMenu: React.FC = ({ onLogout, colla /> ))} + {canUseLiteAdmin && ( +
+ Hide LiteAdmin + +
+ )} diff --git a/ui/litellm-dashboard/src/components/liteadmin/LiteAdmin.integration.test.tsx b/ui/litellm-dashboard/src/components/liteadmin/LiteAdmin.integration.test.tsx index 7031702c53e..eaee2d520ac 100644 --- a/ui/litellm-dashboard/src/components/liteadmin/LiteAdmin.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/liteadmin/LiteAdmin.integration.test.tsx @@ -7,6 +7,9 @@ import { setGlobalLitellmHeaderName, switchToWorkerUrl } from "@/components/netw import { Toaster } from "@/components/ui/sonner"; import { toast } from "@/lib/toast"; import userEvent from "@testing-library/user-event"; +import type { ComponentType } from "react"; +import SidebarAccountMenu from "@/components/SidebarAccountMenu/SidebarAccountMenu"; +import UserDropdown from "@/components/Navbar/UserDropdown/UserDropdown"; import LiteAdmin from "./LiteAdmin"; import { MAX_INPUT_LENGTH } from "./agent"; @@ -18,6 +21,7 @@ const { transport } = vi.hoisted(() => { vi.unmock("@/app/(dashboard)/hooks/useAuthorized"); vi.unmock("@/lib/toast"); +vi.mock("next/navigation", () => ({ useRouter: () => ({ push: vi.fn() }) })); const MANAGEMENT = "https://management.test/proxy"; const INFERENCE = "https://management.test/inference"; @@ -80,13 +84,14 @@ function SessionReady() { return {authLoading ? "Session loading" : "Session ready"}; } -function renderWidget() { +function renderWidget(Menu?: ComponentType<{ onLogout: () => void }>) { const client = new QueryClient({ defaultOptions: { queries: { retry: false, gcTime: 0 } } }); const tree = () => ( + {Menu && undefined} />} @@ -111,6 +116,7 @@ function gateway(replies: (ModelReply | Promise)[], options: Gateway const path = new URL(request.url).pathname; if (path.endsWith("/litellm-ui-config")) return json({ proxy_base_url: MANAGEMENT, server_root_path: "", admin_ui_disabled: false }); + if (path.endsWith("/health/readiness/details")) return json({ status: "healthy" }); if (path.endsWith("/sso/get/ui_settings")) { if (typeof settings === "function") return settings(request); return json({ PROXY_BASE_URL: MANAGEMENT, LITELLM_UI_API_DOC_BASE_URL: settings.target }, settings.status); @@ -176,6 +182,79 @@ afterEach(() => { }); describe("LiteAdmin in the gateway", () => { + it.each([ + ["sidebar", SidebarAccountMenu], + ["navbar", UserDropdown], + ] as const)("persists Hide LiteAdmin from the %s account menu", async (_name, Menu) => { + gateway([]); + const user = userEvent.setup(); + const view = renderWidget(Menu); + await screen.findByRole("button", { name: "LiteAdmin" }); + await user.click(screen.getByRole("button", { name: /account menu/i })); + const toggle = await screen.findByRole("switch", { name: "Toggle hide LiteAdmin" }); + expect(toggle).not.toBeChecked(); + await user.click(toggle); + expect(toggle).toBeChecked(); + expect(screen.queryByRole("button", { name: "LiteAdmin" })).not.toBeInTheDocument(); + + view.unmount(); + const restored = renderWidget(Menu); + await screen.findByText("Session ready"); + await waitFor(() => expect(restored.client.isFetching()).toBe(0)); + expect(screen.queryByRole("button", { name: "LiteAdmin" })).not.toBeInTheDocument(); + await user.click(screen.getByRole("button", { name: /account menu/i })); + const savedToggle = await screen.findByRole("switch", { name: "Toggle hide LiteAdmin" }); + expect(savedToggle).toBeChecked(); + await user.click(savedToggle); + expect(await screen.findByRole("button", { name: "LiteAdmin" })).toBeInTheDocument(); + }); + + it("isolates Hide LiteAdmin by admin and gateway and reacts to another tab clearing it", async () => { + gateway([]); + const user = userEvent.setup(); + const view = renderWidget(SidebarAccountMenu); + await screen.findByRole("button", { name: "LiteAdmin" }); + await user.click(screen.getByRole("button", { name: /account menu/i })); + await user.click(await screen.findByRole("switch", { name: "Toggle hide LiteAdmin" })); + + session("proxy_admin", "second-admin"); + view.refresh(); + expect(await screen.findByRole("button", { name: "LiteAdmin" })).toBeInTheDocument(); + expect(screen.getByRole("switch", { name: "Toggle hide LiteAdmin" })).not.toBeChecked(); + session(); + view.refresh(); + expect(screen.queryByRole("button", { name: "LiteAdmin" })).not.toBeInTheDocument(); + + switchToWorkerUrl("https://other-gateway.test"); + view.refresh(); + expect(await screen.findByRole("button", { name: "LiteAdmin" })).toBeInTheDocument(); + expect(screen.getByRole("switch", { name: "Toggle hide LiteAdmin" })).not.toBeChecked(); + switchToWorkerUrl(MANAGEMENT); + view.refresh(); + expect(screen.queryByRole("button", { name: "LiteAdmin" })).not.toBeInTheDocument(); + + act(() => { + localStorage.clear(); + window.dispatchEvent(new StorageEvent("storage", { key: null })); + }); + expect(await screen.findByRole("button", { name: "LiteAdmin" })).toBeInTheDocument(); + expect(screen.getByRole("switch", { name: "Toggle hide LiteAdmin" })).not.toBeChecked(); + }); + + it.each([ + ["sidebar", SidebarAccountMenu], + ["navbar", UserDropdown], + ] as const)("does not offer Hide LiteAdmin to a view-only admin in the %s menu", async (_name, Menu) => { + session("proxy_admin_viewer"); + gateway([]); + renderWidget(Menu); + await screen.findByText("Session ready"); + const user = userEvent.setup(); + await user.click(screen.getByRole("button", { name: /account menu/i })); + expect(await screen.findByRole("switch", { name: "Toggle hide all prompts" })).toBeInTheDocument(); + expect(screen.queryByRole("switch", { name: "Toggle hide LiteAdmin" })).not.toBeInTheDocument(); + }); + it.each(["proxy_admin_viewer", "internal_user", "internal_user_viewer", "org_admin"])( "does not expose operations to %s", async (role) => { diff --git a/ui/litellm-dashboard/src/components/liteadmin/LiteAdmin.tsx b/ui/litellm-dashboard/src/components/liteadmin/LiteAdmin.tsx index 0395ed99fb3..8092eb2f5e6 100644 --- a/ui/litellm-dashboard/src/components/liteadmin/LiteAdmin.tsx +++ b/ui/litellm-dashboard/src/components/liteadmin/LiteAdmin.tsx @@ -4,6 +4,7 @@ import { useRef, useState, type ReactNode } from "react"; import { useQuery } from "@tanstack/react-query"; import { RotateCcw, Sparkles, X } from "lucide-react"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { useDisableLiteAdmin } from "@/app/(dashboard)/hooks/useDisableLiteAdmin"; import { useProxySettingsQuery } from "@/app/(dashboard)/hooks/proxySettings/useProxySettings"; import { ChatComposer } from "@/app/(dashboard)/playground/components/chat_ui/ChatComposer"; import { EndpointType, isModeCompatibleWithEndpoint } from "@/components/chat_ui/mode_endpoint_mapping"; @@ -34,9 +35,10 @@ type ManagementSession = Omit; export default function LiteAdmin() { const auth = useAuthorized(); + const [disabled] = useDisableLiteAdmin(auth.userId); const sessionReady = !auth.isLoading && auth.isAuthorized; const writableAdmin = !auth.isViewOnly && isProxyAdminRole(auth.userRole); - const allowed = sessionReady && writableAdmin; + const allowed = sessionReady && writableAdmin && !disabled; if (!allowed || !auth.token || !auth.accessToken) return null; const session = { token: auth.token, accessToken: auth.accessToken, managementBaseUrl: getProxyBaseUrl() }; return ( From 0c1c3e18d5250ec3a0e1e3f287e0b93e2906d900 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Wed, 23 Sep 2026 17:33:22 -0700 Subject: [PATCH 05/10] fix(ui): prefer native providers in auto-router presets (#42639) --- .../app/(dashboard)/hooks/models/useModels.ts | 1 + .../add_auto_router_tab.integration.test.tsx | 32 ++- .../components/add_model/auto_setup.test.ts | 38 +++- .../src/lib/autorouter_presets.test.ts | 209 ++++++++++++++++-- .../src/lib/autorouter_presets.ts | 87 ++++++-- 5 files changed, 322 insertions(+), 45 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts index 579ee7ff81a..b5cc329d4c9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts @@ -99,6 +99,7 @@ export interface AutoRouterDeployment extends AutoRouterCandidateDeployment { litellm_params?: { model?: string | null; base_model?: string | null; + custom_llm_provider?: string | null; complexity_router_config?: unknown; complexity_router_default_model?: string | null; auto_router_config?: unknown; diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.integration.test.tsx index 2779eec752f..0337d35d698 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.integration.test.tsx @@ -1660,10 +1660,6 @@ describe("AddAutoRouterTab", () => { const ALL_RENAMED_DEPLOYMENTS = getAllPresets().flatMap((preset) => renamedDeploymentsFor(preset.key)); - const renamedGroupFor = (model: string): string => - ALL_RENAMED_DEPLOYMENTS.find((deployment) => deployment.litellm_params.model === `someprovider/${model}`)! - .model_name; - it("enables a preset whose models exist only under renamed deployments, labeling the match", async () => { mockFetchAvailableModels.mockResolvedValue(groupsFor(ALL_RENAMED_DEPLOYMENTS)); mockFetchAllModelDeployments.mockResolvedValue(ALL_RENAMED_DEPLOYMENTS); @@ -1677,10 +1673,23 @@ describe("AddAutoRouterTab", () => { expect(optionByLabel("Anthropic Family")!).toHaveTextContent(/Matches your deployments/); }); - it("keeps detailed configuration open and prefills the admin's group names on apply", async () => { + it("keeps detailed configuration open and submits native group names when cloud twins are available", async () => { const user = userEvent.setup(); - mockFetchAvailableModels.mockResolvedValue(groupsFor(ALL_RENAMED_DEPLOYMENTS)); - mockFetchAllModelDeployments.mockResolvedValue(ALL_RENAMED_DEPLOYMENTS); + const nativeDeployments = renamedDeploymentsFor("anthropic_family").map((deployment) => ({ + ...deployment, + litellm_params: { model: deployment.litellm_params.model.replace("someprovider/", "anthropic/") }, + })); + const nativeGroupFor = (model: string): string => + nativeDeployments.find((deployment) => deployment.litellm_params.model === `anthropic/${model}`)!.model_name; + const cloudDeployments = nativeDeployments.map((deployment) => ({ + model_name: `a-cloud-${deployment.model_name}`, + litellm_params: { + model: `bedrock/us.anthropic.${deployment.litellm_params.model.split("/")[1]}-v1:0`, + }, + })); + const deployments = [...cloudDeployments, ...nativeDeployments]; + mockFetchAvailableModels.mockResolvedValue(groupsFor(deployments)); + mockFetchAllModelDeployments.mockResolvedValue(deployments); renderWithProviders(); openTemplateDropdown(); @@ -1688,6 +1697,7 @@ describe("AddAutoRouterTab", () => { expect(isOptionDisabled(optionByLabel("Anthropic Family")!)).toBe(false); }); await selectTemplate("Anthropic Family"); + expectTierModel("Complex", nativeGroupFor(ANTHROPIC_TIERS.COMPLEX[0])); openAutoRouterAdvanced("Keyword/Semantic Matching"); @@ -1701,10 +1711,10 @@ describe("AddAutoRouterTab", () => { expect(vi.mocked(handleAddAutoRouterSubmit).mock.calls.at(-1)?.[0]).toMatchObject({ complexity_router_config: { tiers: { - SIMPLE: ANTHROPIC_TIERS.SIMPLE.map(renamedGroupFor), - MEDIUM: ANTHROPIC_TIERS.MEDIUM.map(renamedGroupFor), - COMPLEX: ANTHROPIC_TIERS.COMPLEX.map(renamedGroupFor), - REASONING: ANTHROPIC_TIERS.REASONING.map(renamedGroupFor), + SIMPLE: ANTHROPIC_TIERS.SIMPLE.map(nativeGroupFor), + MEDIUM: ANTHROPIC_TIERS.MEDIUM.map(nativeGroupFor), + COMPLEX: ANTHROPIC_TIERS.COMPLEX.map(nativeGroupFor), + REASONING: ANTHROPIC_TIERS.REASONING.map(nativeGroupFor), }, }, }); diff --git a/ui/litellm-dashboard/src/components/add_model/auto_setup.test.ts b/ui/litellm-dashboard/src/components/add_model/auto_setup.test.ts index c5784db4501..a6e81bbc184 100644 --- a/ui/litellm-dashboard/src/components/add_model/auto_setup.test.ts +++ b/ui/litellm-dashboard/src/components/add_model/auto_setup.test.ts @@ -1,6 +1,6 @@ import { describe, expect, it } from "vitest"; import type { AutoRouterDeployment } from "@/app/(dashboard)/hooks/models/useModels"; -import { buildModelAvailability } from "@/lib/autorouter_presets"; +import { buildModelAvailability, deploymentRefsFromModelInfo } from "@/lib/autorouter_presets"; import { buildAutomaticRouterConfig, buildPreferredTierModels, type PreferredTierModels } from "./auto_setup"; const models = (...names: string[]) => names.map((model_group) => ({ model_group, mode: "chat" })); @@ -54,6 +54,42 @@ describe("buildPreferredTierModels", () => { }); describe("buildAutomaticRouterConfig", () => { + it("selects native Terra and Sol groups with their reasoning settings, retaining cloud fallback", () => { + const modelNames = ["gpt-5.6-terra", "gpt-5.6-sol"]; + const deployments = modelNames.flatMap((model) => [ + deployment(model, `azure/${model}`), + deployment(`z-native-${model}`, `openai/${model}`), + ]); + const available = deployments.map(({ model_name }) => reasoningModel(model_name!, ["none", "high"])); + const availability = buildModelAvailability( + available.map(({ model_group }) => model_group), + deploymentRefsFromModelInfo(deployments), + ); + const preferred = buildPreferredTierModels([], availability); + const config = buildAutomaticRouterConfig(available, deployments, preferred); + + expect(tierModels(config)).toEqual([ + "z-native-gpt-5.6-terra", + "z-native-gpt-5.6-terra", + "z-native-gpt-5.6-sol", + "z-native-gpt-5.6-sol", + ]); + expect(config?.tier_model_params).toEqual({ + REASONING: { "z-native-gpt-5.6-sol": { reasoning_effort: "high" } }, + }); + + const cloudOnly = models(...modelNames); + const cloudAvailability = buildModelAvailability(modelNames, deploymentRefsFromModelInfo(deployments)); + const cloudPreferred = buildPreferredTierModels([], cloudAvailability); + + expect(tierModels(buildAutomaticRouterConfig(cloudOnly, deployments, cloudPreferred))).toEqual([ + "gpt-5.6-terra", + "gpt-5.6-terra", + "gpt-5.6-sol", + "gpt-5.6-sol", + ]); + }); + it("selects one preferred model for each tier", () => { const preferred: PreferredTierModels = { SIMPLE: ["simple"], diff --git a/ui/litellm-dashboard/src/lib/autorouter_presets.test.ts b/ui/litellm-dashboard/src/lib/autorouter_presets.test.ts index 1d75a1dc7d2..c67befe8fa5 100644 --- a/ui/litellm-dashboard/src/lib/autorouter_presets.test.ts +++ b/ui/litellm-dashboard/src/lib/autorouter_presets.test.ts @@ -13,6 +13,7 @@ import { buildModelAvailability, deploymentRefsFromModelInfo, normalizeModelName, + resolveAvailableModel, resolveAvailableModels, } from "./autorouter_presets"; import { DEFAULT_MATCH_THRESHOLD } from "@/components/add_model/SemanticKeywordMatching"; @@ -393,28 +394,140 @@ describe("autorouter_presets", () => { expect(resolveAvailableModels("anthropic/claude-sonnet-5", availability)).toEqual(["a-group", "z-group"]); }); - it("breaks ties between groups serving the same model deterministically, alphabetically", () => { + it.each([ + ["OpenAI", getPresetByKey("openai_family")!.complexity_router_config.tiers.MEDIUM[0], "openai", "azure"], + [ + "Anthropic", + getPresetByKey("anthropic_family")!.complexity_router_config.tiers.COMPLEX[0], + "anthropic", + "bedrock", + ], + ["Gemini", getPresetByKey("gemini_family")!.complexity_router_config.tiers.SIMPLE[0], "gemini", "vertex_ai"], + ["DeepSeek", getPresetByKey("lite")!.complexity_router_config.tiers.SIMPLE[0], "deepseek", "openrouter"], + ["Muse", getPresetByKey("lite")!.complexity_router_config.tiers.MEDIUM[0], "meta", "openrouter"], + ["Kimi", getPresetByKey("lite")!.complexity_router_config.tiers.COMPLEX[0], "moonshot", "openrouter"], + ["Grok", "grok-4.7", "xai", "openrouter"], + ])( + "prefills %s through its native provider and falls back when only the cloud group is available", + (_family, model, native, cloud) => { + const deployments = [ + { modelGroup: "a-cloud", underlyingModels: [`${cloud}/${model}`] }, + { modelGroup: "z-native", underlyingModels: [`${native}/${model}`] }, + ]; + const config = { + tiers: { SIMPLE: [model], MEDIUM: [], COMPLEX: [], REASONING: [] }, + tier_model_configs: { SIMPLE: [{ model_name: model, litellm_params: { reasoning_effort: "high" } }] }, + classifier_type: "llm" as const, + classifier_llm_config: { model, timeout_ms: 3000 }, + classification_mode: "every_request" as const, + session_affinity: false, + deployment_affinity: true, + modality_routing: false, + modality_pin_override: false, + }; + + for (const [groups, selected] of [ + [["a-cloud", "z-native"], "z-native"], + [["a-cloud"], "a-cloud"], + ] as const) { + const availability = buildModelAvailability(groups, deployments); + const prefill = buildPresetPrefill(config, availability).complexityRouterConfig; + + expect(prefill.tiers.SIMPLE).toEqual([selected]); + expect(prefill.tier_model_params).toEqual({ SIMPLE: { [selected]: { reasoning_effort: "high" } } }); + expect(prefill.classifier_llm_config).toEqual({ model: selected, timeout_ms: 3000 }); + } + }, + ); + + it.each(["claude-opus-5-5", "claude-opus-5.5"])( + "prefers a native deployment over the cloud group named %s", + (cloudGroup) => { + const availability = buildModelAvailability( + [cloudGroup, "z-native"], + [ + { modelGroup: cloudGroup, underlyingModels: ["bedrock/us.anthropic.claude-opus-5-5-v1:0"] }, + { modelGroup: "z-native", underlyingModels: ["anthropic/claude-opus-5-5"] }, + ], + ); + + expect(resolveAvailableModel("claude-opus-5-5", availability)).toBe("z-native"); + expect(resolveAvailableModels("claude-opus-5-5", availability)).toEqual([cloudGroup]); + }, + ); + + it("breaks ties between native groups alphabetically regardless of deployment order", () => { const availability = buildModelAvailability( - ["z-group", "a-group"], + ["z-native", "a-native"], [ - { modelGroup: "z-group", underlyingModels: ["anthropic/claude-opus-5"] }, - { modelGroup: "a-group", underlyingModels: ["bedrock/us.anthropic.claude-opus-5-v1:0"] }, + { modelGroup: "z-native", underlyingModels: ["anthropic/claude-opus-5-5"] }, + { modelGroup: "a-native", underlyingModels: ["anthropic/claude-opus-5-5"] }, ], ); - const config = { - tiers: { SIMPLE: ["claude-opus-5"], MEDIUM: [], COMPLEX: [], REASONING: [] }, - classifier_type: "heuristic" as const, - classification_mode: "every_request" as const, - session_affinity: false, - deployment_affinity: true, - }; - expect(buildPresetPrefill(config, availability).complexityRouterConfig.tiers.SIMPLE).toEqual(["a-group"]); + + expect(resolveAvailableModel("claude-opus-5-5", availability)).toBe("a-native"); }); - it("prefers an exact group-name match over the deployment index", () => { + it.each(["gpt-6-sol", "claude-opus-5-5"])("recognizes the native default of bare %s", (model) => { + const availability = buildModelAvailability( + ["a-cloud", "z-native"], + [ + { modelGroup: "a-cloud", underlyingModels: [`openrouter/${model}`] }, + { modelGroup: "z-native", underlyingModels: [model] }, + ], + ); + + expect(resolveAvailableModel(model, availability)).toBe("z-native"); + }); + + it.each(["bedrock/claude-opus-5-5", "unknown-model"])( + "prefers an exclusively native group over one that also routes to %s", + (otherModel) => { + const deployments = [ + { modelGroup: "a-cloud", underlyingModels: ["bedrock/claude-opus-5-5"] }, + { modelGroup: "b-mixed", underlyingModels: ["anthropic/claude-opus-5-5"] }, + { modelGroup: "b-mixed", underlyingModels: [otherModel] }, + { modelGroup: "z-native", underlyingModels: ["anthropic/claude-opus-5-5"] }, + ]; + const availability = buildModelAvailability(["a-cloud", "b-mixed", "z-native"], deployments); + + expect(resolveAvailableModel("claude-opus-5-5", availability)).toBe("z-native"); + const noNativeGroup = buildModelAvailability(["a-cloud", "b-mixed"], deployments); + expect(resolveAvailableModel("claude-opus-5-5", noNativeGroup)).toBe("a-cloud"); + }, + ); + + it.each([ + { model: "azure/opaque-deployment", base_model: "openai/gpt-6-sol" }, + { model: "openai/gpt-6-sol", custom_llm_provider: "openrouter" }, + ])("keeps cloud routing authoritative over native-looking model metadata: %j", (litellmParams) => { + const availability = buildModelAvailability( + ["a-cloud", "z-native"], + deploymentRefsFromModelInfo([ + { model_name: "a-cloud", litellm_params: litellmParams, model_info: { base_model: "openai/gpt-6-sol" } }, + { model_name: "z-native", litellm_params: { model: "openai/gpt-6-sol" } }, + ]), + ); + + expect(resolveAvailableModel("gpt-6-sol", availability)).toBe("z-native"); + }); + + it("recognizes an explicit native provider on an otherwise unqualified model", () => { + const availability = buildModelAvailability( + ["a-cloud", "z-native"], + deploymentRefsFromModelInfo([ + { model_name: "a-cloud", litellm_params: { model: "openrouter/meta/muse-spark-1.3" } }, + { model_name: "z-native", litellm_params: { model: "muse-spark-1.3", custom_llm_provider: "meta" } }, + ]), + ); + + expect(resolveAvailableModel("muse-spark-1.3", availability)).toBe("z-native"); + }); + + it("preserves exact group-name precedence when no known native deployment is available", () => { const availability = buildModelAvailability( ["claude-opus-5", "renamed-opus"], - [{ modelGroup: "renamed-opus", underlyingModels: ["anthropic/claude-opus-5"] }], + [{ modelGroup: "renamed-opus", underlyingModels: ["bedrock/us.anthropic.claude-opus-5-v1:0"] }], ); const config = { tiers: { SIMPLE: ["claude-opus-5"], MEDIUM: [], COMPLEX: [], REASONING: [] }, @@ -573,6 +686,70 @@ describe("autorouter_presets", () => { ]); }); + it.each(["native/*", "*"])("ranks wildcard groups using their routing deployment: %s", (nativePattern) => { + const nativeGroup = nativePattern === "*" ? "openai/gpt-6-sol" : "native/gpt-6-sol"; + const availability = buildModelAvailability( + ["azure/gpt-6-sol", nativeGroup], + [ + { modelGroup: "azure/*", underlyingModels: ["openrouter/*"] }, + { modelGroup: nativePattern, underlyingModels: ["openai/*"] }, + ], + ); + + expect(resolveAvailableModel("gpt-6-sol", availability)).toBe(nativeGroup); + }); + + it("does not treat a native-looking wildcard group as native when its deployment uses the cloud", () => { + const availability = buildModelAvailability( + ["openai/gpt-6-sol", "z-native/gpt-6-sol"], + [ + { modelGroup: "openai/*", underlyingModels: ["azure/*"] }, + { modelGroup: "z-native/*", underlyingModels: ["openai/*"] }, + ], + ); + + expect(resolveAvailableModel("gpt-6-sol", availability)).toBe("z-native/gpt-6-sol"); + }); + + it("keeps literal native deployments ahead of a matching cloud wildcard", () => { + const availability = buildModelAvailability( + ["a-cloud", "team/gpt-6-sol"], + [ + { modelGroup: "a-cloud", underlyingModels: ["azure/gpt-6-sol"] }, + { modelGroup: "team/gpt-6-sol", underlyingModels: ["openai/gpt-6-sol"] }, + { modelGroup: "team/*", underlyingModels: ["azure/*"] }, + ], + ); + + expect(resolveAvailableModel("gpt-6-sol", availability)).toBe("team/gpt-6-sol"); + }); + + it("does not promote a bare-star expansion when its routing group also contains a cloud deployment", () => { + const availability = buildModelAvailability( + ["openai/gpt-6-sol", "z-native"], + [ + { modelGroup: "*", underlyingModels: ["openai/*"] }, + { modelGroup: "*", underlyingModels: ["azure/*"] }, + { modelGroup: "z-native", underlyingModels: ["openai/gpt-6-sol"] }, + ], + ); + + expect(resolveAvailableModel("gpt-6-sol", availability)).toBe("z-native"); + }); + + it("retains fallback ordering when overlapping wildcard routes have different providers", () => { + const availability = buildModelAvailability( + ["a-cloud", "team/gpt-6-sol"], + [ + { modelGroup: "a-cloud", underlyingModels: ["azure/gpt-6-sol"] }, + { modelGroup: "team/*", underlyingModels: ["azure/*"] }, + { modelGroup: "team/gpt-*", underlyingModels: ["openai/gpt-*"] }, + ], + ); + + expect(resolveAvailableModel("gpt-6-sol", availability)).toBe("a-cloud"); + }); + it.each(getAllPresets().map((preset) => [preset.key, preset] as const))( "fully resolves the %s preset through wildcard-expanded groups only", (_key, preset) => { @@ -602,7 +779,9 @@ describe("autorouter_presets", () => { { model_name: "no-underlying", litellm_params: {}, model_info: {} }, { litellm_params: { model: "openai/gpt-5.4" } }, ]); - expect(refs).toEqual([{ modelGroup: "azure-prod", underlyingModels: ["azure/my-deployment", "azure/gpt-5.4"] }]); + expect(refs).toEqual([ + { modelGroup: "azure-prod", underlyingModels: ["azure/my-deployment", "azure/gpt-5.4"], provider: "azure" }, + ]); }); it("lets an azure deployment resolve through base_model declared under litellm_params", () => { diff --git a/ui/litellm-dashboard/src/lib/autorouter_presets.ts b/ui/litellm-dashboard/src/lib/autorouter_presets.ts index 8cf461d77b9..bbbf151e49b 100644 --- a/ui/litellm-dashboard/src/lib/autorouter_presets.ts +++ b/ui/litellm-dashboard/src/lib/autorouter_presets.ts @@ -65,13 +65,37 @@ export const normalizeModelName = (model: string): string => model.replace(/(\d) export interface DeploymentModelRef { modelGroup: string; underlyingModels: readonly string[]; + provider?: string; } export interface ModelAvailability { modelGroups: Set; underlyingIndex: Map; + nativeUnderlyingIndex: Map; } +const NATIVE_MODEL_PROVIDERS: readonly (readonly [RegExp, string])[] = [ + [/^(gpt-|o\d|text-embedding-)/, "openai"], + [/^claude-/, "anthropic"], + [/^gemini-/, "gemini"], + [/^deepseek-/, "deepseek"], + [/^muse-/, "meta"], + [/^kimi-/, "moonshot"], + [/^grok-/, "xai"], +]; + +const nativeModelProvider = (model: string): string | undefined => + NATIVE_MODEL_PROVIDERS.find(([pattern]) => pattern.test(model))?.[1]; + +const routingProvider = (model: string): string => { + if (model.includes("/")) return model.split("/")[0]; + const native = nativeModelProvider(model); + return native === "openai" || native === "anthropic" ? native : ""; +}; + +const deploymentProvider = (deployment: DeploymentModelRef): string => + deployment.provider ?? routingProvider(deployment.underlyingModels[0] ?? ""); + const normalizeUnderlyingModel = (model: string): string | null => { if (model.includes("*")) return null; const ownName = model.slice(model.lastIndexOf("/") + 1).split("@")[0]; @@ -107,32 +131,42 @@ export const buildModelAvailability = ( deployments: readonly DeploymentModelRef[], ): ModelAvailability => { const groups = new Set(modelGroups); + const deploymentGroups = new Set(deployments.map((deployment) => deployment.modelGroup)); + const deploymentProviders = new Map>(); + for (const deployment of deployments) { + const providers = deploymentProviders.get(deployment.modelGroup) ?? new Set(); + providers.add(deploymentProvider(deployment)); + deploymentProviders.set(deployment.modelGroup, providers); + } const literalEntries = deployments .filter((deployment) => groups.has(deployment.modelGroup)) .flatMap((deployment) => deployment.underlyingModels .map(normalizeUnderlyingModel) - .filter((key): key is string => key !== null) - .map((key) => ({ key, modelGroup: deployment.modelGroup })), + .map((key) => ({ key, modelGroup: deployment.modelGroup, sourceGroup: deployment.modelGroup })), ); // Mirrors get_known_models_from_wildcard: a bare "*" model_name expands via its underlying // wildcard (or not at all), and a wildcard without a "/" expands to nothing. - const wildcardPatterns = Array.from( - new Set( - deployments - .flatMap((deployment) => - deployment.modelGroup === "*" ? deployment.underlyingModels : [deployment.modelGroup], - ) - .filter((pattern) => pattern !== "*" && pattern.includes("*") && pattern.includes("/")), - ), + const wildcardPatterns = deployments.flatMap((deployment) => + (deployment.modelGroup === "*" ? deployment.underlyingModels : [deployment.modelGroup]) + .filter((pattern) => pattern !== "*" && pattern.includes("*") && pattern.includes("/")) + .map((pattern) => ({ pattern, sourceGroup: deployment.modelGroup })), ); const wildcardEntries = Array.from(groups) - .filter((group) => !group.includes("*") && wildcardPatterns.some((pattern) => matchesWildcard(pattern, group))) - .map((group) => ({ key: normalizeUnderlyingModel(group), modelGroup: group })) - .filter((entry): entry is { key: string; modelGroup: string } => entry.key !== null); + .filter((group) => !group.includes("*") && !deploymentGroups.has(group)) + .flatMap((group) => + wildcardPatterns + .filter(({ pattern }) => matchesWildcard(pattern, group)) + .map(({ sourceGroup }) => ({ key: normalizeUnderlyingModel(group), modelGroup: group, sourceGroup })), + ); const entries = [...literalEntries, ...wildcardEntries]; const grouped = new Map>(); + const providersByGroup = new Map>(); for (const entry of entries) { + const providers = providersByGroup.get(entry.modelGroup) ?? new Set(); + for (const provider of deploymentProviders.get(entry.sourceGroup) ?? []) providers.add(provider); + providersByGroup.set(entry.modelGroup, providers); + if (entry.key === null) continue; const groupsForKey = grouped.get(entry.key) ?? new Set(); groupsForKey.add(entry.modelGroup); grouped.set(entry.key, groupsForKey); @@ -140,13 +174,23 @@ export const buildModelAvailability = ( const underlyingIndex = new Map( Array.from(grouped, ([key, groupsForKey]) => [key, Array.from(groupsForKey).sort()] as const), ); - return { modelGroups: groups, underlyingIndex }; + const nativeUnderlyingIndex = new Map( + Array.from(underlyingIndex, ([key, matches]) => [ + key, + matches.filter((group) => { + const native = nativeModelProvider(key); + const providers = providersByGroup.get(group); + return native !== undefined && providers?.size === 1 && providers.has(native); + }), + ]), + ); + return { modelGroups: groups, underlyingIndex, nativeUnderlyingIndex }; }; export const deploymentRefsFromModelInfo = ( rows: readonly { model_name?: string | null; - litellm_params?: { model?: string | null; base_model?: string | null } | null; + litellm_params?: { model?: string | null; base_model?: string | null; custom_llm_provider?: string | null } | null; model_info?: { base_model?: string | null } | null; }[], ): DeploymentModelRef[] => @@ -156,7 +200,10 @@ export const deploymentRefsFromModelInfo = ( row.litellm_params?.base_model, row.model_info?.base_model, ].filter((model): model is string => Boolean(model)); - return row.model_name && underlyingModels.length > 0 ? [{ modelGroup: row.model_name, underlyingModels }] : []; + const provider = row.litellm_params?.custom_llm_provider || routingProvider(row.litellm_params?.model ?? ""); + return row.model_name && underlyingModels.length > 0 + ? [{ modelGroup: row.model_name, underlyingModels, provider }] + : []; }); export const resolveAvailableModels = (requiredModel: string, availability: ModelAvailability): readonly string[] => { @@ -169,8 +216,12 @@ export const resolveAvailableModels = (requiredModel: string, availability: Mode return key === null ? [] : underlyingIndex.get(key) ?? []; }; -export const resolveAvailableModel = (requiredModel: string, availability: ModelAvailability): string | undefined => - resolveAvailableModels(requiredModel, availability)[0]; +export const resolveAvailableModel = (requiredModel: string, availability: ModelAvailability): string | undefined => { + const key = normalizeUnderlyingModel(requiredModel); + const nativeMatches = key === null ? [] : availability.nativeUnderlyingIndex.get(key) ?? []; + const matches = resolveAvailableModels(requiredModel, availability); + return matches.find((model) => nativeMatches.includes(model)) ?? nativeMatches[0] ?? matches[0]; +}; export const getMissingModels = ( config: Parameters[0], From 8b001d40246c2d76b510948b2e6b90be52d6065c Mon Sep 17 00:00:00 2001 From: tin-berri Date: Wed, 23 Sep 2026 17:43:54 -0700 Subject: [PATCH 06/10] fix(shadow-eval): replay approved pre-call guardrail snapshots (#42774) --- litellm/integrations/shadow_eval_logger.py | 104 +++++++++-- litellm/litellm_core_utils/litellm_logging.py | 20 +- litellm/proxy/common_request_processing.py | 2 +- litellm/proxy/litellm_pre_call_utils.py | 22 ++- .../integrations/test_shadow_eval_logger.py | 172 +++++++++++++++--- .../test_litellm_logging.py | 96 ++++++++++ .../proxy/test_common_request_processing.py | 102 +++++++---- .../proxy/test_litellm_pre_call_utils.py | 77 +++++++- 8 files changed, 512 insertions(+), 83 deletions(-) diff --git a/litellm/integrations/shadow_eval_logger.py b/litellm/integrations/shadow_eval_logger.py index cdc108a6b4e..84106b77a7d 100644 --- a/litellm/integrations/shadow_eval_logger.py +++ b/litellm/integrations/shadow_eval_logger.py @@ -11,6 +11,7 @@ across pods or stop races; the hook reads active jobs through a short-TTL cache. import asyncio import hashlib +import json import random import traceback from collections.abc import Awaitable, Callable, Mapping, Sequence @@ -28,7 +29,7 @@ from litellm.caching.in_memory_cache import InMemoryCache from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.websearch_interception.tools import is_web_search_tool_responses -from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs +from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs, independent_snapshot from litellm.litellm_core_utils.internal_call_metadata import sanitized_forwardable_call_metadata from litellm.litellm_core_utils.llm_judge import ( default_router_provider, @@ -281,10 +282,8 @@ class _SurfaceOps: request (messages plus translated generation params) and how its response yields the judgeable final text. Membership in this table IS the sampling allowlist; unknown call types fail closed. ``wire_params`` marks the surfaces whose params - come from the proxy's wire-body snapshot, which is taken before the guardrail - pre-call hook: those rows must not sample a request a pre-call guardrail rewrote, - or the shadow call would replay content (tools, unmasked entities) the guardrail - removed.""" + come from the proxy's native request snapshot. Requests rewritten by guardrails + require a post-hook snapshot whose guardrail history is still current.""" __slots__ = ("chat_request", "final_text", "wire_params") @@ -311,19 +310,85 @@ _NON_MUTATING_GUARDRAIL_MODES: Final = frozenset( ) +def _guardrail_is_non_mutating(entry: Mapping[str, object], allowed_modes: frozenset[str]) -> bool: + modes: Final = entry.get("guardrail_mode") + return all( + isinstance(mode, str) and mode in allowed_modes + for mode in (modes if isinstance(modes, list | tuple) else (modes,)) + ) + + def _request_mutating_guardrail_ran(request_metadata: Mapping[str, object]) -> bool: - """Whether a guardrail that can rewrite the outbound request ran on this one, read - from the same guardrail-information entries spend logging uses. str-enum modes - compare equal to their plain-string values, and an entry whose mode is missing or - unrecognized counts as mutating.""" raw: Final = request_metadata.get("standard_logging_guardrail_information") entries: Final = raw if isinstance(raw, Sequence) else () - modes_per_entry: Final = tuple(entry.get("guardrail_mode") for entry in entries if isinstance(entry, Mapping)) return any( - not all( - mode in _NON_MUTATING_GUARDRAIL_MODES for mode in (modes if isinstance(modes, list | tuple) else (modes,)) + not _guardrail_is_non_mutating(entry, _NON_MUTATING_GUARDRAIL_MODES) + for entry in entries + if isinstance(entry, Mapping) + ) + + +def request_guardrail_fingerprint(request_metadata: Mapping[str, object]) -> str | None: + raw: Final = request_metadata.get("standard_logging_guardrail_information") + entries: Final = raw if isinstance(raw, Sequence) else () + replay_safe_modes: Final = _NON_MUTATING_GUARDRAIL_MODES - frozenset(("logging_only",)) + relevant: Final = tuple( + entry + for entry in entries + if isinstance(entry, Mapping) and not _guardrail_is_non_mutating(entry, replay_safe_modes) + ) + try: + serialized: Final = json.dumps(relevant, sort_keys=True, default=str) + except (TypeError, ValueError): + return None + return hashlib.sha256(serialized.encode()).hexdigest() + + +@dataclass(frozen=True, slots=True) +class GuardrailRequestSnapshot: + body: Mapping[str, object] + fingerprint: str + + @staticmethod + def capture(body: Mapping[str, object], metadata: Mapping[str, object]) -> "GuardrailRequestSnapshot | None": + if not _request_mutating_guardrail_ran(metadata): + return None + fingerprint: Final = request_guardrail_fingerprint(metadata) + if fingerprint is None: + return None + return GuardrailRequestSnapshot( + body=MappingProxyType( + _CHAT_REQUEST_ADAPTER.validate_python( + independent_snapshot(dict(body)) # mutable-ok: snapshot helper requires a plain dictionary + ) + ), + fingerprint=fingerprint, ) - for modes in modes_per_entry + + +def _post_guardrail_kwargs( + kwargs: Mapping[str, object], + request_metadata: Mapping[str, object], + ops: _SurfaceOps, + guardrail_snapshot: GuardrailRequestSnapshot | None, +) -> Mapping[str, object] | None: + if guardrail_snapshot is None or guardrail_snapshot.fingerprint != request_guardrail_fingerprint(request_metadata): + return None + raw_params: Final = kwargs.get("litellm_params") + litellm_params: Final = raw_params if isinstance(raw_params, Mapping) else _EMPTY_METADATA + raw_request: Final = litellm_params.get("proxy_server_request") + request: Final = raw_request if isinstance(raw_request, Mapping) else _EMPTY_METADATA + body: Final = guardrail_snapshot.body + return MappingProxyType( + { + **kwargs, + "messages": body.get("input" if ops is _RESPONSES_OPS else "messages"), + "system": body.get("system"), + "instructions": body.get("instructions"), + "litellm_params": MappingProxyType( + {**litellm_params, "proxy_server_request": MappingProxyType({**request, "body": body})} + ), + } ) @@ -881,6 +946,8 @@ class ShadowEvalLogger(CustomLogger): response_obj: object, start_time: object, end_time: object, + *, + guardrail_snapshot: GuardrailRequestSnapshot | None = None, ) -> None: try: payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object") # pyright: ignore[reportAssignmentType] # untyped callback kwargs @@ -914,8 +981,13 @@ class ShadowEvalLogger(CustomLogger): ops: Final = _SURFACE_OPS.get(str(payload.get("call_type") or "")) if ops is None: return # only surfaces this table can normalize are comparable; unknown types fail closed - if ops.wire_params and _request_mutating_guardrail_ran(request_metadata): - return # the wire-body snapshot predates the rewrite; replaying it would resurrect stripped content + sample_kwargs: Final = ( + _post_guardrail_kwargs(kwargs, request_metadata, ops, guardrail_snapshot) + if ops.wire_params and _request_mutating_guardrail_ran(request_metadata) + else kwargs + ) + if sample_kwargs is None: + return active_jobs: Final = await self._active_jobs() eligible: Final = self._sampled_jobs( tuple(job for target in targets for job in active_jobs.get(target, ())), @@ -927,7 +999,7 @@ class ShadowEvalLogger(CustomLogger): return sample: Final = _judgeable_sample( ops, - kwargs, + sample_kwargs, MappingProxyType(dict(payload.get("model_parameters") or {})), # mutable-ok: frozen snapshot response_obj, ) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 8e28a0d543d..b038762a476 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -226,6 +226,7 @@ if TYPE_CHECKING: from litellm.integrations.otel.logger import OpenTelemetryV2 from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config + from litellm.integrations.shadow_eval_logger import GuardrailRequestSnapshot from litellm.litellm_core_utils.llm_cost_calc.utils import BilledTokenRates from litellm.llms.base_llm.passthrough.transformation import PassthroughStreamCollector from litellm.proxy.hooks.autorouter_baseline_cache import BaselineCacheContext, CapturedBaselineObservation @@ -714,6 +715,7 @@ class Logging(LiteLLMLoggingBaseClass): self._defer_async_logging: bool = False self._enqueue_deferred_logging: Callable[[], None] | None = None self._on_detached_stream_failure: Callable[[Exception], Awaitable[None]] | None = None + self.shadow_eval_request_snapshot: GuardrailRequestSnapshot | None = None def set_response_timing_metrics(self, timing_metrics: Mapping[str, float]) -> None: """Keep ``_response_ms`` / ``litellm_overhead_time_ms`` for a result that has no ``_hidden_params``.""" @@ -2825,6 +2827,7 @@ class Logging(LiteLLMLoggingBaseClass): ): continue + self.shadow_eval_request_snapshot = None self.model_call_details, result = callback.logging_hook( kwargs=self.model_call_details, result=result, @@ -3391,6 +3394,7 @@ class Logging(LiteLLMLoggingBaseClass): ): continue + self.shadow_eval_request_snapshot = None self.model_call_details, result = await callback.async_logging_hook( kwargs=self.model_call_details, result=result, @@ -3450,6 +3454,8 @@ class Logging(LiteLLMLoggingBaseClass): ) if isinstance(callback, CustomLogger): # custom logger class + from litellm.integrations.shadow_eval_logger import ShadowEvalLogger + model_call_details: dict = self.model_call_details ################################## # call redaction hook for custom logger @@ -3460,7 +3466,19 @@ class Logging(LiteLLMLoggingBaseClass): model_call_details=model_call_details, custom_logger=callback ) ################################## - if self.stream is True: + if isinstance(callback, ShadowEvalLogger) and ( + not self.stream or "async_complete_streaming_response" in model_call_details + ): + await callback.async_log_success_event( + kwargs=model_call_details, + response_obj=model_call_details["async_complete_streaming_response"] + if self.stream + else result, + start_time=start_time, + end_time=end_time, + guardrail_snapshot=self.shadow_eval_request_snapshot, + ) + elif self.stream is True: if "async_complete_streaming_response" in model_call_details: await callback.async_log_success_event( kwargs=model_call_details, diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index cbf4786affe..2e43c6b0d22 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -2208,7 +2208,7 @@ class ProxyBaseLLMRequestProcessing: # Refresh AFTER pre_call_hook: guardrails (e.g. Presidio PII masking) may # have mutated `self.data` in place, and the audit-trail snapshot taken in # add_litellm_data_to_request predates that mutation. - refresh_proxy_server_request_body_snapshot(self.data) + refresh_proxy_server_request_body_snapshot(self.data, guardrails_applied=True) verbose_proxy_logger.debug("receiving data: %s", self.data) if "messages" in self.data and self.data["messages"]: diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 90bba82aa84..3bf6fea62f7 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1923,6 +1923,8 @@ class LiteLLMProxyRequestSetup: def refresh_proxy_server_request_body_snapshot( data: MutableMapping[str, object], + *, + guardrails_applied: bool = False, ) -> None: """ Re-snapshot ``data["proxy_server_request"]["body"]`` from the current state of ``data``. @@ -1938,13 +1940,27 @@ def refresh_proxy_server_request_body_snapshot( ``Logging`` instance, so it must be excluded here the same way ``secret_fields`` and ``proxy_server_request`` are. """ - proxy_server_request = data.get("proxy_server_request") + from litellm.integrations.shadow_eval_logger import GuardrailRequestSnapshot + from litellm.litellm_core_utils.litellm_logging import Logging + + logging_obj: Final = data.get("litellm_logging_obj") + if isinstance(logging_obj, Logging): + logging_obj.shadow_eval_request_snapshot = None + proxy_server_request: Final = data.get("proxy_server_request") if not isinstance(proxy_server_request, dict): return - _body_snapshot_exclude = ( + _body_snapshot_exclude: Final = ( frozenset({"secret_fields", "proxy_server_request", "litellm_logging_obj"}) | _TRANSPORT_ONLY_CREDENTIAL_KEYS ) - proxy_server_request["body"] = {k: v for k, v in data.items() if k not in _body_snapshot_exclude} + body: Final = { # mutable-ok: audit JSON serialization requires a dict with shared nested messages + k: v for k, v in data.items() if k not in _body_snapshot_exclude + } + proxy_server_request["body"] = body + if guardrails_applied and isinstance(logging_obj, Logging): + metadata: Final = data.get(get_metadata_variable_name_from_kwargs(data)) + logging_obj.shadow_eval_request_snapshot = GuardrailRequestSnapshot.capture( + body, metadata if isinstance(metadata, Mapping) else MappingProxyType({}) + ) async def add_litellm_data_to_request( diff --git a/tests/test_litellm/integrations/test_shadow_eval_logger.py b/tests/test_litellm/integrations/test_shadow_eval_logger.py index 76efe9c8576..e3f059a7941 100644 --- a/tests/test_litellm/integrations/test_shadow_eval_logger.py +++ b/tests/test_litellm/integrations/test_shadow_eval_logger.py @@ -4,7 +4,7 @@ the detached pipeline's single attempt-row write, and the cache-first job lookup import asyncio from collections.abc import Mapping from datetime import datetime, timedelta, timezone -from typing import Final +from typing import Final, Literal from unittest.mock import AsyncMock, MagicMock import pytest @@ -19,11 +19,13 @@ from litellm.integrations.shadow_eval_logger import ( JUDGE_MAX_OUTPUT_TOKENS, PAIRWISE_JUDGE_RESPONSE_FORMAT, ActiveShadowEvalJob, + GuardrailRequestSnapshot, ShadowEvalLogger, _failure_detail, _judge_user_prompt, _sample_hits, _unmask_preference, + request_guardrail_fingerprint, ) from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import ( @@ -35,6 +37,15 @@ from litellm.types.utils import ( ) +def test_guardrail_fingerprint_excludes_auth_metadata() -> None: + history: Final = [{"guardrail_name": "mask", "guardrail_mode": "pre_call"}] + metadata: Final = {"standard_logging_guardrail_information": history} + fingerprint: Final = request_guardrail_fingerprint(metadata) + assert fingerprint == request_guardrail_fingerprint({**metadata, "user_api_key": "first-test-credential"}) + assert fingerprint == request_guardrail_fingerprint({**metadata, "user_api_key": "second-test-credential"}) + assert fingerprint != request_guardrail_fingerprint({"standard_logging_guardrail_information": []}) + + def _job(**overrides) -> ActiveShadowEvalJob: defaults = dict( id="job-1", @@ -312,11 +323,19 @@ class TestSurfaceNormalization: """/v1/messages and /v1/responses arms: the hook normalizes each surface's logged request through litellm's own transformations and judges only text-final turns.""" - async def _drive(self, hook_kwargs, response_obj): - prisma = _prisma() - router = _router() - logger = _logger(router=router, prisma=prisma, jobs=(_job(),)) - await logger.async_log_success_event(hook_kwargs, response_obj, None, None) + async def _drive( + self, + hook_kwargs: Mapping[str, object], + response_obj: object, + *, + guardrail_snapshot: GuardrailRequestSnapshot | None = None, + ) -> tuple[MagicMock, MagicMock]: + prisma: Final = _prisma() + router: Final = _router() + logger: Final = _logger(router=router, prisma=prisma, jobs=(_job(),)) + await logger.async_log_success_event( + hook_kwargs, response_obj, None, None, guardrail_snapshot=guardrail_snapshot + ) await _drain(logger) return prisma, router @@ -711,41 +730,142 @@ class TestSurfaceNormalization: prisma.db.litellm_shadowevalattempt.create.assert_not_called() @pytest.mark.parametrize( - "call_type,guardrail_mode,sampled", + "call_type,guardrail_mode,checkpoint,later_mode,sampled", [ - ("anthropic_messages", ["logging_only", "pre_call"], False), - ("aresponses", GuardrailEventHooks.pre_call, False), - ("anthropic_messages", "post_call", True), - ("acompletion", "pre_call", True), + ("anthropic_messages", "pre_call", "absent", None, False), + ("aresponses", "pre_call", "corrupt", None, False), + ("anthropic_messages", ["logging_only", "pre_call"], "missing", None, False), + ("aresponses", GuardrailEventHooks.pre_call, "missing", None, False), + ("anthropic_messages", "pre_call", "unapproved", None, False), + ("aresponses", "pre_call", "unapproved", None, False), + ("anthropic_messages", ["logging_only", "pre_call"], "approved", None, True), + ("aresponses", GuardrailEventHooks.pre_call, "approved", None, True), + ("anthropic_messages", "pre_call", "approved", "pre_call", False), + ("aresponses", "pre_call", "approved", "pre_call", False), + ("anthropic_messages", "pre_call", "approved", "logging_only", False), + ("aresponses", "pre_call", "approved", "logging_only", False), + ("anthropic_messages", "pre_call", "approved", "post_call", True), + ("aresponses", "pre_call", "approved", "post_call", True), + ("anthropic_messages", "post_call", "missing", None, True), + ("acompletion", "pre_call", "missing", None, True), ], - ids=["anthropic-pre-call-list", "responses-pre-call-enum", "anthropic-post-call-only", "chat-pre-call"], ) - async def test_guardrail_rewritten_requests_never_replay_the_wire_body(self, call_type, guardrail_mode, sampled): - """The proxy snapshots the wire body before the guardrail pre-call hook, so the - wire-sourced surfaces skip requests a request-mutating guardrail ran on rather - than replay stripped tools or unmasked content; chat sources the dispatched - call and keeps sampling, as do requests only response-mode guardrails touched.""" - hook_kwargs = _success_kwargs( + async def test_guardrail_replay_requires_current_approved_snapshot( + self, + call_type: str, + guardrail_mode: str | list[str], + checkpoint: Literal["absent", "corrupt", "missing", "unapproved", "approved"], + later_mode: str | None, + sampled: bool, + ) -> None: + history: Final[list[dict[str, object]]] = [{"guardrail_name": "g", "guardrail_mode": guardrail_mode}] + if checkpoint == "corrupt": + history[0]["guardrail_response"] = history + body: Final[dict[str, object]] = { + "model": "model", + "messages": [{"role": "user", "content": "approved input"}], + "input": "approved input", + } + snapshot: Final = ( + GuardrailRequestSnapshot.capture(body, {"standard_logging_guardrail_information": history}) + if checkpoint in ("approved", "corrupt") else None + ) + if checkpoint == "corrupt": + assert snapshot is None + base_kwargs: Final = _success_kwargs( call_type=call_type, request_metadata={ - "standard_logging_guardrail_information": [{"guardrail_name": "g", "guardrail_mode": guardrail_mode}] + "standard_logging_guardrail_information": history + + ([{"guardrail_name": "g", "guardrail_mode": later_mode}] if later_mode else []) }, ) - response = RESPONSE - if call_type == "anthropic_messages": - hook_kwargs["messages"] = [{"role": "user", "content": [{"type": "text", "text": "hi"}]}] - elif call_type == "aresponses": - hook_kwargs["messages"] = "hi" - response = RESPONSES_API_RESPONSE + hook_kwargs: Final = { + **base_kwargs, + "messages": "hi" if call_type == "aresponses" else base_kwargs["messages"], + "litellm_params": { + **base_kwargs["litellm_params"], + "proxy_server_request": None if checkpoint == "absent" else {"body": body}, + }, + } - prisma, router = await self._drive(hook_kwargs, response) + prisma, router = await self._drive( + hook_kwargs, + RESPONSES_API_RESPONSE if call_type == "aresponses" else RESPONSE, + guardrail_snapshot=snapshot, + ) if sampled: + assert router.acompletion.call_count == 2 prisma.db.litellm_shadowevalattempt.create.assert_called_once() else: router.acompletion.assert_not_called() prisma.db.litellm_shadowevalattempt.create.assert_not_called() + @pytest.mark.parametrize("call_type", ["anthropic_messages", "aresponses"]) + @pytest.mark.parametrize("remove_optional_fields", [False, True]) + async def test_approved_guardrail_snapshot_replays_independent_native_input( + self, call_type: str, remove_optional_fields: bool + ) -> None: + is_responses: Final = call_type == "aresponses" + metadata: Final = { + "standard_logging_guardrail_information": [{"guardrail_name": "g", "guardrail_mode": "pre_call"}] + } + live_message: Final = {"role": "user", "content": "approved input"} + live_tool: Final = { + "name": "approved_tool", + "description": "approved tool", + "strict": False, + "parameters" if is_responses else "input_schema": {"type": "object", "properties": {}}, + **({"type": "function"} if is_responses else {}), + } + data: Final[dict[str, object]] = { + "model": "model", + "input" if is_responses else "messages": [live_message], + "max_output_tokens" if is_responses else "max_tokens": 123, + **({} if remove_optional_fields else { + "instructions" if is_responses else "system": "approved system", + "tools": [live_tool], + "temperature": 0.2, + }), + } + snapshot: Final = GuardrailRequestSnapshot.capture(data, metadata) + assert snapshot is not None + live_message["content"] = "changed after checkpoint" + live_tool["name"] = "changed_after_checkpoint" + base_kwargs: Final = _success_kwargs(call_type=call_type, request_metadata=metadata) + hook_kwargs: Final = { + **base_kwargs, + "messages": "stale input" if is_responses else [{"role": "user", "content": "stale input"}], + "system": "stale system", + "instructions": "stale system", + "standard_logging_object": { + **base_kwargs["standard_logging_object"], + "model_parameters": {"tools": [{"name": "stale_tool"}], "temperature": 0.9, "max_tokens": 999}, + }, + "litellm_params": {**base_kwargs["litellm_params"], "proxy_server_request": {"body": data}}, + } + + prisma, router = await self._drive( + hook_kwargs, RESPONSES_API_RESPONSE if is_responses else RESPONSE, guardrail_snapshot=snapshot + ) + + assert router.acompletion.call_count == 2 + shadow_call: Final = router.acompletion.call_args_list[0].kwargs + assert shadow_call["messages"] == ( + [] if remove_optional_fields else [{"role": "system", "content": "approved system"}] + ) + [{"role": "user", "content": "approved input"}] + assert shadow_call["max_tokens"] == 123 + assert {key: shadow_call[key] for key in ("tools", "temperature") if key in shadow_call} == ( + {} if remove_optional_fields else { + "temperature": 0.2, + "tools": [{"type": "function", "function": { + "name": "approved_tool", "description": "approved tool", "strict": False, + "parameters": {"type": "object", "properties": {}}, + }}], + } + ) + prisma.db.litellm_shadowevalattempt.create.assert_called_once() + @pytest.mark.parametrize( "call_type,messages,response_obj", [ diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 3d5c38c3acd..23c01841b1b 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -2473,6 +2473,7 @@ def test_success_handler_skips_guardrail_logging_hook_when_disabled(logging_obj) from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger + from litellm.integrations.shadow_eval_logger import GuardrailRequestSnapshot from litellm.types.guardrails import GuardrailEventHooks class DummyGuardrail(CustomGuardrail): @@ -2482,6 +2483,12 @@ def test_success_handler_skips_guardrail_logging_hook_when_disabled(logging_obj) pass logging_obj.stream = False + snapshot: Final = GuardrailRequestSnapshot.capture( + {"messages": [{"role": "user", "content": "approved"}]}, + {"standard_logging_guardrail_information": [{"guardrail_mode": "pre_call"}]}, + ) + assert snapshot is not None + logging_obj.shadow_eval_request_snapshot = snapshot model_response = ModelResponse( id="resp-guardrail-skip", @@ -2523,6 +2530,7 @@ def test_success_handler_skips_guardrail_logging_hook_when_disabled(logging_obj) assert guardrail_call_kwargs["event_type"] == GuardrailEventHooks.logging_only guardrail.logging_hook.assert_not_called() dummy_logger.logging_hook.assert_called_once() + assert logging_obj.shadow_eval_request_snapshot is snapshot def test_success_handler_runs_guardrail_logging_hook_when_enabled(logging_obj): @@ -2530,12 +2538,18 @@ def test_success_handler_runs_guardrail_logging_hook_when_enabled(logging_obj): import datetime from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.integrations.shadow_eval_logger import GuardrailRequestSnapshot from litellm.types.guardrails import GuardrailEventHooks class DummyGuardrail(CustomGuardrail): pass logging_obj.stream = False + logging_obj.shadow_eval_request_snapshot = GuardrailRequestSnapshot.capture( + {"messages": [{"role": "user", "content": "approved"}]}, + {"standard_logging_guardrail_information": [{"guardrail_mode": "pre_call"}]}, + ) + assert logging_obj.shadow_eval_request_snapshot is not None model_response = ModelResponse( id="resp-guardrail-run", @@ -2580,6 +2594,88 @@ def test_success_handler_runs_guardrail_logging_hook_when_enabled(logging_obj): assert guardrail_call_kwargs["event_type"] == GuardrailEventHooks.logging_only guardrail.logging_hook.assert_called_once() assert logging_obj.model_call_details.get("guardrail_hook_ran") is True + assert logging_obj.shadow_eval_request_snapshot is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("hook_mode", ["disabled", "mask", "raises"]) +@pytest.mark.parametrize("stream", [False, True]) +async def test_shadow_snapshot_stays_private_and_is_invalidated_before_logging_guardrails( + monkeypatch: pytest.MonkeyPatch, hook_mode: Literal["disabled", "mask", "raises"], stream: bool +) -> None: + from litellm.caching.in_memory_cache import InMemoryCache + from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.integrations.shadow_eval_logger import GuardrailRequestSnapshot, ShadowEvalLogger + from litellm.types.guardrails import GuardrailEventHooks + + shadow_snapshots: Final[list[GuardrailRequestSnapshot | None]] = [] + hook_snapshots: Final[list[GuardrailRequestSnapshot | None]] = [] + other_payloads: Final[list[Mapping[str, object]]] = [] + prisma_reads: Final[list[bool]] = [] + + def no_prisma() -> None: + prisma_reads.append(True) + + class RecordingShadowLogger(ShadowEvalLogger): + async def async_log_success_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: object, + end_time: object, *, guardrail_snapshot: GuardrailRequestSnapshot | None = None, + ) -> None: + shadow_snapshots.append(guardrail_snapshot) + await super().async_log_success_event( + kwargs, response_obj, start_time, end_time, guardrail_snapshot=guardrail_snapshot + ) + + class RecordingLogger(CustomLogger): + async def async_log_success_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object, + ) -> None: + other_payloads.append(kwargs) + + class LoggingGuardrail(CustomGuardrail): + async def async_logging_hook( + self, kwargs: dict[str, object], result: object, call_type: str, + ) -> tuple[dict[str, object], object]: + hook_snapshots.append(logging_obj.shadow_eval_request_snapshot) + if hook_mode == "raises": + raise RuntimeError("logging guardrail failed without recording history") + return {**kwargs, "messages": [{"role": "user", "content": "masked"}]}, result + + metadata: Final = { + "standard_logging_guardrail_information": [{"guardrail_mode": "pre_call"}], + "user_api_key_hash": "test-key", + } + snapshot: Final = GuardrailRequestSnapshot.capture( + {"messages": [{"role": "user", "content": "snapshot-only"}]}, metadata, + ) + assert snapshot is not None + shadow: Final = RecordingShadowLogger(prisma_provider=no_prisma, jobs_cache=InMemoryCache()) + guardrail: Final = LoggingGuardrail( + guardrail_name="late-mask", default_on=True, + event_hook=GuardrailEventHooks.pre_call if hook_mode == "disabled" else GuardrailEventHooks.logging_only, + ) + monkeypatch.setattr(litellm, "_async_success_callback", []) + logging_obj: Final = LitellmLogging( + model="test-model", messages=[], stream=stream, call_type="anthropic_messages", + start_time=datetime.datetime.now(), litellm_call_id="private-snapshot", function_id="private-snapshot", + dynamic_async_success_callbacks=[shadow, RecordingLogger(), guardrail], + ) + logging_obj.update_messages([{"role": "user", "content": "logged input"}]) + logging_obj.update_environment_variables(litellm_params={"metadata": metadata}, optional_params={}) + logging_obj.shadow_eval_request_snapshot = snapshot + payload: Final = { + "id": "private-snapshot", "call_type": "anthropic_messages", "metadata": metadata, + "model_group": "test-model", "model_parameters": {}, + } + + await logging_obj.async_success_handler(result=ModelResponse(), standard_logging_object=payload) + + assert shadow_snapshots == ([snapshot] if hook_mode == "disabled" else [None]) + assert hook_snapshots == ([] if hook_mode == "disabled" else [None]) + assert prisma_reads == ([True] if hook_mode == "disabled" else []) + assert len(other_payloads) == 1 + assert "snapshot-only" not in json.dumps(other_payloads[0], default=str) + assert "snapshot-only" not in json.dumps(logging_obj.model_call_details, default=str) def test_get_user_agent_tags(): diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 0ef02f76a4f..7b74e69685c 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -3,7 +3,7 @@ import copy import datetime import json from types import MappingProxyType, SimpleNamespace -from typing import AsyncGenerator, Callable, Final, Iterator, Optional, Sequence +from typing import AsyncGenerator, Callable, Final, Iterator, Literal, Optional, Sequence from urllib.parse import unquote_plus from unittest.mock import AsyncMock, MagicMock, patch @@ -495,60 +495,92 @@ class TestProxyBaseLLMRequestProcessing: add_litellm_data_to_request.assert_not_awaited() @pytest.mark.asyncio + @pytest.mark.parametrize("safe_memory_mode", [False, True]) + @pytest.mark.parametrize( + "route_type,input_key,system_key,token_key", + [ + ("acompletion", "messages", "system", "max_tokens"), + ("anthropic_messages", "messages", "system", "max_tokens"), + ("aresponses", "input", "instructions", "max_output_tokens"), + ], + ) async def test_common_processing_pre_call_logic_refreshes_proxy_server_request_body_after_guardrails( - self, monkeypatch - ): - """ - A guardrail (e.g. Presidio PII masking) mutates data["messages"] in place inside - pre_call_hook. The proxy_server_request.body snapshot is taken before that hook - runs, so it must be refreshed afterward or SpendLogs (when store_prompts_in_spend_logs - is enabled) persists the raw pre-guardrail body, bypassing the masking entirely. - """ - processing_obj = ProxyBaseLLMRequestProcessing(data={}) - mock_request = MagicMock(spec=Request) + self, + monkeypatch: pytest.MonkeyPatch, + safe_memory_mode: bool, + route_type: Literal["acompletion", "anthropic_messages", "aresponses"], + input_key: str, + system_key: str, + token_key: str, + ) -> None: + from litellm.integrations.shadow_eval_logger import request_guardrail_fingerprint + + monkeypatch.setattr(litellm, "safe_memory_mode", safe_memory_mode) + processing_obj: Final = ProxyBaseLLMRequestProcessing(data={}) + mock_request: Final = MagicMock(spec=Request) mock_request.headers = {} + metadata_key: Final = "metadata" if route_type == "acompletion" else "litellm_metadata" + raw_body: Final = { + input_key: [{"role": "user", "content": "private input"}], + system_key: "private system", + "tools": [{"name": "private", "description": "private tool"}], + "tool_choice": {"type": "tool", "name": "private"}, + token_key: 100, + } + approved_messages: Final = [{"role": "user", "content": ""}] + approved_tools: Final = [{"name": "allowed", "description": ""}] + approved_body: Final = {input_key: approved_messages, "tools": approved_tools, token_key: 64} + recorded: Final = [{"guardrail_name": "mask", "guardrail_mode": "pre_call", "guardrail_status": "success"}] - raw_messages = [{"role": "user", "content": "my ssn is 123-45-6789"}] - - async def mock_add_litellm_data_to_request(*args, **kwargs): + async def mock_pre_call_hook( + user_api_key_dict: UserAPIKeyAuth, + data: dict[str, object], + call_type: str, + skip_guardrails: bool = False, + ) -> dict[str, object]: + logging_obj: Final = data["litellm_logging_obj"] + assert isinstance(logging_obj, LiteLLMLoggingObj) + assert logging_obj.shadow_eval_request_snapshot is None return { - "messages": raw_messages, - "proxy_server_request": { - "url": "http://testserver/chat/completions", - "method": "POST", - "body": {"messages": raw_messages}, - }, + **{key: value for key, value in data.items() if key not in (system_key, "tool_choice")}, + **approved_body, + metadata_key: {"standard_logging_guardrail_information": recorded}, } - async def mock_pre_call_hook(user_api_key_dict, data, call_type, skip_guardrails=False): - data["messages"] = [{"role": "user", "content": "my ssn is "}] - return data - - mock_proxy_logging_obj = MagicMock(spec=ProxyLogging) + mock_proxy_logging_obj: Final = MagicMock(spec=ProxyLogging) mock_proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=mock_pre_call_hook) monkeypatch.setattr( litellm.proxy.common_request_processing, "add_litellm_data_to_request", - mock_add_litellm_data_to_request, + AsyncMock(return_value={**raw_body, metadata_key: {}, "proxy_server_request": {"body": raw_body}}), ) - returned_data, _ = await processing_obj.common_processing_pre_call_logic( + returned_data, logging_obj = await processing_obj.common_processing_pre_call_logic( request=mock_request, general_settings={}, user_api_key_dict=MagicMock(spec=UserAPIKeyAuth), proxy_logging_obj=mock_proxy_logging_obj, proxy_config=MagicMock(spec=ProxyConfig), - route_type="acompletion", + route_type=route_type, ) - persisted_body = returned_data["proxy_server_request"]["body"] - assert persisted_body["messages"] == returned_data["messages"] - assert "123-45-6789" not in json.dumps(persisted_body["messages"]) - # litellm_logging_obj is stamped onto `data` by function_setup between the - # initial snapshot and pre_call_hook; it must never leak into the persisted - # audit body, which needs to stay plain-JSON-serializable end to end. + proxy_request: Final = returned_data["proxy_server_request"] + persisted_body: Final = proxy_request["body"] + snapshot: Final = logging_obj.shadow_eval_request_snapshot + expected_content: Final = copy.deepcopy(approved_body) + assert snapshot is not None + assert {key: persisted_body[key] for key in raw_body if key in persisted_body} == expected_content + assert {key: snapshot.body[key] for key in raw_body if key in snapshot.body} == expected_content + assert snapshot.fingerprint == request_guardrail_fingerprint( + {"standard_logging_guardrail_information": recorded} + ) assert "litellm_logging_obj" not in persisted_body - json.dumps(persisted_body) + assert "private" not in json.dumps(persisted_body) + approved_messages[0]["content"] = "later input mutation" + approved_tools[0]["description"] = "later tool mutation" + assert {key: snapshot.body[key] for key in raw_body if key in snapshot.body} == expected_content + assert persisted_body[input_key][0]["content"] == "later input mutation" + assert persisted_body["tools"][0]["description"] == "later tool mutation" @staticmethod def _guardrail_tag_budget_harness( diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 0d4b9e8d21f..b2241191ced 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -5,6 +5,7 @@ import os import time from datetime import datetime, timezone from types import SimpleNamespace +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -807,6 +808,9 @@ async def test_add_litellm_data_to_request_body_snapshot_excludes_proxy_server_r "model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "hello"}], "api_key": "request-key", + "proxy_server_request": { + "body": {"messages": [{"role": "user", "content": "forged"}]}, + }, } user_api_key_dict = UserAPIKeyAuth( @@ -836,6 +840,77 @@ async def test_add_litellm_data_to_request_body_snapshot_excludes_proxy_server_r ) assert "api_key" not in snapshot_body assert updated["proxy_server_request"]["credential_fields"] == ("api_key",) + assert snapshot_body["messages"] == [{"role": "user", "content": "hello"}] + + +def test_initial_snapshot_refresh_clears_a_previous_guardrail_checkpoint() -> None: + from litellm.integrations.shadow_eval_logger import GuardrailRequestSnapshot + from litellm.litellm_core_utils.litellm_logging import Logging + from litellm.proxy.litellm_pre_call_utils import refresh_proxy_server_request_body_snapshot + + logging_obj: Final = Logging( + model="test-model", messages=[], stream=False, call_type="acompletion", + start_time=datetime.now(), litellm_call_id="new-request", function_id="new-request", + ) + logging_obj.shadow_eval_request_snapshot = GuardrailRequestSnapshot.capture( + {"messages": [{"role": "user", "content": "previous request"}]}, + {"standard_logging_guardrail_information": [{"guardrail_mode": "pre_call"}]}, + ) + assert logging_obj.shadow_eval_request_snapshot is not None + proxy_request: Final = {"body": {}} + data: Final = { + "messages": [{"role": "user", "content": "new request"}], + "proxy_server_request": proxy_request, + "litellm_logging_obj": logging_obj, + } + + refresh_proxy_server_request_body_snapshot(data) + + assert logging_obj.shadow_eval_request_snapshot is None + assert proxy_request == {"body": {"messages": [{"role": "user", "content": "new request"}]}} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("pre_call_ran", [False, True]) +async def test_post_guardrail_snapshot_preserves_logging_only_masking_in_spend_logs( + monkeypatch: pytest.MonkeyPatch, pre_call_ran: bool +) -> None: + from litellm.litellm_core_utils.litellm_logging import Logging + from litellm.proxy.guardrails.guardrail_hooks.presidio import _OPTIONAL_PresidioPIIMasking + from litellm.proxy.litellm_pre_call_utils import refresh_proxy_server_request_body_snapshot + from litellm.proxy.spend_tracking.spend_tracking_utils import _get_proxy_server_request_for_spend_logs_payload + + monkeypatch.setenv("STORE_PROMPTS_IN_SPEND_LOGS", "true") + messages: Final = [{"role": "user", "content": "email probe@example.invalid"}] + metadata: Final = { + "standard_logging_guardrail_information": [{"guardrail_mode": "pre_call"}] if pre_call_ran else [] + } + data: Final = {"messages": messages, "metadata": metadata, "proxy_server_request": {}} + logging_obj: Final = Logging( + model="test-model", messages=messages, stream=False, call_type="acompletion", + start_time=datetime.now(), litellm_call_id="mask-spend", function_id="mask-spend", kwargs=data, + ) + data["litellm_logging_obj"] = logging_obj + refresh_proxy_server_request_body_snapshot(data, guardrails_applied=True) + logging_obj.update_messages(messages) + snapshot: Final = logging_obj.shadow_eval_request_snapshot + assert (snapshot is not None) is pre_call_ran + guardrail: Final = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, logging_only=True, mock_redacted_text={"text": "email [EMAIL]", "items": []} + ) + + kwargs, _ = await guardrail.async_logging_hook( + kwargs=logging_obj.model_call_details, result=None, call_type="acompletion" + ) + stored: Final = json.loads(_get_proxy_server_request_for_spend_logs_payload( + metadata={}, litellm_params=kwargs["litellm_params"], kwargs=kwargs, + )) + + assert kwargs["messages"] == [{"role": "user", "content": "email [EMAIL]"}] + assert stored["messages"] == kwargs["messages"] + if snapshot is not None: + assert snapshot.body["messages"] == [{"role": "user", "content": "email probe@example.invalid"}] + assert "probe@example.invalid" not in json.dumps(stored) def test_refresh_proxy_server_request_body_snapshot_picks_up_guardrail_masking(): @@ -2850,7 +2925,7 @@ def test_add_headers_to_llm_call_by_model_group_existing_headers_in_data(): litellm.model_group_settings = original_model_group_settings -from typing import Final, Optional +from typing import Optional from fastapi.responses import Response From b1194aa37320e074b473a4c62d89afc3cedc8097 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Wed, 23 Sep 2026 17:49:11 -0700 Subject: [PATCH 07/10] fix(ui): show ten prompt caching requests per page (#42638) --- ...tCachingRequestsTable.integration.test.tsx | 34 +++++++++++++++++-- .../PromptCachingRequestsTable.tsx | 2 +- 2 files changed, 33 insertions(+), 3 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingRequestsTable.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingRequestsTable.integration.test.tsx index 833a46ce16f..e7594a18105 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingRequestsTable.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingRequestsTable.integration.test.tsx @@ -24,7 +24,7 @@ const request = (overrides: Partial = {}): CacheRequest => ({ ...overrides, }); const response = (requests: CacheRequest[], nextCursor: RequestsResponse["next_cursor"] = null) => { - const body: RequestsResponse = { requests, has_more: nextCursor !== null, next_cursor: nextCursor, page_size: 50 }; + const body: RequestsResponse = { requests, has_more: nextCursor !== null, next_cursor: nextCursor, page_size: 10 }; return Response.json(body); }; const lastQuery = () => new URL(String(fetchMock.mock.calls.at(-1)?.[0]), "http://localhost").searchParams; @@ -82,6 +82,36 @@ describe("PromptCachingRequestsTable", () => { expect(fetchMock.mock.calls[0][1]?.headers).toEqual(expect.objectContaining({ Authorization: "Bearer token-a" })); }); + it("shows ten requests per page and keeps the remaining request reachable", async () => { + const rows = Array.from({ length: 11 }, (_, index) => request({ request_id: `request-${index + 1}` })); + fetchMock.mockImplementation(async (input) => { + const query = new URL(String(input), "http://localhost").searchParams; + const start = rows.findIndex((row) => row.request_id === query.get("cursor_request_id")) + 1; + const end = start + Number(query.get("page_size")); + const page = rows.slice(start, end); + const last = page.at(-1); + return response( + page, + end < rows.length && last ? { start_time: last.start_time, request_id: last.request_id } : null, + ); + }); + renderWithProviders(); + + const table = await screen.findByRole("table", { name: "Prompt caching requests" }); + expect(within(table).getAllByRole("link")).toHaveLength(10); + expect(within(table).queryByRole("link", { name: "request-11" })).not.toBeInTheDocument(); + fireEvent.click(screen.getByRole("button", { name: "Next" })); + await screen.findByRole("link", { name: "request-11" }); + expect(within(screen.getByRole("table", { name: "Prompt caching requests" })).getAllByRole("link")).toHaveLength(1); + expect(screen.getByRole("button", { name: "Next" })).toBeDisabled(); + fireEvent.click(screen.getByRole("button", { name: "Previous" })); + await screen.findByRole("link", { name: "request-1" }); + expect(within(screen.getByRole("table", { name: "Prompt caching requests" })).getAllByRole("link")).toHaveLength( + 10, + ); + expect(screen.getByRole("button", { name: "Previous" })).toBeDisabled(); + }); + it("forwards complete server cursors, goes back to prior cursors, and clears them for each caching filter", async () => { fetchMock.mockImplementation(async (input) => { const query = new URL(String(input), "http://localhost").searchParams; @@ -141,7 +171,7 @@ describe("PromptCachingRequestsTable", () => { fireEvent.click(screen.getByRole("tab", { name: "Cache hits" })); await screen.findByRole("link", { name: "hits-1" }); expect(lastQuery().get("filter")).toBe("hits"); - expect(lastQuery().get("page_size")).toBe("50"); + expect(lastQuery().get("page_size")).toBe("10"); expect(screen.getByText("Page 1")).toBeInTheDocument(); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingRequestsTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingRequestsTable.tsx index 29aa9252e7b..140c11d2318 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingRequestsTable.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingRequestsTable.tsx @@ -52,7 +52,7 @@ export default function PromptCachingRequestsTable({ accessToken, dateValue }: P start_date: startDate, end_date: endDate, filter, - page_size: 50, + page_size: 10, cursor_start_time: cursor?.start_time, cursor_request_id: cursor?.request_id, }; From 1175559c397b9f35faaca8622b5165d786d2a23b Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 17:50:09 -0700 Subject: [PATCH 08/10] feat(lint): add LIT013 flagging *-ok suppressions that suppress nothing and remove the 240 stale ones (#42793) --- litellm/_logging.py | 4 +- .../transformation.py | 4 +- litellm/experimental_mcp_client/client.py | 6 +- litellm/integrations/custom_guardrail.py | 2 +- .../integrations/newrelic/newrelic_metrics.py | 4 +- .../integrations/otel/plumbing/providers.py | 4 +- .../integrations/otel/presets/destinations.py | 2 +- litellm/integrations/prometheus.py | 2 +- litellm/integrations/shadow_eval_logger.py | 7 +- .../interactions/background_cost_polling.py | 4 +- litellm/litellm_core_utils/core_helpers.py | 2 +- .../get_supported_openai_params.py | 4 +- .../json_fragment_accumulator.py | 32 +- litellm/litellm_core_utils/litellm_logging.py | 2 +- .../convert_dict_to_response.py | 4 +- .../prompt_templates/common_utils.py | 4 +- .../litellm_core_utils/provider_affinity.py | 2 +- .../streaming_chunk_builder_utils.py | 4 +- .../chat/guardrail_translation/handler.py | 44 +- litellm/llms/anthropic/common_utils.py | 2 +- .../messages/response_cache.py | 4 +- .../messages/streaming_iterator.py | 2 +- .../responses_adapters/transformation.py | 6 +- .../llms/azure/passthrough/transformation.py | 4 +- .../image_generation/flux_transformation.py | 4 +- .../base_llm/passthrough/transformation.py | 4 +- .../bedrock/messages/mantle_transformation.py | 2 +- litellm/llms/bedrock/realtime/handler.py | 4 +- .../flux_lora_depth_transformation.py | 6 +- .../llms/fal_ai/image_edit/transformation.py | 6 +- .../gpt_image_2_transformation.py | 8 +- litellm/llms/gigachat/authenticator.py | 2 +- litellm/llms/gigachat/chat/streaming.py | 8 +- .../llms/mistral/batches/transformation.py | 2 +- .../mongodb/vector_stores/transformation.py | 2 +- .../nvidia_nim/passthrough/transformation.py | 2 +- .../llms/nvidia_nim/rerank/transformation.py | 2 +- .../llms/openai/chat/gpt_transformation.py | 4 +- litellm/llms/openai/openai.py | 8 +- litellm/llms/snowflake/chat/transformation.py | 34 +- .../text_to_speech/transformation.py | 16 +- .../xai/audio_transcription/transformation.py | 4 +- litellm/ocr/main.py | 4 +- litellm/passthrough/main.py | 12 +- .../_experimental/mcp_server/contracts.py | 16 +- litellm/proxy/_experimental/mcp_server/db.py | 2 - .../mcp_server/legacy_callbacks.py | 4 +- .../_experimental/mcp_server/mcp_debug.py | 8 +- .../_experimental/mcp_server/operations.py | 4 +- .../_experimental/mcp_server/tool_search.py | 6 +- .../proxy/agent_endpoints/agent_registry.py | 8 +- litellm/proxy/auth/auth_utils.py | 2 +- litellm/proxy/client/cli/commands/agents.py | 4 +- .../client/cli/commands/claude_settings.py | 2 +- .../client/cli/commands/codex_settings.py | 1 - litellm/proxy/client/cli/commands/pi.py | 6 +- .../auth_cache_invalidation_pubsub.py | 8 +- .../proxy/common_utils/reset_budget_job.py | 20 +- litellm/proxy/common_utils/sse_keepalive.py | 8 +- litellm/proxy/db/baseline_accounting.py | 6 +- litellm/proxy/db/shadow_eval_funnel.py | 2 +- .../guardrails/guardrail_hooks/alice/alice.py | 2 +- .../guardrail_hooks/bedrock_guardrails.py | 10 +- .../guardrails/guardrail_hooks/presidio.py | 2 +- .../proxy/hooks/autorouter_baseline_cache.py | 6 +- litellm/proxy/hooks/batch_rate_limiter.py | 6 +- .../hooks/parallel_request_limiter_v3.py | 4 +- litellm/proxy/litellm_pre_call_utils.py | 8 +- .../auto_router_endpoints.py | 6 +- .../config_override_endpoints.py | 8 +- .../management_v1/budgets.py | 2 +- .../mcp_management_endpoints.py | 8 +- .../management_endpoints/scim/scim_v2.py | 1 - .../management_endpoints/team_endpoints.py | 6 +- .../management_helpers/bulk_user_creation.py | 2 +- .../batch_guardrails.py | 6 +- .../llm_passthrough_endpoints.py | 16 +- .../vertex_passthrough_logging_handler.py | 4 +- .../managed_id_rewriter.py | 8 +- .../pass_through_endpoints.py | 4 +- .../streaming_handler.py | 4 +- .../proxy/policy_engine/pipeline_executor.py | 10 +- litellm/proxy/proxy_server.py | 10 +- litellm/proxy/rag_endpoints/endpoints.py | 2 +- .../proxy/response_api_endpoints/endpoints.py | 6 +- .../spend_tracking/carried_budget_state.py | 4 +- litellm/proxy/utils.py | 21 +- litellm/rerank_api/main.py | 4 +- litellm/responses/additional_tools.py | 7 +- .../transformation.py | 4 +- litellm/responses/streaming_iterator.py | 12 +- litellm/responses/utils.py | 2 +- litellm/router.py | 30 +- .../complexity_router/complexity_router.py | 4 +- litellm/router_strategy/tag_based_routing.py | 4 +- .../auto_router_tuning_baseline.py | 2 +- .../router_utils/fallback_event_handlers.py | 2 +- litellm/rust_bridge/lifecycle.py | 2 +- litellm/rust_bridge/logger.py | 2 +- .../auto_router_endpoints.py | 2 +- .../managed_id_rewriter.py | 4 +- litellm/types/utils.py | 8 +- litellm/utils.py | 12 +- scripts/check_type_discipline.py | 419 ++++++++++-------- scripts/type_discipline_gate.py | 36 +- tests/e2e/lifecycle.py | 7 +- tests/e2e/load/proxy_usage.py | 2 +- tests/integration/conftest.py | 1 - .../test_exception_handler_reconnect_retry.py | 21 +- .../guardrail_hooks/test_conduct.py | 8 +- .../test_check_type_discipline.py | 104 +++-- tests/test_litellm_rust/support/isolation.py | 4 +- tests/unit/messages/test_dispatch.py | 4 +- type-discipline-budget.json | 3 + 114 files changed, 540 insertions(+), 732 deletions(-) diff --git a/litellm/_logging.py b/litellm/_logging.py index 802b01b2e90..c65795babff 100644 --- a/litellm/_logging.py +++ b/litellm/_logging.py @@ -631,9 +631,9 @@ class LevelRoutingStreamHandler(logging.StreamHandler): ) preferred: Final = sys.stdout if is_stdout_record else sys.stderr if preferred is None or getattr(preferred, "closed", False): - self.stream = sys.stderr # rebind-ok: fall back to the pre-fix stream rather than raising per record + self.stream = sys.stderr else: - self.stream = preferred # rebind-ok: StreamHandler.emit writes self.stream under the handler lock + self.stream = preferred super().emit(record) diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 4c321b12573..31af5a144eb 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -191,8 +191,6 @@ def _as_chat_reasoning_items( ) -> list[ChatCompletionReasoningItem] | None: if not reasoning_items: return None - # cast-ok: _BuiltReasoningItem is the structural shape ChatCompletionReasoningItem - # describes, and TypedDict invariance is what stops the two from unifying here. return cast(list[ChatCompletionReasoningItem], list(reasoning_items)) @@ -1370,7 +1368,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): if tool_call_index_map is None: return output_index if output_index not in tool_call_index_map: - tool_call_index_map[output_index] = len(tool_call_index_map) # mutable-ok: per-stream accumulator state + tool_call_index_map[output_index] = len(tool_call_index_map) return tool_call_index_map[output_index] @staticmethod diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 1ccae8de35f..1206f9abcbd 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -764,7 +764,7 @@ class MCPClient: follow_redirects=True, event_hooks=MappingProxyType( {"response": [capture_upstream_error_response], "request": [guard] if guard else []} - ), # mutable-ok: httpx types require lists of hooks + ), ) return factory @@ -921,9 +921,7 @@ class MCPClient: with anyio.fail_after(max(self.timeout, MCP_TOOL_LISTING_TIMEOUT)): for page_index in range(MCP_TOOL_LISTING_MAX_PAGES): try: - page = await fetch_page( # rebind-ok: each SDK page replaces the previous one - None if cursor is None else PaginatedRequestParams(cursor=cursor) - ) + page = await fetch_page(None if cursor is None else PaginatedRequestParams(cursor=cursor)) except MCPError as error: if page_index > 0 and error.error.code == METHOD_NOT_FOUND: raise RuntimeError("MCP list operation became unavailable during pagination") from error diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index ffa0bc36f6b..5d64eff526b 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -1641,5 +1641,5 @@ def log_guardrail_information(func): return async_wrapper(*args, **kwargs) return sync_wrapper(*args, **kwargs) - vars(wrapper)[LOGS_GUARDRAIL_INFORMATION_MARKER] = True # rebind-ok: stamps the wrapper this call just built + vars(wrapper)[LOGS_GUARDRAIL_INFORMATION_MARKER] = True return wrapper diff --git a/litellm/integrations/newrelic/newrelic_metrics.py b/litellm/integrations/newrelic/newrelic_metrics.py index da952b78d3f..0a45a7e52c3 100644 --- a/litellm/integrations/newrelic/newrelic_metrics.py +++ b/litellm/integrations/newrelic/newrelic_metrics.py @@ -366,9 +366,7 @@ class NewRelicMetricsLogger(CustomBatchLogger): error to keep the client-error path (drop) distinct from 5xx (retry).""" payload: Final = build_metric_payload(records=batch, window_start=window_start, now=time.time()) try: - status = ( - await self.async_send_compressed_data(payload) - ).status_code # rebind-ok: reassigned from the raised HTTPStatusError below + status = (await self.async_send_compressed_data(payload)).status_code except HTTPStatusError as e: status = e.response.status_code except Exception as e: # noqa: BLE001 # transport/network failure re-queues the batch diff --git a/litellm/integrations/otel/plumbing/providers.py b/litellm/integrations/otel/plumbing/providers.py index e3474edaf14..8bac36aad76 100644 --- a/litellm/integrations/otel/plumbing/providers.py +++ b/litellm/integrations/otel/plumbing/providers.py @@ -353,7 +353,7 @@ class _DrainPool: def _drain_until_closed(self) -> None: while True: - processor: SpanProcessor | None = self._pending.get() # rebind-ok: loop variable + processor: SpanProcessor | None = self._pending.get() if processor is None: return _shutdown_quietly(processor) @@ -572,7 +572,7 @@ class TenantFanOutSpanProcessor(SpanProcessor): span, destination.span_scope ): continue - processor = self._acquire(destination) # rebind-ok: loop variable; pyright forbids Final in a loop + processor = self._acquire(destination) if processor is None: continue try: diff --git a/litellm/integrations/otel/presets/destinations.py b/litellm/integrations/otel/presets/destinations.py index 63801e623af..e6cb775af1d 100644 --- a/litellm/integrations/otel/presets/destinations.py +++ b/litellm/integrations/otel/presets/destinations.py @@ -151,7 +151,7 @@ def destination_for( endpoint, protocol = resolved return OtelDestination( endpoint=endpoint, - headers=MappingProxyType(dict(headers)), # mutable-ok: MappingProxyType needs a concrete mapping to wrap + headers=MappingProxyType(dict(headers)), resource_attributes=MappingProxyType({"service.name": service_name}) if service_name else _NO_ATTRS, callback_name=callback_name, protocol=protocol, diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index d62a6c3427a..2fdcb8ef745 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -131,7 +131,7 @@ def _paginated_table(repository: BaseRepository[_TableRowT]) -> _PaginatedPrisma """View a repository's prisma table through the pagination surface budget metrics need.""" return cast( _PaginatedPrismaTable[_TableRowT], - repository.table, # cast-ok: prisma rows carry the budget columns the domain model declares + repository.table, ) diff --git a/litellm/integrations/shadow_eval_logger.py b/litellm/integrations/shadow_eval_logger.py index 84106b77a7d..19d9bee7493 100644 --- a/litellm/integrations/shadow_eval_logger.py +++ b/litellm/integrations/shadow_eval_logger.py @@ -873,7 +873,6 @@ class ShadowEvalLogger(CustomLogger): await prisma.db.litellm_shadowevalattempt.group_by( by=["job_id"], count=True, - # mutable-ok: Prisma aggregate spec sum={"judge_cost": True, "shadow_cost": True, "shadow_classifier_cost": True}, where={"job_id": {"in": [str(record.id) for record in records]}}, # mutable-ok: Prisma filter ) @@ -901,7 +900,7 @@ class ShadowEvalLogger(CustomLogger): {target: tuple(job for _, job in group) for target, group in groupby(by_target, key=itemgetter(0))} ) await self._jobs_cache.async_set_cache(_JOBS_CACHE_KEY, jobs) - self._job_starts = {} # rebind-ok: new generation, counts absorbed into the fill + self._job_starts = {} return jobs except Exception as e: # noqa: BLE001 # a DB blip must never break request logging verbose_logger.debug("shadow_eval: active-job read failed: %s", e) @@ -1033,7 +1032,7 @@ class ShadowEvalLogger(CustomLogger): real_cache_hit=real_cache_hit, control_tier=control_tier, shadow_params=shadow_params, - parent_metadata=MappingProxyType(dict(request_metadata)), # mutable-ok: frozen snapshot + parent_metadata=MappingProxyType(dict(request_metadata)), ) ).add_done_callback(self._release_shadow_slot) except Exception as e: # noqa: BLE001 # logging hooks must never fail the request @@ -1347,7 +1346,7 @@ class ShadowEvalLogger(CustomLogger): { "role": "user", "content": _judge_user_prompt(conversation, response_a, response_b, _tool_definitions_text(tools)), - }, # mutable-ok: SDK message + }, ] try: response: Final = await judge_acompletion( diff --git a/litellm/interactions/background_cost_polling.py b/litellm/interactions/background_cost_polling.py index 51325354e7d..b48c7c03573 100644 --- a/litellm/interactions/background_cost_polling.py +++ b/litellm/interactions/background_cost_polling.py @@ -79,9 +79,7 @@ async def _fetch_interaction(context: BackgroundInteractionPollContext) -> Inter custom_llm_provider=context.custom_llm_provider, api_key=context.api_key, api_base=context.api_base, - **{ - "no-log": True - }, # mutable-ok: "no-log" is not a valid identifier, so it can only be passed through a mapping + **{"no-log": True}, ) diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index d7fbe9f7e09..a2d40279c49 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -764,4 +764,4 @@ def set_response_cost_in_hidden_params(response: _CarriesHiddenParams, cost: flo **(additional_headers if isinstance(additional_headers, Mapping) else _NO_HEADERS), RESPONSE_COST_HEADER: cost, } - hidden_params["additional_headers"] = merged # rebind-ok: the caller's record is the point + hidden_params["additional_headers"] = merged diff --git a/litellm/litellm_core_utils/get_supported_openai_params.py b/litellm/litellm_core_utils/get_supported_openai_params.py index 08b8816e17d..680f31a797f 100644 --- a/litellm/litellm_core_utils/get_supported_openai_params.py +++ b/litellm/litellm_core_utils/get_supported_openai_params.py @@ -32,9 +32,7 @@ def get_supported_openai_params( - None if unmapped """ if not custom_llm_provider: - custom_llm_provider = declared_authenticating_provider( - model - ) # rebind-ok: resolving would run the provider's OAuth flow + custom_llm_provider = declared_authenticating_provider(model) if not custom_llm_provider: try: custom_llm_provider = litellm.get_llm_provider(model=model)[1] diff --git a/litellm/litellm_core_utils/json_fragment_accumulator.py b/litellm/litellm_core_utils/json_fragment_accumulator.py index 81d18dd0119..e262f05932c 100644 --- a/litellm/litellm_core_utils/json_fragment_accumulator.py +++ b/litellm/litellm_core_utils/json_fragment_accumulator.py @@ -21,20 +21,18 @@ class JSONFragmentAccumulator: def __init__(self) -> None: self._chunks: list[str] = [] # mutable-ok: O(1) append; string concat would copy the buffer each time - self._buffer: str = ( - "" # mutable-ok: lazily materialized join of _chunks, rebuilt only when _chunks is non-empty - ) - self._offset: int = 0 # mutable-ok: cursor past already-consumed values; avoids re-slicing on every pop - self._could_close: bool = False # mutable-ok: cached heuristic; rescanning past fragments was itself O(n^2) + self._buffer: str = "" + self._offset: int = 0 + self._could_close: bool = False def __bool__(self) -> bool: return bool(self._chunks) or self._offset < len(self._buffer) def append(self, fragment: str) -> None: - self._chunks.append(fragment) # mutable-ok: see __init__ + self._chunks.append(fragment) stripped: Final = fragment.rstrip() if stripped: - self._could_close = stripped[-1] in ("}", "]") # mutable-ok: see __init__ + self._could_close = stripped[-1] in ("}", "]") def could_close_json(self) -> bool: """ @@ -50,8 +48,8 @@ class JSONFragmentAccumulator: if not self._chunks: return unconsumed: Final = self._buffer[self._offset :] - self._buffer = unconsumed + "".join(self._chunks) # mutable-ok: merge pending fragments, once per append batch - self._offset = 0 # mutable-ok: see __init__ + self._buffer = unconsumed + "".join(self._chunks) + self._offset = 0 self._chunks = [] # mutable-ok: see __init__ def pop_next_value(self) -> tuple[bool, object]: @@ -69,7 +67,7 @@ class JSONFragmentAccumulator: while start < length and self._buffer[start].isspace(): start += 1 if start >= length: - self._offset = start # mutable-ok: see __init__ + self._offset = start return False, None decoder: Final = json.JSONDecoder() try: @@ -77,11 +75,11 @@ class JSONFragmentAccumulator: except json.JSONDecodeError: return False, None decoded, end_index = cast("tuple[object, int]", raw_value) # cast-ok: raw_decode returns tuple[Any, int] - self._offset = end_index # mutable-ok: see __init__ + self._offset = end_index if self._offset >= len(self._buffer): - self._buffer = "" # mutable-ok: see __init__ - self._offset = 0 # mutable-ok: see __init__ - self._could_close = False # mutable-ok: buffer is empty, nothing can close + self._buffer = "" + self._offset = 0 + self._could_close = False return True, decoded def snapshot(self) -> str: @@ -91,7 +89,7 @@ class JSONFragmentAccumulator: def set(self, value: str) -> None: """Replace the buffer's contents with a single fragment.""" self._chunks = [] # mutable-ok: see __init__ - self._buffer = value # mutable-ok: see __init__ - self._offset = 0 # mutable-ok: see __init__ + self._buffer = value + self._offset = 0 stripped: Final = value.rstrip() - self._could_close = bool(stripped) and stripped[-1] in ("}", "]") # mutable-ok: see __init__ + self._could_close = bool(stripped) and stripped[-1] in ("}", "]") diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index b038762a476..0603414cabd 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -6667,7 +6667,7 @@ def get_standard_logging_object_payload( "version": 3, "status": "unknown", "reason": "pending_projection", - } # mutable-ok: spend-log JSON serialization requires plain mappings + } if captured_baseline is not None else ( { # mutable-ok: spend-log JSON serialization requires plain mappings diff --git a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py index 87524d86c61..9ea730a873f 100644 --- a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py +++ b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py @@ -372,9 +372,7 @@ from collections import defaultdict def _handle_invalid_parallel_tool_calls( - tool_calls: list[ - ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall - ], # mutable-ok: patched in place via slice assignment + tool_calls: list[ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall], ): """ Handle hallucinated parallel tool call from openai - https://community.openai.com/t/model-tries-to-call-unknown-function-multi-tool-use-parallel/490653 diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 6c45622649f..378295e1b7a 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -208,7 +208,7 @@ def _content_parts_contain_image(parts: Sequence[object]) -> bool: for _ in range(_IMAGE_SCAN_MAX_DEPTH): if any(isinstance(part, Mapping) and part.get("type") in _IMAGE_CONTENT_PART_TYPES for part in frontier): return True - frontier = tuple( # rebind-ok: depth-bounded frontier walk + frontier = tuple( nested for part in frontier if isinstance(part, Mapping) @@ -2020,7 +2020,7 @@ def _anthropic_content_lists(messages: Sequence[object]) -> Iterator[object]: def _strip_encrypted_reasoning_from_blocks(content: object) -> None: blocks: Final = cast(list[object], content) # cast-ok: narrowed by the caller's isinstance kept: Final = tuple(block for block in blocks if not is_encrypted_reasoning_block(block)) - blocks[:] = kept # rebind-ok: shared with fallback snapshot + blocks[:] = kept def _reasoning_replay_group_key(indexed_block: tuple[int, Mapping[str, object]]) -> str: diff --git a/litellm/litellm_core_utils/provider_affinity.py b/litellm/litellm_core_utils/provider_affinity.py index 31cd9a7ff69..33bf2ee7079 100644 --- a/litellm/litellm_core_utils/provider_affinity.py +++ b/litellm/litellm_core_utils/provider_affinity.py @@ -83,7 +83,7 @@ def get_stable_session_id(litellm_params: object | None) -> str | None: return None -def add_provider_affinity_header( # mutable-ok: downstream handlers add auth and signing headers +def add_provider_affinity_header( headers: Mapping[str, object], litellm_params: object | None ) -> dict[str, object]: # mutable-ok: downstream handlers add auth and signing headers header_name: Final = _get_provider_affinity_header_name(litellm_params) diff --git a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py index aa4e0cf5495..d975c3551f3 100644 --- a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py +++ b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py @@ -475,9 +475,7 @@ class ChunkProcessor: def get_combined_tool_content( self, tool_call_chunks: Sequence["_ToolCallChunk"] - ) -> list[ - ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall - ]: # mutable-ok: assigned verbatim to Message.tool_calls, a list field + ) -> list[ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall]: tool_calls_list: list[ ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall ] = [] # mutable-ok: see return type diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 24ff63c9433..a78f633f5d7 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -199,9 +199,7 @@ def _write_back_system_block(system: object, block_idx: int, response: str) -> N return text_blocks: Final = tuple(block for block in system if isinstance(block, dict) and block.get("type") == "text") if block_idx < len(text_blocks): - text_blocks[block_idx]["text"] = ( - response # mutable-ok: guardrails rewrite the caller's request payload in place - ) + text_blocks[block_idx]["text"] = response def _write_back_message_text(message: _WritableMessage, target: MessageTextTarget, response: str) -> None: @@ -211,22 +209,16 @@ def _write_back_message_text(message: _WritableMessage, target: MessageTextTarge match target: case MessageContentTarget(): if isinstance(content, str): - message["content"] = response # mutable-ok: guardrails rewrite the caller's request payload in place + message["content"] = response case ContentBlockTextTarget(content_idx=content_idx): if isinstance(content, list): - content[content_idx]["text"] = ( - response # mutable-ok: guardrails rewrite the caller's request payload in place - ) + content[content_idx]["text"] = response case ToolResultStringTarget(content_idx=content_idx): if isinstance(content, list): - content[content_idx]["content"] = ( - response # mutable-ok: guardrails rewrite the caller's request payload in place - ) + content[content_idx]["content"] = response case ToolResultBlockTextTarget(content_idx=content_idx, block_idx=block_idx): if isinstance(content, list): - content[content_idx]["content"][block_idx]["text"] = ( - response # mutable-ok: guardrails rewrite the caller's request payload in place - ) + content[content_idx]["content"][block_idx]["text"] = response case _: assert_never(target) @@ -248,9 +240,9 @@ def _write_back_tool_use( block: Final = content[target.content_idx] if isinstance(content, list) else None if not isinstance(block, dict): return - block["input"] = rewritten_input # mutable-ok: guardrails rewrite the caller's request payload in place + block["input"] = rewritten_input if shape.name is not None and shape.name != block.get("name"): - block["name"] = shape.name # mutable-ok: guardrails rewrite the caller's request payload in place + block["name"] = shape.name @dataclass(frozen=True, slots=True) @@ -603,13 +595,9 @@ class AnthropicMessagesHandler(BaseTranslation): *(item for one_message in extracted for item in one_message.scanned), ) texts_to_check: Final = [item.text for item in scanned] # mutable-ok: GenericGuardrailAPIInputs takes list[str] - images_to_check: Final = [ - image for one_message in extracted for image in one_message.images - ] # mutable-ok: GenericGuardrailAPIInputs takes list[str] + images_to_check: Final = [image for one_message in extracted for image in one_message.images] scanned_tool_calls: Final = tuple(item for one_message in extracted for item in one_message.tool_calls) - tool_calls_to_check: Final = [ - item.tool_call for item in scanned_tool_calls - ] # mutable-ok: GenericGuardrailAPIInputs takes list[ChatCompletionToolCallChunk] + tool_calls_to_check: Final = [item.tool_call for item in scanned_tool_calls] pre_guardrail_tool_calls: Final = _tool_call_shapes(tool_calls_to_check) # Step 2: Apply guardrail to all texts and tool calls in batch @@ -697,9 +685,7 @@ class AnthropicMessagesHandler(BaseTranslation): return data - def _hoisted_top_level_system_message( - self, data: dict - ) -> AllMessageValues | None: # mutable-ok: API message payload + def _hoisted_top_level_system_message(self, data: dict) -> AllMessageValues | None: """Return the system message produced by translating the top-level prompt.""" system: Final = data.get("system") if not system: @@ -736,7 +722,7 @@ class AnthropicMessagesHandler(BaseTranslation): if isinstance(content, str): return ( {"role": "system", "content": content} if content else None # mutable-ok: API message payload - ) # mutable-ok: API message payload + ) if not isinstance(content, list): return None blocks: Final[list[dict[str, object]]] = [] # mutable-ok: API message payload @@ -749,14 +735,14 @@ class AnthropicMessagesHandler(BaseTranslation): anthropic_block: dict[str, object] = { # mutable-ok: API message payload "type": "text", "text": text, - } # mutable-ok: API message payload + } cache_control = block.get("cache_control") if cache_control: anthropic_block["cache_control"] = deepcopy(cache_control) blocks.append(anthropic_block) return ( {"role": "system", "content": blocks} if blocks else None # mutable-ok: API message payload - ) # mutable-ok: API message payload + ) @staticmethod def _fold_leading_systems_into_top_level( @@ -1098,9 +1084,7 @@ class AnthropicMessagesHandler(BaseTranslation): match item.target: case SystemStringTarget(): if isinstance(data.get("system"), str): - data["system"] = ( - guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place - ) + data["system"] = guardrail_response case SystemBlockTextTarget(block_idx=block_idx): _write_back_system_block(data.get("system"), block_idx, guardrail_response) case ( diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index c0e6006633e..bf2d588dd3a 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -1591,7 +1591,7 @@ def _flatten_web_search_results_in_message(message: object) -> object: return {**message, "content": [b for b in rewritten if b is not None]} # mutable-ok: JSON wire format -def flatten_unencrypted_web_search_results_in_anthropic_messages( # mutable-ok: as sibling sanitizers +def flatten_unencrypted_web_search_results_in_anthropic_messages( messages: list[Any], ) -> list[Any]: """ diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/response_cache.py b/litellm/llms/anthropic/experimental_pass_through/messages/response_cache.py index 86dfe8ff451..dc2d4408c20 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/response_cache.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/response_cache.py @@ -88,9 +88,7 @@ class AnthropicMessagesStreamCacheWriter: try: events: Final = _split_sse_events(collected_stream.decode("utf-8")) - cached_payload: Final = { - CACHED_STREAM_EVENTS_KEY: events - } # mutable-ok: cache backends serialize plain dicts + cached_payload: Final = {CACHED_STREAM_EVENTS_KEY: events} await litellm.cache.async_add_cache( cached_payload, dynamic_cache_object=self.caching_handler.dual_cache, diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py index 98c5c6d6d4e..0bd46382fef 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py @@ -186,7 +186,7 @@ def _sse_event(event_type: str, payload: Mapping[str, object]) -> bytes: def _incomplete_stream_error_sse_event() -> bytes: - return _sse_event( # mutable-ok: one-shot JSON payload, never mutated after construction + return _sse_event( "error", {"type": "error", "error": {"type": "api_error", "message": INCOMPLETE_STREAM_ERROR_MESSAGE}}, ) diff --git a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py index 1fdb0318bab..e3d3425f8a6 100644 --- a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py @@ -148,13 +148,11 @@ class LiteLLMAnthropicToResponsesAPIAdapter: if isinstance(content, str): return ( [{"type": "input_text", "text": content}] if content else [] # mutable-ok: API message payload - ) # mutable-ok: API message payload + ) if not isinstance(content, list): return [] # mutable-ok: API message payload return [ # mutable-ok: API message payload - with_prompt_cache_breakpoint( - {"type": "input_text", "text": text}, block.get("prompt_cache_breakpoint") - ) # mutable-ok: API message payload + with_prompt_cache_breakpoint({"type": "input_text", "text": text}, block.get("prompt_cache_breakpoint")) for block in content if isinstance(block, dict) and block.get("type") == "text" and (text := block.get("text")) # pyright: ignore[reportUnnecessaryIsInstance] # untrusted client payload ] diff --git a/litellm/llms/azure/passthrough/transformation.py b/litellm/llms/azure/passthrough/transformation.py index c40cefecdd0..a648a24f5e3 100644 --- a/litellm/llms/azure/passthrough/transformation.py +++ b/litellm/llms/azure/passthrough/transformation.py @@ -59,9 +59,7 @@ def logged_responses_stream(all_chunks: Sequence[str], logging_obj: Logging) -> terminal_event: Final = OpenAIResponsesAPIConfig.parse_terminal_event_from_stream_chunks(all_chunks=all_chunks) if terminal_event is None: return None - logging_obj.call_type = ( - RESPONSES_RELAY_SHAPE.call_type.value - ) # rebind-ok: routes cost calculation to the relayed shape's pricing path + logging_obj.call_type = RESPONSES_RELAY_SHAPE.call_type.value return terminal_event diff --git a/litellm/llms/azure_ai/image_generation/flux_transformation.py b/litellm/llms/azure_ai/image_generation/flux_transformation.py index ac9ec24420b..b6a9caf147b 100644 --- a/litellm/llms/azure_ai/image_generation/flux_transformation.py +++ b/litellm/llms/azure_ai/image_generation/flux_transformation.py @@ -73,9 +73,7 @@ class AzureFoundryFluxImageGenerationConfig(GPTImageGenerationConfig): normalized_model: Final = model.lower().replace(".", "-").replace("_", "-") return "flux-2-flex" if "flux-2-flex" in normalized_model else "flux-2-pro" - def get_supported_openai_params( # mutable-ok: inherited config contract returns a list - self, model: str - ) -> list[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: if not self.is_flux2_model(model): return super().get_supported_openai_params(model) return [ # mutable-ok: BaseImageGenerationConfig requires a list diff --git a/litellm/llms/base_llm/passthrough/transformation.py b/litellm/llms/base_llm/passthrough/transformation.py index f2a12c3f22d..84cbd4204e3 100644 --- a/litellm/llms/base_llm/passthrough/transformation.py +++ b/litellm/llms/base_llm/passthrough/transformation.py @@ -95,9 +95,7 @@ def logged_relay_shape( parsed: Final = shape.parse(body) except ValidationError: return None - logging_obj.call_type = ( - shape.call_type.value - ) # rebind-ok: routes cost calculation to the relayed shape's pricing path + logging_obj.call_type = shape.call_type.value return parsed diff --git a/litellm/llms/bedrock/messages/mantle_transformation.py b/litellm/llms/bedrock/messages/mantle_transformation.py index 052eb90a833..66744275778 100644 --- a/litellm/llms/bedrock/messages/mantle_transformation.py +++ b/litellm/llms/bedrock/messages/mantle_transformation.py @@ -45,7 +45,7 @@ def _move_betas_into_header(request: Mapping[str, object], headers: dict[str, st if betas: headers["anthropic-beta"] = ",".join(betas) # rebind-ok: the handler signs and sends this same dict return - headers.pop("anthropic-beta", None) # rebind-ok: a caller header Mantle rejects in full must not reach it + headers.pop("anthropic-beta", None) class AmazonMantleMessagesConfig(AmazonAnthropicClaudeMessagesConfig): diff --git a/litellm/llms/bedrock/realtime/handler.py b/litellm/llms/bedrock/realtime/handler.py index d17590bdaaa..049313c3c96 100644 --- a/litellm/llms/bedrock/realtime/handler.py +++ b/litellm/llms/bedrock/realtime/handler.py @@ -489,9 +489,7 @@ class BedrockRealtime(BaseAWSLLM): parsed_client_message = _parse_client_message(message) is_session_update = _json_str(parsed_client_message.get("type")) == "session.update" if is_session_update: - client_ws.scope[BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY] = ( - message # rebind-ok: scope outlives the attempt - ) + client_ws.scope[BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY] = message transformed_messages = transformation_config.transform_realtime_request( message=message, diff --git a/litellm/llms/fal_ai/image_edit/flux_lora_depth_transformation.py b/litellm/llms/fal_ai/image_edit/flux_lora_depth_transformation.py index fa469d638d2..0b6205ff302 100644 --- a/litellm/llms/fal_ai/image_edit/flux_lora_depth_transformation.py +++ b/litellm/llms/fal_ai/image_edit/flux_lora_depth_transformation.py @@ -27,7 +27,7 @@ class FalAIFluxLoraDepthEditConfig(FalAIImageEditConfig): def get_supported_openai_params(self, model: str) -> list: # mutable-ok: base class contract returns a list return list(SUPPORTED_OPENAI_PARAMS) # mutable-ok: base class contract returns a list - def map_openai_params( # mutable-ok: base class contract returns a dict + def map_openai_params( self, image_edit_optional_params: ImageEditOptionalRequestParams, model: str, @@ -63,9 +63,7 @@ class FalAIFluxLoraDepthEditConfig(FalAIImageEditConfig): if len(images) > 1: raise ValueError(f"{FLUX_LORA_DEPTH_ENDPOINT} accepts exactly one control image") provider_params: Final[Mapping[str, object]] = MappingProxyType( - { - key: value for key, value in image_edit_optional_request_params.items() if key != "mask" - } # mutable-ok: frozen by MappingProxyType + {key: value for key, value in image_edit_optional_request_params.items() if key != "mask"} ) request_body: Final[dict[str, object]] = { # mutable-ok: base class contract returns a dict "prompt": prompt, diff --git a/litellm/llms/fal_ai/image_edit/transformation.py b/litellm/llms/fal_ai/image_edit/transformation.py index 6e6a872839a..839c15c4c28 100644 --- a/litellm/llms/fal_ai/image_edit/transformation.py +++ b/litellm/llms/fal_ai/image_edit/transformation.py @@ -84,7 +84,7 @@ class FalAIImageEditConfig(BaseImageEditConfig): def get_supported_openai_params(self, model: str) -> list: # mutable-ok: base class contract returns a list return list(SUPPORTED_OPENAI_PARAMS) # mutable-ok: base class contract returns a list - def map_openai_params( # mutable-ok: base class contract returns a dict + def map_openai_params( self, image_edit_optional_params: ImageEditOptionalRequestParams, model: str, @@ -146,9 +146,7 @@ class FalAIImageEditConfig(BaseImageEditConfig): MappingProxyType({"mask_url": to_data_url(mask)}) if mask is not None else MappingProxyType({}) ) provider_params: Final[Mapping[str, object]] = MappingProxyType( - { - key: value for key, value in image_edit_optional_request_params.items() if key != "mask" - } # mutable-ok: frozen by MappingProxyType + {key: value for key, value in image_edit_optional_request_params.items() if key != "mask"} ) request_body: Final[dict[str, object]] = { # mutable-ok: base class contract returns a dict "prompt": prompt, diff --git a/litellm/llms/fal_ai/image_generation/gpt_image_2_transformation.py b/litellm/llms/fal_ai/image_generation/gpt_image_2_transformation.py index ca301662cf8..0d008555f8b 100644 --- a/litellm/llms/fal_ai/image_generation/gpt_image_2_transformation.py +++ b/litellm/llms/fal_ai/image_generation/gpt_image_2_transformation.py @@ -101,12 +101,10 @@ class FalAIGPTImage2Config(FalAIBaseConfig): endpoint: Final[str] = model if model.startswith(self.MODEL_PREFIX) else f"{self.MODEL_PREFIX}{model}" return f"{base_url}/{endpoint}" - def get_supported_openai_params( # mutable-ok: base class contract returns a list - self, model: str - ) -> list[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: return list(SUPPORTED_OPENAI_PARAMS) # mutable-ok: base class contract returns a list - def map_openai_params( # mutable-ok: base class contract returns a dict + def map_openai_params( self, non_default_params: Mapping[str, object], optional_params: Mapping[str, object], @@ -138,7 +136,7 @@ class FalAIGPTImage2Config(FalAIBaseConfig): return map_gpt_image_quality(value, model) return value - def transform_image_generation_request( # mutable-ok: base class contract returns a dict + def transform_image_generation_request( self, model: str, prompt: str, diff --git a/litellm/llms/gigachat/authenticator.py b/litellm/llms/gigachat/authenticator.py index 73086ba395b..a85dcd9c70d 100644 --- a/litellm/llms/gigachat/authenticator.py +++ b/litellm/llms/gigachat/authenticator.py @@ -243,7 +243,7 @@ def _parse_token_response(response: httpx.Response) -> tuple[str, int]: ) # expires_at is in milliseconds - expires_at: int # rebind-ok: conditionally assigned from str or int + expires_at: int if isinstance(expires_at_raw, str): expires_at = int(expires_at_raw) # rebind-ok: conditionally assigned from str or int else: diff --git a/litellm/llms/gigachat/chat/streaming.py b/litellm/llms/gigachat/chat/streaming.py index 0a4cbd8e520..908412d9c31 100644 --- a/litellm/llms/gigachat/chat/streaming.py +++ b/litellm/llms/gigachat/chat/streaming.py @@ -30,7 +30,7 @@ class GigaChatModelResponseIterator: def chunk_parser(self, chunk: Mapping[str, object]) -> GenericStreamingChunk: """Parse a single streaming chunk from GigaChat.""" - choices: Sequence = chunk.get("choices") or () # mutable-ok: tuple literal as default + choices: Sequence = chunk.get("choices") or () if not choices: return GenericStreamingChunk( text="", @@ -56,7 +56,7 @@ class GigaChatModelResponseIterator: if chunk_finish_reason == "function_call" and isinstance(raw_function_call, Mapping) and raw_function_call: func_call: Final[Mapping[str, object]] = raw_function_call args_raw: Final[object] = func_call.get("arguments") or {} - args_str: str # rebind-ok: conditionally assigned from dict or str + args_str: str if isinstance(args_raw, dict): args_str = json.dumps(args_raw, ensure_ascii=False) # rebind-ok: build from dict else: @@ -80,10 +80,10 @@ class GigaChatModelResponseIterator: usage = convert_usage(validated_usage) _prompt_details: dict | None = ( usage.prompt_tokens_details.model_dump() if usage.prompt_tokens_details else None - ) # rebind-ok: conditional + ) _completion_details: dict | None = ( usage.completion_tokens_details.model_dump() if usage.completion_tokens_details else None - ) # rebind-ok: conditional + ) usage_block = ChatCompletionUsageBlock( # pyright: ignore[reportCallIssue] # TypedDict kwarg constructor prompt_tokens=usage.prompt_tokens, completion_tokens=usage.completion_tokens, diff --git a/litellm/llms/mistral/batches/transformation.py b/litellm/llms/mistral/batches/transformation.py index ef9ee5ff503..d3ed6a3af62 100644 --- a/litellm/llms/mistral/batches/transformation.py +++ b/litellm/llms/mistral/batches/transformation.py @@ -33,7 +33,7 @@ OpenAIBatchStatus: TypeAlias = Literal[ "validating", "failed", "in_progress", "finalizing", "completed", "expired", "cancelling", "cancelled" ] -_NO_HEADERS: Final[Mapping[str, str]] = MappingProxyType({}) # mutable-ok: frozen at module scope +_NO_HEADERS: Final[Mapping[str, str]] = MappingProxyType({}) _STATUS_MAP: Final[MappingProxyType[MistralBatchStatus, OpenAIBatchStatus]] = MappingProxyType( { "QUEUED": "validating", diff --git a/litellm/llms/mongodb/vector_stores/transformation.py b/litellm/llms/mongodb/vector_stores/transformation.py index a59f39d3be8..94d3aef48cc 100644 --- a/litellm/llms/mongodb/vector_stores/transformation.py +++ b/litellm/llms/mongodb/vector_stores/transformation.py @@ -197,7 +197,7 @@ class MongoDBVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig): **headers, "Authorization": f"Bearer {api_key}", "Content-Type": "application/json", - } # mutable-ok: writable HTTP headers + } def get_complete_url(self, api_base: str | None, litellm_params: Mapping[str, object]) -> str: if not api_base: diff --git a/litellm/llms/nvidia_nim/passthrough/transformation.py b/litellm/llms/nvidia_nim/passthrough/transformation.py index 7de1ce4d631..e8e7da8e10b 100644 --- a/litellm/llms/nvidia_nim/passthrough/transformation.py +++ b/litellm/llms/nvidia_nim/passthrough/transformation.py @@ -110,7 +110,7 @@ class NvidiaNimPassthroughConfig(BasePassthroughConfig): return { **headers, "Authorization": f"Bearer {api_key}", - } # mutable-ok: base class contract returns dict for httpx + } @staticmethod def get_api_base(api_base: str | None = None) -> str | None: diff --git a/litellm/llms/nvidia_nim/rerank/transformation.py b/litellm/llms/nvidia_nim/rerank/transformation.py index 93e00dad9a1..15cfdb6bece 100644 --- a/litellm/llms/nvidia_nim/rerank/transformation.py +++ b/litellm/llms/nvidia_nim/rerank/transformation.py @@ -215,7 +215,7 @@ class NvidiaNimRerankConfig(BaseRerankConfig): elif isinstance(doc, dict): # Preserve only the structured passage fields supported by the # selected rerank route. - supported_fields: NvidiaNimPassageObject = {} # mutable-ok: assembling a request TypedDict + supported_fields: NvidiaNimPassageObject = {} if "text" in self.SUPPORTED_PASSAGE_FIELDS and "text" in doc: supported_fields["text"] = doc["text"] if "image" in self.SUPPORTED_PASSAGE_FIELDS and "image" in doc: diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index b63684db782..62351d8e39a 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -596,9 +596,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): for choice in choices: ## HANDLE JSON MODE - anthropic returns single function call] tool_calls = choice["message"].get("tool_calls", None) - new_tool_calls: list[ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall] | None = ( - None # mutable-ok: holds _handle_invalid_parallel_tool_calls' list; Message.__init__ expects list - ) + new_tool_calls: list[ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall] | None = None message_content = choice["message"].get("content", None) if tool_calls is not None: _openai_tool_calls = [] diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 7ac0d988074..63874ca9619 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -1427,9 +1427,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): }, ) - request_data: Final = ( # mutable-ok: the OpenAI SDK takes the request body as a dict - {**data, "extra_headers": headers} if headers else data - ) + request_data: Final = {**data, "extra_headers": headers} if headers else data response = await openai_aclient.images.generate(**request_data, timeout=timeout) stringified_response: Final = response.model_dump() ## LOGGING @@ -1513,9 +1511,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): ) ## COMPLETION CALL - request_data: Final = ( # mutable-ok: the OpenAI SDK takes the request body as a dict - {**data, "extra_headers": headers} if headers else data - ) + request_data: Final = {**data, "extra_headers": headers} if headers else data _response: Final = openai_client.images.generate(**request_data, timeout=timeout) response: Final = _response.model_dump() diff --git a/litellm/llms/snowflake/chat/transformation.py b/litellm/llms/snowflake/chat/transformation.py index aeff902f655..734d0e20818 100644 --- a/litellm/llms/snowflake/chat/transformation.py +++ b/litellm/llms/snowflake/chat/transformation.py @@ -118,7 +118,7 @@ def _convert_image_url_to_anthropic(block: Mapping[str, object]) -> object: anthropic_process_openai_file_message({"type": "file", "file": {"file_data": url}}) if select_anthropic_content_block_type_for_file(_data_uri_media_type(url)) == "document" else create_anthropic_image_param( - image_url if isinstance(image_url, dict) else url, # mutable-ok: caller's JSON block + image_url if isinstance(image_url, dict) else url, format=_image_url_field(image_url, "format"), is_bedrock_invoke=True, ) @@ -191,12 +191,8 @@ def _signed_thinking_blocks(msg: object) -> list[dict[str, object]]: # mutable- ] -def _clean_input_schema(schema: object) -> object: # mutable-ok: JSON schema copy - return ( - {key: value for key, value in schema.items() if key != "$schema"} - if isinstance(schema, Mapping) - else schema # mutable-ok: JSON schema copy - ) # mutable-ok: JSON schema copy +def _clean_input_schema(schema: object) -> object: + return {key: value for key, value in schema.items() if key != "$schema"} if isinstance(schema, Mapping) else schema class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): @@ -299,9 +295,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): ) return anthropic_tools - def _extract_system_and_messages( # mutable-ok: JSON wire messages - self, messages: list[AllMessageValues] - ) -> tuple[list[dict] | None, list[dict]]: + def _extract_system_and_messages(self, messages: list[AllMessageValues]) -> tuple[list[dict] | None, list[dict]]: """ Split messages into system prompt and conversation turns for Anthropic format. @@ -330,9 +324,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): { # mutable-ok: JSON wire system block "type": "text", "text": block.get("text", ""), - **( - {"cache_control": block["cache_control"]} if "cache_control" in block else {} - ), # mutable-ok: JSON wire block + **({"cache_control": block["cache_control"]} if "cache_control" in block else {}), } for block in content if isinstance(block, Mapping) and block.get("type") == "text" @@ -372,7 +364,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): ] if isinstance(content, list) else [*thinking_blocks, *([{"type": "text", "text": content}] if content else [])] - ) # rebind-ok: loop-local normalized content + ) conversation.append({"role": "assistant", "content": thinking_content}) else: conversation.append({"role": "assistant", "content": content}) @@ -380,9 +372,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): tool_call_id_value = ( msg.get("tool_call_id", "") if isinstance(msg, dict) else getattr(msg, "tool_call_id", "") ) - tool_call_id = ( - tool_call_id_value if isinstance(tool_call_id_value, str) else "" - ) # rebind-ok: normalized loop value + tool_call_id = tool_call_id_value if isinstance(tool_call_id_value, str) else "" tool_result_block = _convert_tool_result_to_anthropic(content, tool_call_id, msg_cache_control) if ( conversation @@ -395,13 +385,13 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): else: conversation.append( {"role": "user", "content": [tool_result_block]} # mutable-ok: JSON wire message - ) # mutable-ok: JSON wire message + ) else: - conversation.append( # mutable-ok: JSON wire message + conversation.append( { # mutable-ok: JSON wire message "role": role, "content": _convert_image_url_blocks_to_anthropic(content), - } # mutable-ok: JSON wire message + } ) system: Final[list[dict] | None] = system_parts if system_parts else None # mutable-ok: JSON wire messages @@ -516,11 +506,11 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): "messages": conversation, "stream": stream, **optional_params, - **extra_body, # mutable-ok: JSON wire body + **extra_body, } ) if system is not None: - body["system"] = normalize_cache_control_in_anthropic_payload( # mutable-ok: JSON wire payload + body["system"] = normalize_cache_control_in_anthropic_payload( {"system": system} # mutable-ok: JSON wire payload )["system"] diff --git a/litellm/llms/vertex_ai/text_to_speech/transformation.py b/litellm/llms/vertex_ai/text_to_speech/transformation.py index d382f43495f..6c2c59d98e1 100644 --- a/litellm/llms/vertex_ai/text_to_speech/transformation.py +++ b/litellm/llms/vertex_ai/text_to_speech/transformation.py @@ -43,9 +43,7 @@ else: LiteLLMLoggingObj = Any HttpxBinaryResponseContent = Any -_LyriaVoice: TypeAlias = ( - str | dict | None -) # mutable-ok: inherited interface supports structured provider voice dictionaries +_LyriaVoice: TypeAlias = str | dict | None class VertexAITextToSpeechConfig(BaseTextToSpeechConfig, VertexBase): @@ -664,21 +662,15 @@ class VertexAILyriaTextToSpeechConfig(VertexAITextToSpeechConfig): if model_info["vertex_ai_audio_api"] == "lyria_predict": predictions: Final = response_json.get("predictions") or () if predictions: - audio_data = predictions[0].get("audioContent") or predictions[0].get( - "bytesBase64Encoded" - ) # rebind-ok: predict response supplies the generated audio value + audio_data = predictions[0].get("audioContent") or predictions[0].get("bytesBase64Encoded") mime_type = predictions[0].get("mimeType") # rebind-ok: predict response supplies its audio MIME type else: for step in response_json.get("steps") or response_json.get("outputs") or (): content_items = step.get("content") or () if step.get("type") == "model_output" else (step,) for content in content_items: if content.get("type") == "audio" and content.get("data"): - audio_data = content[ - "data" - ] # rebind-ok: interactions response supplies the generated audio value - mime_type = content.get( - "mime_type" - ) # rebind-ok: interactions response supplies its audio MIME type + audio_data = content["data"] + mime_type = content.get("mime_type") if audio_data is None: raise ValueError(f"No generated audio found in Vertex AI {base_model} response") binary_data: Final = base64.b64decode(audio_data) diff --git a/litellm/llms/xai/audio_transcription/transformation.py b/litellm/llms/xai/audio_transcription/transformation.py index feeabed0d9c..49447413a37 100644 --- a/litellm/llms/xai/audio_transcription/transformation.py +++ b/litellm/llms/xai/audio_transcription/transformation.py @@ -168,9 +168,7 @@ class XAIAudioTranscriptionConfig(BaseAudioTranscriptionConfig): for word in payload.words ] - hidden_params: Final[dict[str, object]] = dict( - payload.model_dump(mode="json") - ) # mutable-ok: TranscriptionResponse._hidden_params is a dict + hidden_params: Final[dict[str, object]] = dict(payload.model_dump(mode="json")) if payload.duration is not None: hidden_params["audio_transcription_duration"] = payload.duration response._hidden_params = hidden_params # pyright: ignore[reportPrivateUsage] # TranscriptionResponse exposes no public hidden-params setter diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 06830ed4b53..3ca6c1295c5 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -173,9 +173,7 @@ def _prepare_ocr_request( custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, provider_config=ocr_provider_config, - optional_params=cast( - dict[str, object], optional_params - ), # cast-ok: provider configs return heterogeneous OCR options + optional_params=cast(dict[str, object], optional_params), litellm_params=dict(litellm_params), effective_timeout=effective_timeout, litellm_logging_obj=litellm_logging_obj, diff --git a/litellm/passthrough/main.py b/litellm/passthrough/main.py index 73d8bab686b..ef931827d85 100644 --- a/litellm/passthrough/main.py +++ b/litellm/passthrough/main.py @@ -428,9 +428,7 @@ def llm_passthrough_route( _is_async: Final = bool(kwargs.get("allm_passthrough_route", False)) - litellm_logging_obj: Final = cast( - LiteLLMLoggingObj, kwargs.get("litellm_logging_obj") - ) # cast-ok: logging obj is constructed upstream; tests inject mocks + litellm_logging_obj: Final = cast(LiteLLMLoggingObj, kwargs.get("litellm_logging_obj")) model, custom_llm_provider, api_key, api_base = get_llm_provider( model=model, @@ -516,9 +514,7 @@ def llm_passthrough_route( forward_headers=False, ) - _request_data: dict | None = ( - data if isinstance(data, dict) else (json if isinstance(json, dict) else None) - ) # rebind-ok: conditional + _request_data: dict | None = data if isinstance(data, dict) else (json if isinstance(json, dict) else None) headers, signed_json_body = provider_config.sign_request( headers=headers, litellm_params=litellm_params_dict, @@ -544,9 +540,7 @@ def llm_passthrough_route( ) ## IS STREAMING REQUEST - _streaming_request_data: dict = ( - data if isinstance(data, dict) else (json if isinstance(json, dict) else {}) - ) # rebind-ok: conditional + _streaming_request_data: dict = data if isinstance(data, dict) else (json if isinstance(json, dict) else {}) is_streaming_request: Final = provider_config.is_streaming_request( endpoint=endpoint, request_data=_streaming_request_data, diff --git a/litellm/proxy/_experimental/mcp_server/contracts.py b/litellm/proxy/_experimental/mcp_server/contracts.py index c3129d171ad..1879e285789 100644 --- a/litellm/proxy/_experimental/mcp_server/contracts.py +++ b/litellm/proxy/_experimental/mcp_server/contracts.py @@ -57,24 +57,20 @@ class OperationContext: ) -> tuple[ UserAPIKeyAuth | None, str | None, - list[str] | None, # mutable-ok: detached legacy server-list payload - dict[str, dict[str, str]] | None, # mutable-ok: legacy auth dispatch requires concrete dict headers - dict[str, str] | None, # mutable-ok: detached legacy header payload - dict[str, str] | None, # mutable-ok: detached legacy header payload + list[str] | None, + dict[str, dict[str, str]] | None, + dict[str, str] | None, + dict[str, str] | None, str | None, ]: return ( self.user_api_key_auth, self.mcp_auth_header, list(self.mcp_servers) if self.mcp_servers is not None else None, # mutable-ok: legacy policy list input - { - key: dict(value) for key, value in self.mcp_server_auth_headers.items() - } # mutable-ok: legacy auth dispatch checks concrete dict headers + {key: dict(value) for key, value in self.mcp_server_auth_headers.items()} if self.mcp_server_auth_headers is not None else None, - dict(self.oauth2_headers) - if self.oauth2_headers is not None - else None, # mutable-ok: legacy OAuth header input + dict(self.oauth2_headers) if self.oauth2_headers is not None else None, dict(self.raw_headers) if self.raw_headers is not None else None, # mutable-ok: legacy request header input self.client_ip, ) diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 30ee8b7a4fc..18723a76b2b 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -636,8 +636,6 @@ async def get_all_mcp_servers( where: Final[prisma_db_types.LiteLLM_MCPServerTableWhereInput] = ( {"approval_status": approval_status} if approval_status is not None - # mutable-ok: prisma where-inputs must be plain dicts, and both `NOT` and `not` drop - # NULL rows (measured), so the OR is the only NULL-preserving way to exclude drafts else {"OR": [{"approval_status": None}, {"approval_status": {"not": MCPApprovalStatus.draft}}]} ) mcp_servers: Final = await _db_find_mcp_server_rows(prisma_client, where) diff --git a/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py b/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py index 9e321062643..f52d2a006d2 100644 --- a/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py +++ b/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py @@ -55,9 +55,7 @@ def create_sampling_callback( params=params, default_model=getattr(litellm, "default_mcp_sampling_model", None), user_api_key_auth=captured.user_api_key_auth, - raw_headers=dict(captured.raw_headers) - if captured.raw_headers is not None - else None, # mutable-ok: handler consumes an owned request header dict + raw_headers=dict(captured.raw_headers) if captured.raw_headers is not None else None, client_ip=captured.client_ip, ) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_debug.py b/litellm/proxy/_experimental/mcp_server/mcp_debug.py index ff482b80b50..70da73fa045 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_debug.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_debug.py @@ -174,9 +174,7 @@ class MCPAuthDiagnostics: { "x-mcp-debug-auth-resolution": AuthResolution.multiple.value, "x-mcp-debug-auth-resolutions": json.dumps( - { - server_id: source.value for server_id, source in self._outcomes[:32] - }, # mutable-ok: JSON encoder requires a concrete dict + {server_id: source.value for server_id, source in self._outcomes[:32]}, separators=(",", ":"), ensure_ascii=True, ), @@ -597,9 +595,7 @@ async def capture_upstream_error_response(response: httpx.Response | httpx2.Resp ) except (asyncio.TimeoutError, httpx.HTTPError, httpx.StreamError, httpx2.HTTPError, httpx2.StreamError): response._content = b"" # pyright: ignore[reportPrivateUsage] # rebind-ok: httpx auth retries must survive diagnostic read failures - response.extensions[_CAPTURE_EXTENSION] = ( - "(unavailable: error body read failed)" # rebind-ok: httpx response hooks communicate through extensions - ) + response.extensions[_CAPTURE_EXTENSION] = "(unavailable: error body read failed)" return response.extensions[_CAPTURE_EXTENSION] = preview # rebind-ok: httpx response hooks communicate through extensions diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 26bf68d9932..dcab43bdc76 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -3103,9 +3103,7 @@ class GatewayOperations: return await _execute_mcp_tool( name=operation.name, arguments=dict(operation.arguments), # mutable-ok: existing tool hooks own mutable argument data - allowed_mcp_servers=list( - operation.allowed_mcp_servers - ), # mutable-ok: legacy dispatch list contract + allowed_mcp_servers=list(operation.allowed_mcp_servers), start_time=operation.start_time, user_api_key_auth=auth, mcp_auth_header=token, diff --git a/litellm/proxy/_experimental/mcp_server/tool_search.py b/litellm/proxy/_experimental/mcp_server/tool_search.py index 3650c722103..3c060752934 100644 --- a/litellm/proxy/_experimental/mcp_server/tool_search.py +++ b/litellm/proxy/_experimental/mcp_server/tool_search.py @@ -103,7 +103,7 @@ def _tool_result(tool: Tool) -> ToolSearchResult: "name": tool.name, "description": tool.description or "", "inputSchema": tool.input_schema, - } # mutable-ok: wire schema payload + } def _scored_result(tool: Tool, score: float) -> ToolSearchResult: @@ -112,7 +112,7 @@ def _scored_result(tool: Tool, score: float) -> ToolSearchResult: "description": tool.description or "", "inputSchema": tool.input_schema, "score": score, - } # mutable-ok: wire schema payload + } _MCP_PROXY_IDENTITY_META_KEY: Final[str] = "litellm.ai/proxy_tool_identity" @@ -120,7 +120,7 @@ _MCP_PROXY_IDENTITY_META_KEY: Final[str] = "litellm.ai/proxy_tool_identity" def with_mcp_proxy_identity(tool: Tool, server_id: str) -> Tool: identity: Final[MCPProxyToolIdentity] = {"server_id": server_id, "tool_name": tool.name} - return tool.model_copy( # mutable-ok: Pydantic requires mutable update and metadata mappings + return tool.model_copy( update={ # mutable-ok: Pydantic update payload "meta": {**(tool.meta or {}), _MCP_PROXY_IDENTITY_META_KEY: identity} # mutable-ok: metadata mapping } diff --git a/litellm/proxy/agent_endpoints/agent_registry.py b/litellm/proxy/agent_endpoints/agent_registry.py index d6b12e830e1..3d56c2b5326 100644 --- a/litellm/proxy/agent_endpoints/agent_registry.py +++ b/litellm/proxy/agent_endpoints/agent_registry.py @@ -135,9 +135,7 @@ def _dump_agent_params(raw: Mapping[str, object]) -> dict[str, object]: _AGENT_PARAMS_MASKER: Final = SensitiveDataMasker() _REDACT_AGENT_PARAMS_MAX_DEPTH: Final = 10 -_AGENT_PARAMS_ADAPTER: Final[TypeAdapter[dict[str, object]]] = TypeAdapter( - dict[str, object] -) # mutable-ok: safe_dumps() and AgentResponse.litellm_params both require a real dict, not a Mapping +_AGENT_PARAMS_ADAPTER: Final[TypeAdapter[dict[str, object]]] = TypeAdapter(dict[str, object]) _AGENT_PARAMS_SEQUENCE_ADAPTER: Final[TypeAdapter[tuple[object, ...]]] = TypeAdapter(tuple[object, ...]) _EMPTY_LITELLM_PARAMS: Final[Mapping[str, object]] = MappingProxyType({}) @@ -189,7 +187,7 @@ def _redact_agent_params_tree(value: object, _depth: int) -> object: else _redact_agent_params_tree(nested_value, _depth + 1) ) for key, nested_value in typed_params.items() - } # mutable-ok: consumed by json.dumps()/AgentResponse.litellm_params, both of which require a real dict + } def parse_agent_litellm_params(value: object) -> Mapping[str, object]: @@ -318,7 +316,7 @@ def _restore_redacted_litellm_params( key: value for key in all_keys if (value := _resolved_agent_param_value(key, incoming, existing, _depth)) is not _MISSING_AGENT_PARAM - } # mutable-ok: fed to safe_dumps() for JSON-column storage, which requires a real dict + } class GrantMigrationResult(NamedTuple): diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index c002b2b9508..c0123ae45a3 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -1410,7 +1410,7 @@ def log_once_if_budget_reservation_disabled( "Set disable_budget_reservation to False or remove it to restore " "hard per-request budget enforcement." ) - constants.budget_reservation_disabled_info_emitted = True # rebind-ok: process-wide one-shot sentinel + constants.budget_reservation_disabled_info_emitted = True def is_pass_through_provider_route(route: str) -> bool: diff --git a/litellm/proxy/client/cli/commands/agents.py b/litellm/proxy/client/cli/commands/agents.py index 15b111ff016..7d50131cb88 100644 --- a/litellm/proxy/client/cli/commands/agents.py +++ b/litellm/proxy/client/cli/commands/agents.py @@ -241,9 +241,7 @@ def prepare_codex( _Preparer: TypeAlias = Callable[[str, str, Mapping[str, str]], Sequence[str]] -_PREPARERS: Final[Mapping[str, _Preparer]] = MappingProxyType( - {"pi": prepare_pi, "codex": prepare_codex} # mutable-ok: MappingProxyType freezes the provider registry -) +_PREPARERS: Final[Mapping[str, _Preparer]] = MappingProxyType({"pi": prepare_pi, "codex": prepare_codex}) def agent_launch_args(command: str, base_url: str) -> list[str]: diff --git a/litellm/proxy/client/cli/commands/claude_settings.py b/litellm/proxy/client/cli/commands/claude_settings.py index b3fdc4695cb..13ed483586e 100644 --- a/litellm/proxy/client/cli/commands/claude_settings.py +++ b/litellm/proxy/client/cli/commands/claude_settings.py @@ -621,7 +621,7 @@ def unconfigure_claude_settings( ) target: Final = _write_target(settings_path) file_removed: Final = not settings and not (receipt.file_existed and target.exists()) - kept_receipt: Final = ( # mutable-ok: pydantic serializes the update as given and rejects a mappingproxy + kept_receipt: Final = ( receipt.model_copy(update={"written": {item.key: _fingerprint(absent) for item in withheld}}) if withheld else None diff --git a/litellm/proxy/client/cli/commands/codex_settings.py b/litellm/proxy/client/cli/commands/codex_settings.py index 686eaa47ff0..5b01c31683a 100644 --- a/litellm/proxy/client/cli/commands/codex_settings.py +++ b/litellm/proxy/client/cli/commands/codex_settings.py @@ -106,7 +106,6 @@ def _with(document: TOMLDocument, path: str, snapshot: str | None) -> TOMLDocume if section and section not in document and snapshot is not None: contents: Final = tomlkit.parse(tomlkit.dumps(MappingProxyType({key: tomlkit.parse(snapshot).item("value")}))) return tomlkit.parse(document.as_string() + "\n" + tomlkit.dumps(MappingProxyType({section: contents}))) - # mutable-ok: TOMLKit editing requires private node mutation to preserve comments and order updated: Final = tomlkit.parse(document.as_string()) parent: Final = _table(_mapping(updated).get(section)) if section else updated if parent is None: diff --git a/litellm/proxy/client/cli/commands/pi.py b/litellm/proxy/client/cli/commands/pi.py index 5c749959638..f5834f94fb8 100644 --- a/litellm/proxy/client/cli/commands/pi.py +++ b/litellm/proxy/client/cli/commands/pi.py @@ -175,7 +175,7 @@ def _model_entry( ) output: Final[dict[str, JsonValue]] = ( # mutable-ok: JSON field {"maxTokens": limit.max_tokens} if limit and limit.max_tokens else {} - ) # mutable-ok: JSON field + ) return {"id": model_id, **context, **output} # mutable-ok: JSON serialization requires a mutable object @@ -208,9 +208,7 @@ def sync_models_json( ) -> PiSyncError | None: """Replace only the litellm provider entry, leaving the rest of the file intact.""" try: - current: Final = ( # mutable-ok: JSON object default - _MODELS_FILE_ADAPTER.validate_json(path.read_text()) if path.exists() else {} - ) + current: Final = _MODELS_FILE_ADAPTER.validate_json(path.read_text()) if path.exists() else {} except (OSError, ValidationError) as e: return PiSyncError(f"Could not read {path} as a JSON object: {e}. Fix or move the file, then retry.") existing_providers: Final = current.get("providers", {}) # mutable-ok: JSON object default diff --git a/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py b/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py index 2bb53c7723d..11cb66d1a7f 100644 --- a/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py +++ b/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py @@ -184,17 +184,17 @@ class AuthCacheInvalidationSubscriber: backoff_seconds = _BACKOFF_INITIAL_SECONDS # rebind-ok: exponential backoff accumulator across reconnects while True: try: - client = _pubsub_capable_client(self._redis_cache) # rebind-ok: re-resolved on every reconnect + client = _pubsub_capable_client(self._redis_cache) if client is None: verbose_proxy_logger.warning( "auth cache invalidation subscriber disabled: cluster redis client has no pub/sub support; " "cross-worker eviction falls back to the local cache TTL" ) return - pubsub = client.pubsub() # rebind-ok: fresh pubsub per reconnect + pubsub = client.pubsub() try: await pubsub.subscribe(auth_cache_invalidation_channel(self._redis_cache)) - backoff_seconds = _BACKOFF_INITIAL_SECONDS # rebind-ok: reset after successful subscribe + backoff_seconds = _BACKOFF_INITIAL_SECONDS await self._consume(pubsub) finally: await self._close_pubsub(pubsub) @@ -207,7 +207,7 @@ class AuthCacheInvalidationSubscriber: backoff_seconds, ) await asyncio.sleep(backoff_seconds) - backoff_seconds = min(backoff_seconds * 2, _BACKOFF_MAX_SECONDS) # rebind-ok: backoff accumulator + backoff_seconds = min(backoff_seconds * 2, _BACKOFF_MAX_SECONDS) async def _consume(self, pubsub: _ConfigSyncPubSub) -> None: while True: diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index 3efc189a475..b35b876b475 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -240,12 +240,8 @@ def _queue_budget_linked_resets( one transaction, so the reverse order lets the zero re-match a row the decrement just moved into the (0, cap] range and erase its carried spend.""" for budget_id, cap in cascade.rollover_caps.items(): - writes.queue_spend_zero( - where={"budget_id": budget_id, **extra, "spend": {"gt": 0, "lte": cap}} - ) # mutable-ok: prisma where filter must be a dict - writes.queue_spend_decrement( - where={"budget_id": budget_id, **extra, "spend": {"gt": cap}}, amount=cap - ) # mutable-ok: prisma where filter must be a dict + writes.queue_spend_zero(where={"budget_id": budget_id, **extra, "spend": {"gt": 0, "lte": cap}}) + writes.queue_spend_decrement(where={"budget_id": budget_id, **extra, "spend": {"gt": cap}}, amount=cap) plain_ids: Final = tuple(bid for bid in cascade.budget_ids if bid not in cascade.rollover_caps) if plain_ids: writes.queue_spend_zero(where=_budget_link_where(plain_ids, extra)) @@ -267,16 +263,10 @@ def _queue_enduser_resets(writes: LinkedSpendResetWrites, cascade: "_BudgetCasca return cap: Final = cascade.rollover_caps.get(default_budget_id) if cap is None: - writes.queue_spend_zero( - where={"budget_id": None, **_SPENT_ROWS_WHERE} - ) # mutable-ok: prisma where filter must be a dict + writes.queue_spend_zero(where={"budget_id": None, **_SPENT_ROWS_WHERE}) return - writes.queue_spend_zero( - where={"budget_id": None, "spend": {"gt": 0, "lte": cap}} - ) # mutable-ok: prisma where filter must be a dict - writes.queue_spend_decrement( - where={"budget_id": None, "spend": {"gt": cap}}, amount=cap - ) # mutable-ok: prisma where filter must be a dict + writes.queue_spend_zero(where={"budget_id": None, "spend": {"gt": 0, "lte": cap}}) + writes.queue_spend_decrement(where={"budget_id": None, "spend": {"gt": cap}}, amount=cap) @dataclass(frozen=True, slots=True) diff --git a/litellm/proxy/common_utils/sse_keepalive.py b/litellm/proxy/common_utils/sse_keepalive.py index 26fccf8ee82..cf98a7e9224 100644 --- a/litellm/proxy/common_utils/sse_keepalive.py +++ b/litellm/proxy/common_utils/sse_keepalive.py @@ -65,9 +65,7 @@ async def _keepalive_ping_stream( ping_interval_seconds: float, ping_chunk: str, ) -> AsyncGenerator[str, None]: - pending = asyncio.ensure_future( - stream.__anext__() - ) # rebind-ok: re-armed with the next __anext__ after each delivered chunk + pending = asyncio.ensure_future(stream.__anext__()) try: while True: await asyncio.wait({pending}, timeout=ping_interval_seconds) @@ -125,9 +123,7 @@ async def _keepalive_ping_byte_stream( stream: AsyncGenerator[bytes, None], ping_interval_seconds: float, ) -> AsyncGenerator[bytes, None]: - pending = asyncio.ensure_future( - stream.__anext__() - ) # rebind-ok: re-armed with the next __anext__ after each delivered chunk + pending = asyncio.ensure_future(stream.__anext__()) # The tail of the bytes relayed so far, long enough to hold any delimiter. # Seeded as a delimiter because a stream starts at a frame boundary, and kept # across chunks because a delimiter can be split between two transport reads, diff --git a/litellm/proxy/db/baseline_accounting.py b/litellm/proxy/db/baseline_accounting.py index 4219102d9aa..78036237993 100644 --- a/litellm/proxy/db/baseline_accounting.py +++ b/litellm/proxy/db/baseline_accounting.py @@ -491,7 +491,7 @@ class BaselineAccountingStore: tuple(await db.query_raw(_READ_PAGE, scope, after_revision, cursor, _PAGE_TIMESTAMPS, withdraw_from)) ): yield page - cursor = page[-1].started_at # rebind-ok: keyset pagination advances after each complete timestamp group + cursor = page[-1].started_at async def _withdraw(self, db: SupportsRawQueries, scope: str, started_at: float) -> None: async for page in self._pages(db, scope, 0, withdraw_from=started_at): @@ -623,9 +623,7 @@ async def flush_baseline_accounting(client: PrismaClient) -> None: store: Final = BaselineAccountingStore.for_client(client) async with client.baseline_accounting_lock: batch: Final = tuple(client.baseline_accounting_transactions[:32]) - client.baseline_accounting_transactions = client.baseline_accounting_transactions[ - 32: - ] # rebind-ok: drain under lock + client.baseline_accounting_transactions = client.baseline_accounting_transactions[32:] more_queued: Final = bool(client.baseline_accounting_transactions) try: remaining: Final = await asyncio.wait_for(_flush_records(store, batch), timeout=5) diff --git a/litellm/proxy/db/shadow_eval_funnel.py b/litellm/proxy/db/shadow_eval_funnel.py index 9181d3f5035..3578c3def7e 100644 --- a/litellm/proxy/db/shadow_eval_funnel.py +++ b/litellm/proxy/db/shadow_eval_funnel.py @@ -40,7 +40,7 @@ def pending_shadow_eval_funnel_events() -> int: def record_shadow_eval_funnel_event(job_id: str, stage: ShadowEvalFunnelStage) -> None: """Count one skipped request for one job leg; synchronous so the hook's read-modify- write cannot interleave with the flush's snapshot on the shared event loop.""" - counters: Final = _pending.setdefault(job_id, dict.fromkeys(FUNNEL_STAGES, 0)) # mutable-ok: queue entry + counters: Final = _pending.setdefault(job_id, dict.fromkeys(FUNNEL_STAGES, 0)) counters[stage] += 1 diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice/alice.py b/litellm/proxy/guardrails/guardrail_hooks/alice/alice.py index bcc35e7a22f..287031c3528 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice/alice.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice/alice.py @@ -286,7 +286,7 @@ class AliceGuardrail(CustomGuardrail): text = replacement.get("text") if not (isinstance(index, int) and isinstance(text, str) and 0 <= index < len(texts)): raise self._mask_rejected(verdict) - texts[index] = text # mutable-ok: item assignment into the local working copy above + texts[index] = text inputs["texts"] = texts diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 434c52ca6f3..228b31604a3 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -1218,7 +1218,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): bedrock_request_data: Final = { # mutable-ok: outbound JSON request body **base_request_data, "content": content, - } # mutable-ok: outbound JSON request body + } prepared_request: Final = await run_aws_signing( self._prepare_request, credentials=credentials, @@ -1266,9 +1266,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ) response_usage: Final = bedrock_guardrail_response.get("usage") if isinstance(response_usage, dict): - completed_chunk_usages.append( - response_usage - ) # rebind-ok: accumulator threaded from make_bedrock_api_request, recording this billed call + completed_chunk_usages.append(response_usage) return bedrock_guardrail_response status_code, detail_message = self._parse_bedrock_guardrail_error_response(httpx_response) @@ -2860,9 +2858,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): return except ModifyResponseException as e: if raw_sse: - e.model = _pre_block_response.model or e.model # rebind-ok: exc.model defaults to the guardrail + e.model = _pre_block_response.model or e.model if e.original_response is None: - e.original_response = _pre_block_response # rebind-ok: the block builder reads usage off this + e.original_response = _pre_block_response for block_chunk in AnthropicMessagesHandler().build_block_sse_chunks(e, stream_started=False): yield block_chunk return diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index fe91d6d7a28..eb07c19a580 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -168,7 +168,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): # Per-loop semaphores bounding chunked-analyze fan-out across ALL # concurrent oversized blocks/requests on this instance, not per call - self._loop_chunk_semaphores: _LoopSemaphores = {} # mutable-ok: per-loop semaphore cache + self._loop_chunk_semaphores: _LoopSemaphores = {} if mock_testing is True: # for testing purposes only return diff --git a/litellm/proxy/hooks/autorouter_baseline_cache.py b/litellm/proxy/hooks/autorouter_baseline_cache.py index 8cea7d0e364..0c006730dba 100644 --- a/litellm/proxy/hooks/autorouter_baseline_cache.py +++ b/litellm/proxy/hooks/autorouter_baseline_cache.py @@ -230,12 +230,10 @@ class AutoRouterBaselineCache(CustomLogger): async def invalidate_baseline_cache(logging_obj: Logging, reason: str, *, completed: bool = False) -> None: context: Final = logging_obj.baseline_cache_context if context is not None: - logging_obj.baseline_cache_context = replace( - context, invalidated=reason - ) # rebind-ok: request-owned retry marker + logging_obj.baseline_cache_context = replace(context, invalidated=reason) logging_obj.baseline_observation = context.capture.model_copy( update=MappingProxyType( - { # rebind-ok: capture uncertainty for failure logging + { "observation": context.capture.observation.model_copy( update=MappingProxyType( { diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index a5b6cabf519..22a17bd4cd8 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -114,9 +114,7 @@ class BatchFileUsage(BaseModel): # each target a different model, so the project's per-model ITPM/OTPM # quota for a row's actual model must be charged with that row's own # tokens -- see `_create_project_io_descriptors_for_models`. - per_model_usage: dict[str, dict[str, int]] = Field( - default_factory=dict - ) # mutable-ok: accumulated incrementally per row while parsing the batch file + per_model_usage: dict[str, dict[str, int]] = Field(default_factory=dict) class _PROXY_BatchRateLimiter(CustomLogger): @@ -465,7 +463,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): body: Final[Mapping[str, object]] = ( MappingProxyType(_BATCH_BODY_ADAPTER.validate_python(raw_body)) if isinstance(raw_body, Mapping) - else MappingProxyType({}) # mutable-ok: immediately frozen empty fallback + else MappingProxyType({}) ) # `max_tokens`/`max_completion_tokens` cap chat completions; `/v1/responses` # rows cap output with `max_output_tokens` instead -- omitting it here diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 33744906b13..17cf7382246 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -3210,7 +3210,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): filtered_content = [ # mutable-ok: token_counter requires list content blocks block for block in content if not (isinstance(block, dict) and block.get("type") == "input_audio") ] - sanitized.append( # mutable-ok: token_counter requires mutable message dicts + sanitized.append( {**message, "content": filtered_content} # mutable-ok: token_counter requires message dicts ) return sanitized @@ -3572,7 +3572,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): try: await asyncio.shield(cleanup) except asyncio.CancelledError as exc: - cancellation = exc # rebind-ok: retain the latest cancellation without interrupting slot release + cancellation = exc cleanup.result() if cancellation is not None: raise cancellation diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 3bf6fea62f7..8fc5faee2c9 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -137,7 +137,7 @@ def add_otel_trace_id_to_request( return data["litellm_trace_id"] = trace_id # rebind-ok: data is an out-param if isinstance(metadata, dict): - metadata["trace_id"] = trace_id # rebind-ok: metadata is the request's own out-param dict + metadata["trace_id"] = trace_id def _session_id_from_baggage(baggage: str) -> str | None: @@ -3142,11 +3142,7 @@ async def move_guardrails_to_metadata( - Moves include_guardrail_response into request metadata before provider dispatch """ if "include_guardrail_response" in data: - data[_metadata_variable_name][ - "include_guardrail_response" - ] = ( # rebind-ok: pre-call hooks mutate the shared request dict in place - data.pop("include_guardrail_response") is True - ) + data[_metadata_variable_name]["include_guardrail_response"] = data.pop("include_guardrail_response") is True # Early-out: skip all guardrails processing when nothing is configured key_metadata: Final = user_api_key_dict.metadata diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index 2ae8639fe61..9708161397a 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -1448,7 +1448,7 @@ def _target_labels( """Display labels by (target_type, target_id): a key's (alias, masked name), a team's (alias, None), a user's (email, None).""" return MappingProxyType( - { # mutable-ok: MappingProxyType needs a dict to wrap + { key: value for key, value in chain( ((("key", row.token), (row.key_alias, row.key_name)) for row in key_rows), @@ -1548,7 +1548,7 @@ async def _shadow_eval_results( await _query_raw(prisma_client, _ATTEMPT_AGG_BY_LEG_SQL, leg_ids) or () ) verdicts_by_target: Final[Mapping[tuple[str, str], ShadowEvalSlice]] = MappingProxyType( - { # mutable-ok: MappingProxyType needs a dict to wrap + { target_by_leg[slice.group]: slice.model_copy( update={"group": target_by_leg[slice.group][1]} # mutable-ok: pydantic update payload ) @@ -1760,7 +1760,7 @@ async def start_shadow_eval( "id": leg_id, "target_type": target_type, "target_id": target_id, - } # mutable-ok: Prisma payload + } for leg_id, (target_type, target_id) in zip(leg_ids, requested_targets) ] ) diff --git a/litellm/proxy/management_endpoints/config_override_endpoints.py b/litellm/proxy/management_endpoints/config_override_endpoints.py index b095ecc1fe5..9d182d4e259 100644 --- a/litellm/proxy/management_endpoints/config_override_endpoints.py +++ b/litellm/proxy/management_endpoints/config_override_endpoints.py @@ -796,9 +796,7 @@ async def get_cyberark_config( field_schema: Final = _build_field_schema(CyberArkConfig) - db_record: Final = await _config_overrides_table(prisma_client).find_unique( - where={"config_type": "cyberark"} - ) # mutable-ok: prisma where clause + db_record: Final = await _config_overrides_table(prisma_client).find_unique(where={"config_type": "cyberark"}) if db_record is not None and db_record.config_value is not None: config_data: Final = _parse_config_value(db_record.config_value) @@ -860,9 +858,7 @@ async def delete_cyberark_config( deleted = False # rebind-ok: set true once the DB row is removed try: - await _config_overrides_table(prisma_client).delete( - where={"config_type": "cyberark"} - ) # mutable-ok: prisma where clause + await _config_overrides_table(prisma_client).delete(where={"config_type": "cyberark"}) deleted = True # rebind-ok: set true once the DB row is removed except RecordNotFoundError: verbose_proxy_logger.debug("No existing CyberArk config record to delete") diff --git a/litellm/proxy/management_endpoints/management_v1/budgets.py b/litellm/proxy/management_endpoints/management_v1/budgets.py index ea13e4547bd..106dbfaf7b7 100644 --- a/litellm/proxy/management_endpoints/management_v1/budgets.py +++ b/litellm/proxy/management_endpoints/management_v1/budgets.py @@ -115,7 +115,7 @@ def _scope(caller: UserAPIKeyAuth) -> Scope: # budget_duration is deliberately absent from `sortable`: the column holds strings # like "7d" and "30d", so a lexicographic ORDER BY puts "30d" ahead of "7d". BUDGET_FILTERS: Final[Mapping[str, FilterSpec]] = MappingProxyType( - { # mutable-ok: an immutable mapping has no literal form; MappingProxyType freezes this one and it never escapes + { "budget_duration": FilterSpec(type=str, ops=frozenset(("in", "is_null"))), "max_budget": FilterSpec(type=float, ops=frozenset(("gte", "lte", "is_null"))), "created_at": FilterSpec(type=datetime, ops=frozenset(("gte", "lte"))), diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 9ad78876043..c22081ebb3a 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -706,9 +706,7 @@ if MCP_AVAILABLE: if not caller_user_id: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, - detail={ - "error": "User ID not found in token" - }, # mutable-ok: FastAPI HTTPException detail requires a plain dict + detail={"error": "User ID not found in token"}, ) return caller_user_id @@ -1865,9 +1863,7 @@ if MCP_AVAILABLE: classified: Final = tuple(_classify(index, conversion) for index, conversion in enumerate(conversions)) outcomes: Final = tuple( - [ - await _create(entry) if isinstance(entry, ConvertedConnector) else entry for entry in classified - ] # mutable-ok: await is illegal in a generator expression here + [await _create(entry) if isinstance(entry, ConvertedConnector) else entry for entry in classified] ) imported: Final = tuple(entry for entry in outcomes if isinstance(entry, MCPConnectorImportResult)) diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 3292a0141d1..0cf201b3a00 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -582,7 +582,6 @@ async def _users_named_by_member_value( subject: Final = value.strip() email: Final[_CaseInsensitiveMatch] = {"equals": subject, "mode": "insensitive"} rows: Final = await _table(UserRepository(prisma_client)).find_many( - # mutable-ok: the Prisma serializer requires concrete dicts and a concrete list where={"OR": [{"sso_user_id": subject}, {"user_email": email}]}, take=take, ) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index d092fc2fbb7..09c6c05e22b 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -2996,7 +2996,7 @@ async def _update_team_members_list( # extend() consumes the generator as it appends, so a member already added by this # same call is seen by the next _member_already_in_team check - the batch dedupes # against itself exactly as the append-one-at-a-time loop this replaced did. - complete_team_data.members_with_roles.extend( # rebind-ok: this helper's contract is to grow the caller's roster in place + complete_team_data.members_with_roles.extend( m for m in resolved_members if not _member_already_in_team(m, complete_team_data) ) @@ -4137,9 +4137,7 @@ async def reset_team_member_budget_fn( team_default_budget_id: Final = await _existing_team_default_budget_id(team_obj, prisma_client) budget_link: Final = ( - { - "connect": {"budget_id": team_default_budget_id} - } # mutable-ok: prisma client requires a plain dict data= argument + {"connect": {"budget_id": team_default_budget_id}} if team_default_budget_id is not None else {"disconnect": True} # mutable-ok: same prisma data= argument ) diff --git a/litellm/proxy/management_helpers/bulk_user_creation.py b/litellm/proxy/management_helpers/bulk_user_creation.py index 56abe3b6a3f..dd3f4ff1b12 100644 --- a/litellm/proxy/management_helpers/bulk_user_creation.py +++ b/litellm/proxy/management_helpers/bulk_user_creation.py @@ -543,7 +543,7 @@ async def _write_team_roster( already_present: Final = frozenset(member.user_id for member in roster if member.user_id) new_members: Final = tuple(member for member in members if member.user_id not in already_present) budget_ids: Final = tuple( - [ # mutable-ok: budgets are created one at a time on the transaction's single connection + [ await _resolve_member_budget_id( prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, diff --git a/litellm/proxy/openai_files_endpoints/batch_guardrails.py b/litellm/proxy/openai_files_endpoints/batch_guardrails.py index 53d51db2b7f..1db4474fc40 100644 --- a/litellm/proxy/openai_files_endpoints/batch_guardrails.py +++ b/litellm/proxy/openai_files_endpoints/batch_guardrails.py @@ -355,9 +355,7 @@ def build_scan_metadata(request_metadata: Mapping[str, object]) -> Mapping[str, Passing the whole thing through would carry values that cannot be copied, such as the parent OTel span, and would hand every record proxy state it has no business seeing. """ - return MappingProxyType( - {key: value for key, value in request_metadata.items() if key in _SCAN_METADATA_KEYS} - ) # mutable-ok: MappingProxyType freezes the comprehension + return MappingProxyType({key: value for key, value in request_metadata.items() if key in _SCAN_METADATA_KEYS}) async def _scan_record( @@ -546,7 +544,7 @@ def rewrite_batch_input_file(file_source: BinaryIO, result: BatchScanResult) -> """ redacted: Final = MappingProxyType( {change.line_number: change for change in result.changes if isinstance(change, RecordRedacted)} - ) # mutable-ok: MappingProxyType freezes the lookup table + ) dropped: Final = frozenset(change.line_number for change in result.changes if isinstance(change, RecordDropped)) output: Final = tempfile.SpooledTemporaryFile( # noqa: SIM115 # the caller uploads this handle diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 2d071343844..4e60c318f03 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -468,9 +468,7 @@ async def fal_ai_proxy_route( endpoint_func: Final = create_pass_through_route( endpoint=endpoint, target=str(updated_url), - custom_headers={ - "Authorization": f"Key {fal_ai_api_key}" - }, # mutable-ok: pass-through request headers require a mutable mapping + custom_headers={"Authorization": f"Key {fal_ai_api_key}"}, custom_llm_provider="fal_ai", is_streaming_request=False, ) @@ -3801,13 +3799,9 @@ async def gigachat_proxy_route( raw_model: Final = request_body.get("model") model: Final = raw_model if isinstance(raw_model, str) else None if model: - is_router_model = is_passthrough_request_using_router_model( - request_body, llm_router - ) # rebind-ok: conditionally set to True + is_router_model = is_passthrough_request_using_router_model(request_body, llm_router) elif any(word in endpoint for word in ("completions", "embeddings")): - raise HTTPException( - status_code=400, detail={"error": "Model is required in request body"} - ) # mutable-ok: HTTPException detail dict + raise HTTPException(status_code=400, detail={"error": "Model is required in request body"}) # If router model, use dedicated router passthrough handler # This uses the same common processing path as non-router models @@ -3908,9 +3902,7 @@ async def handle_gigachat_passthrough_router_model( is_streaming: Final = request_body.get("stream", False) # pyright: ignore[reportUnknownVariableType] # request_body is dict[Unknown, Unknown] - data: dict[str, Any] = await _read_request_body( - request=request - ) # mutable-ok: mutated in place by proxy pipeline; pyright: ignore[reportExplicitAny] # Any needed for proxy pipeline + data: dict[str, Any] = await _read_request_body(request=request) # Any needed for proxy pipeline if user_api_key_dict is not None: auth_metadata: Final = { metadata_key: value diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py index 48c1ced47ae..2cdeddbea30 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py @@ -447,9 +447,7 @@ class VertexPassthroughLoggingHandler: kwargs["model"] = model # rebind-ok: callback metadata records the resolved model kwargs["custom_llm_provider"] = "vertex_ai" # rebind-ok: callback metadata records the resolved provider - standard_pass_through_response_object: Final[ - StandardPassThroughResponseObject - ] = { # mutable-ok: callback contract requires a concrete response dictionary + standard_pass_through_response_object: Final[StandardPassThroughResponseObject] = { "response": json_response, } return { # mutable-ok: passthrough logging contract requires a concrete result dictionary diff --git a/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py b/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py index 567d8375737..8e1dba928af 100644 --- a/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py +++ b/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py @@ -81,9 +81,7 @@ if TYPE_CHECKING: from litellm.integrations.custom_logger import CustomLogger from litellm.proxy.utils import PrismaClient -_RowT = TypeVar( - "_RowT", bound=ManagedResourceRow -) # rebind-ok: TypeVar declarations must stay bare assignments for pyright +_RowT = TypeVar("_RowT", bound=ManagedResourceRow) # --------------------------------------------------------------------------- # Field map @@ -998,9 +996,7 @@ async def _build_list_where_with_cursor( params: Final = query_params or {} after_id: Final[str | None] = params.get("after") before_id: Final[str | None] = params.get("before") - where: PrismaWhere = dict( - owner_filter - ) # rebind-ok: narrowed with the cursor boundary when a valid cursor row exists + where: PrismaWhere = dict(owner_filter) fetch_order: SortOrder = "desc" # rebind-ok: flipped to asc when paging backwards from a before cursor cursor_id: Final = after_id or before_id diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index c2874ac948f..f985c1d49d1 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -819,9 +819,7 @@ def _resolve_team_callback_wiring( user_api_key_dict=user_api_key_dict, proxy_config=proxy_config ) if callback_settings_obj and callback_settings_obj.callback_vars: - for ( - item - ) in callback_settings_obj.callback_vars.items(): # rebind-ok: dict.items iteration for env-ref validation + for item in callback_settings_obj.callback_vars.items(): validate_no_callback_env_reference(item[0], item[1], source="key/team callback metadata") except Exception: # noqa: BLE001 - a broken logging config must never fail the passthrough request verbose_proxy_logger.exception( diff --git a/litellm/proxy/pass_through_endpoints/streaming_handler.py b/litellm/proxy/pass_through_endpoints/streaming_handler.py index e1f13f2bee0..8f0f87e6e69 100644 --- a/litellm/proxy/pass_through_endpoints/streaming_handler.py +++ b/litellm/proxy/pass_through_endpoints/streaming_handler.py @@ -217,9 +217,7 @@ class PassThroughStreamingHandler: async for chunk in response.aiter_bytes(): raw_bytes.append(chunk) PassThroughStreamingHandler._stamp_first_chunk_if_needed(litellm_logging_obj) - complete_frames, pending = split_complete_sse_frames( - pending + chunk - ) # rebind-ok: SSE frame reassembly buffer across transport chunks + complete_frames, pending = split_complete_sse_frames(pending + chunk) if complete_frames: yield ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection( complete_frames, resolved_model_name, litellm_logging_obj diff --git a/litellm/proxy/policy_engine/pipeline_executor.py b/litellm/proxy/policy_engine/pipeline_executor.py index e9d23436b59..6c05ca0b22c 100644 --- a/litellm/proxy/policy_engine/pipeline_executor.py +++ b/litellm/proxy/policy_engine/pipeline_executor.py @@ -108,7 +108,7 @@ _GuardrailMethodT = TypeVar("_GuardrailMethodT", bound=Callable[..., object]) def _logged_by_inner_guardrail(method: _GuardrailMethodT) -> _GuardrailMethodT: - vars(method)[LOGS_GUARDRAIL_INFORMATION_MARKER] = True # rebind-ok: stamps the method the class body just defined + vars(method)[LOGS_GUARDRAIL_INFORMATION_MARKER] = True return method @@ -278,9 +278,7 @@ def _prepare_hook_input( guardrail loops do this.""" if "metadata" not in data: data["metadata"] = {} # mutable-ok: request metadata bucket, hooks mutate it - data["metadata"]["guardrails"] = [ - step.guardrail - ] # mutable-ok: guardrails list is part of the request-payload shape + data["metadata"]["guardrails"] = [step.guardrail] scans_raw_request: Final = callback.scan_raw_request hook_input: Final[dict] = ( # mutable-ok: same request-payload shape as data @@ -456,7 +454,7 @@ class PipelineExecutor: observer: Final = _StreamRewriteObserver(scanner) deliver_rewrites: Final = type(endpoint_translation).delivers_ended_stream_rewrites originals: Final = copy.deepcopy(streaming_chunks) - hook_input.pop("response", None) # rebind-ok: an earlier step's stored response goes so this step's is stored + hook_input.pop("response", None) try: if deliver_rewrites: await endpoint_translation.process_output_streaming_response( @@ -582,7 +580,7 @@ class PipelineExecutor: {"response": response}, None, None, - ) # mutable-ok: modified-data contract is a plain dict + ) return ("pass", response if isinstance(response, dict) else None, None, None) except Exception as e: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 9a42ee75c51..26a423d7162 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -5261,9 +5261,7 @@ class ProxyConfig: return with open(f"{user_config_file_path}", "w") as config_file: - yaml.dump( - dict(new_config), config_file, default_flow_style=False - ) # mutable-ok: YAML must serialize a plain dict + yaml.dump(dict(new_config), config_file, default_flow_style=False) async def _save_changed_config_section( self, @@ -10137,7 +10135,7 @@ class ProxyStartupEvent: str(identity): str(fingerprint) for identity, fingerprint in (decoded.items() if isinstance(decoded, Mapping) else ()) } - ) # mutable-ok: MappingProxyType owns the completed immutable baseline + ) snapshot: Final = snapshot_tuning_baselines(deployments) try: await config_table.create( @@ -10163,7 +10161,7 @@ class ProxyStartupEvent: competing_decoded.items() if isinstance(competing_decoded, Mapping) else () ) } - ) # mutable-ok: MappingProxyType owns the completed immutable baseline + ) except Exception as e: # noqa: BLE001 # enforcement is skipped for this boot; refusing every tuned router on a DB blip is the one outcome the gate forbids verbose_proxy_logger.warning("Heuristic-v1 tuning baseline unavailable, gate not enforced this boot: %s", e) return None @@ -10199,7 +10197,7 @@ class ProxyStartupEvent: proxy_logging_obj: ProxyLogging, ) -> ProxyWorkerHeartbeat: """Initializes scheduled background jobs""" - global heuristic_v1_tuning_baselines, store_model_in_db, scheduler, scheduler_executor # rebind-ok: startup publishes the one read-only baseline snapshot + global heuristic_v1_tuning_baselines, store_model_in_db, scheduler, scheduler_executor # MEMORY LEAK FIX: Configure scheduler with optimized settings # Memray analysis showed APScheduler's normalize() and _apply_jitter() causing diff --git a/litellm/proxy/rag_endpoints/endpoints.py b/litellm/proxy/rag_endpoints/endpoints.py index 4f0c9f42421..974bff6338a 100644 --- a/litellm/proxy/rag_endpoints/endpoints.py +++ b/litellm/proxy/rag_endpoints/endpoints.py @@ -824,7 +824,7 @@ async def rag_query( merged_retrieval_config: Final = { **retrieval_config, **store_data, - } # mutable-ok: litellm.aquery requires a plain dict payload + } # Add litellm data request_data: dict[str, object] = {} diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 36b7a3a4a8a..69c7f0a09ed 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -97,11 +97,7 @@ def _normalize_tool_dialect( tools: Final = data.get("tools") tool_choice: Final = data.get("tool_choice") normalized_tools: Final = ( - [ - _convert_tool_envelope(tool, to_chat=to_chat) for tool in tools - ] # mutable-ok: body's tools stays a plain JSON list - if isinstance(tools, list) - else tools + [_convert_tool_envelope(tool, to_chat=to_chat) for tool in tools] if isinstance(tools, list) else tools ) normalized_choice: Final = _convert_tool_envelope(tool_choice, to_chat=to_chat) if normalized_tools == tools and normalized_choice == tool_choice: diff --git a/litellm/proxy/spend_tracking/carried_budget_state.py b/litellm/proxy/spend_tracking/carried_budget_state.py index da8bf60ebda..0dfb38272b0 100644 --- a/litellm/proxy/spend_tracking/carried_budget_state.py +++ b/litellm/proxy/spend_tracking/carried_budget_state.py @@ -36,9 +36,7 @@ def carry_team_and_user_budget_state( def carry_organization_budget_state(valid_token: UserAPIKeyAuth, org_table: LiteLLM_OrganizationTable) -> None: budget_table: Final = org_table.litellm_budget_table - valid_token.organization_alias = ( - org_table.organization_alias - ) # rebind-ok: the request credential is pinned in place + valid_token.organization_alias = org_table.organization_alias valid_token.org_budget_snapshot = OrgBudgetSnapshot( # rebind-ok: same object the caller keeps using spend=org_table.spend, max_budget=budget_table.max_budget if budget_table is not None else None, diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index e25b3bed757..78d6c25a336 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -602,15 +602,11 @@ def _partition_post_call_callbacks() -> tuple[tuple[CustomGuardrail, ...], tuple return (guardrails, others) -def _merge_pipeline_metadata_bucket( - data: dict, bucket_key: str, modified_bucket_value: object -) -> None: # mutable-ok: request payload dict, written in place +def _merge_pipeline_metadata_bucket(data: dict, bucket_key: str, modified_bucket_value: object) -> None: if not isinstance(modified_bucket_value, dict): return modified_bucket: Final = cast("dict[str, object]", modified_bucket_value) # cast-ok: metadata buckets are str-keyed - surviving_writes: Final = { - key: value for key, value in modified_bucket.items() if key != "guardrails" - } # mutable-ok: merged into the live request metadata bucket in place + surviving_writes: Final = {key: value for key, value in modified_bucket.items() if key != "guardrails"} existing_bucket: Final = data.get(bucket_key) if isinstance(existing_bucket, dict): cast("dict[str, object]", existing_bucket).update(surviving_writes) # cast-ok: metadata buckets are str-keyed @@ -618,9 +614,7 @@ def _merge_pipeline_metadata_bucket( data[bucket_key] = surviving_writes -def _merge_pipeline_metadata_writes( - data: dict, modified_data: Mapping[str, object] -) -> None: # mutable-ok: request payload dict, written in place +def _merge_pipeline_metadata_writes(data: dict, modified_data: Mapping[str, object]) -> None: """ Copy metadata-bucket writes from a pipeline's working copy back onto the request. @@ -1052,7 +1046,6 @@ def _deployment_attribution_for_model_group(model_group: object, team_id: str | ) return MappingProxyType( { - # mutable-ok: frozen immediately by the outer MappingProxyType **({"custom_llm_provider": shared_provider} if shared_provider is not None else {}), **( { # mutable-ok: frozen immediately by the outer MappingProxyType @@ -1976,9 +1969,7 @@ class ProxyLogging: """ scans_raw_request: Final = callback.scan_raw_request should_use_raw_snapshot: Final = scans_raw_request and raw_request_snapshot is not None - input_data: Final = ( # mutable-ok: same request-payload shape as data - independent_snapshot(raw_request_snapshot) if should_use_raw_snapshot else data - ) + input_data: Final = independent_snapshot(raw_request_snapshot) if should_use_raw_snapshot else data # _process_guardrail_callback always calls mark_pre_call_hook_ran on a # successful run, which unconditionally stamps bookkeeping metadata onto # the dict regardless of whether the guardrail's own hook mutated @@ -2169,9 +2160,7 @@ class ProxyLogging: if pipeline.mode != event_hook: continue - step_input: dict = ( - {**data, "response": current_response} if current_response is not None else data - ) # mutable-ok: same request-payload shape as data + step_input: dict = {**data, "response": current_response} if current_response is not None else data result: PipelineExecutionResult = await PipelineExecutor.execute_steps( steps=pipeline.steps, diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py index 37ca989b8d3..18ecc250b64 100644 --- a/litellm/rerank_api/main.py +++ b/litellm/rerank_api/main.py @@ -44,9 +44,7 @@ async def arerank( """ Async: Reranks a list of documents based on their relevance to the query """ - _custom_llm_provider: str | None = ( - None # rebind-ok: set by the declared-provider guard or the get_llm_provider unpack; read in the except - ) + _custom_llm_provider: str | None = None try: loop: Final = asyncio.get_event_loop() kwargs["arerank"] = True diff --git a/litellm/responses/additional_tools.py b/litellm/responses/additional_tools.py index ea0d7af350c..5239bd395cc 100644 --- a/litellm/responses/additional_tools.py +++ b/litellm/responses/additional_tools.py @@ -37,12 +37,7 @@ def _tools_of_item(item: object) -> tuple[ALL_RESPONSES_API_TOOL_PARAMS, ...]: parsed: Final = _AdditionalToolsItem.model_validate(item) except ValidationError: return () - return tuple( - cast( - "ALL_RESPONSES_API_TOOL_PARAMS", tool - ) # cast-ok: nested tools carry the same raw tool JSON as top-level tools - for tool in parsed.tools - ) + return tuple(cast("ALL_RESPONSES_API_TOOL_PARAMS", tool) for tool in parsed.tools) def hoist_additional_tools( diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index 3ca2cc28c9a..e421cae0724 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -866,14 +866,14 @@ class LiteLLMCompletionResponsesConfig: elif pending: # Not followed by an assistant message — keep the reasoning # standalone instead of dropping it. - merged.extend( # mutable-ok: append reasoning messages + merged.extend( [_standalone(text, blocks) for text, blocks in pending] # mutable-ok: append reasoning messages ) pending = [] # mutable-ok: reset accumulator merged.append(msg) - merged.extend( # mutable-ok: append trailing reasoning + merged.extend( [_standalone(text, blocks) for text, blocks in pending] # mutable-ok: append trailing reasoning ) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 59655800af6..64989c4cf1c 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -170,7 +170,7 @@ def _log_background_task_failure(task: asyncio.Task[object], *, task_name: str) _ERROR_CODE_HTTP_STATUS: Final[Mapping[str, int]] = MappingProxyType( - { # mutable-ok: immediately frozen by MappingProxyType + { "server_error": 500, "rate_limit_exceeded": 429, "insufficient_quota": 429, @@ -1633,9 +1633,7 @@ def _extract_frame_quota_estimate_inputs(msg_obj: Mapping[str, object]) -> tuple params: Final[Mapping[str, object]] = ( nested if _is_json_object(nested) and nested - else MappingProxyType( # mutable-ok: immediately frozen filtered frame - {k: v for k, v in msg_obj.items() if k != "type"} - ) + else MappingProxyType({k: v for k, v in msg_obj.items() if k != "type"}) ) text_parts: Final[list[str]] = [] # mutable-ok: local accumulator built in one pass, not shared pending: Final[list[object]] = [ # mutable-ok: explicit worklist avoids recursion @@ -2297,7 +2295,7 @@ class ResponsesWebSocketStreaming: except RateLimitError as e: try: await self.websocket.send_text( - json.dumps( # mutable-ok: WebSocket wire payload requires JSON objects + json.dumps( { # mutable-ok: WebSocket wire payload requires JSON objects "type": "error", "error": { # mutable-ok: nested WebSocket error object @@ -2743,9 +2741,7 @@ class ManagedResponsesWebSocketHandler: directly (before serialization) to avoid a redundant JSON round-trip on every chunk. Returns the completed event dict, or ``None``. """ - completed_event: _MutableJsonObject | None = ( - None # rebind-ok: captures the completed event once the stream yields it - ) + completed_event: _MutableJsonObject | None = None stream_response: Final = await litellm.aresponses(model=model, **call_kwargs) async for chunk in stream_response: if chunk is None: diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py index a2642795cea..9b0d259eb8a 100644 --- a/litellm/responses/utils.py +++ b/litellm/responses/utils.py @@ -566,7 +566,7 @@ class ResponsesAPIRequestUtils: return items: Final = cast(list[object], request_input) # cast-ok: untyped client json stripped: Final = tuple(ResponsesAPIRequestUtils._without_encrypted_reasoning(item) for item in items) - items[:] = (item for item in stripped if item is not None) # rebind-ok: list shared with fallback snapshot + items[:] = (item for item in stripped if item is not None) @staticmethod def _without_encrypted_reasoning(item: object) -> object | None: diff --git a/litellm/router.py b/litellm/router.py index 7267f6eb3ba..62042e1c969 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -675,7 +675,7 @@ class RoutingArgs(enum.Enum): # entries their deployments own. Weak so a router nothing references any more, such # as the per-request one built from a caller-supplied user_config, drops out on its # own rather than leaving entries behind that nothing can withdraw. -_live_routers: Final["weakref.WeakSet[Router]"] = weakref.WeakSet() # mutable-ok: identity set of live routers +_live_routers: Final["weakref.WeakSet[Router]"] = weakref.WeakSet() def _replay_live_router_model_cost() -> None: @@ -2954,10 +2954,10 @@ class Router: fallback_headers_are_settled = False async for fallback_item in fallback_response: if not fallback_headers_are_settled: - fallback_headers_are_settled = True # rebind-ok: one-shot latch + fallback_headers_are_settled = True # a fallback that failed over again only repoints itself once it yields - prepared_fallback_hidden_params = ( # rebind-ok: re-read once the fallback yields - Router._adopt_fallback_response_headers(wrapper_ref, fallback_response) + prepared_fallback_hidden_params = Router._adopt_fallback_response_headers( + wrapper_ref, fallback_response ) Router._apply_fallback_hidden_params_to_item(fallback_item, prepared_fallback_hidden_params) if ( @@ -3513,10 +3513,10 @@ class Router: fallback_headers_are_settled = False for fallback_item in fallback_response: if not fallback_headers_are_settled: - fallback_headers_are_settled = True # rebind-ok: one-shot latch + fallback_headers_are_settled = True # a fallback that failed over again only repoints itself once it yields - prepared_fallback_hidden_params = ( # rebind-ok: re-read once the fallback yields - Router._adopt_fallback_response_headers(wrapper_ref, fallback_response) + prepared_fallback_hidden_params = Router._adopt_fallback_response_headers( + wrapper_ref, fallback_response ) Router._apply_fallback_hidden_params_to_item(fallback_item, prepared_fallback_hidden_params) if ( @@ -5459,23 +5459,23 @@ class Router: if _anthropic_stream_should_drop_pre_content_ping(chunk, has_generated_content): continue if _anthropic_stream_commits_now(chunk, has_generated_content, len(buffered_lifecycle_chunks)): - has_generated_content = True # rebind-ok: real content seen, or the buffer cap was hit + has_generated_content = True # A transport can split one SSE data line across byte chunks, so pre-content # detection parses the accumulated buffer plus the current chunk, never the # chunk alone; the buffer is already capped, which bounds this window too. - parse_window = ( # rebind-ok: freshly computed each iteration, never carried over + parse_window = ( b"".join(c for c in (*buffered_lifecycle_chunks, chunk) if isinstance(c, (bytes, bytearray))) # pyright: ignore[reportUnnecessaryIsInstance] # bridge-path chunks are not always bytes at runtime if not has_generated_content and isinstance(chunk, (bytes, bytearray)) # pyright: ignore[reportUnnecessaryIsInstance] # bridge-path chunks are not always bytes at runtime else chunk ) error_event = parse_anthropic_error_event(parse_window) - retriable_pending_error = ( # rebind-ok: freshly computed each iteration, never carried over + retriable_pending_error = ( not has_generated_content and error_event is not None and _is_retriable_anthropic_status(error_event[2]) and not _anthropic_stream_error_is_gateway_verdict(chunk) ) - refusal_stop_details = ( # rebind-ok: freshly computed each iteration, never carried over + refusal_stop_details = ( parse_anthropic_refusal_stop_details(parse_window) if not has_generated_content and error_event is None else None @@ -5493,7 +5493,7 @@ class Router: buffered_lifecycle_chunks = (*buffered_lifecycle_chunks, chunk) continue if retriable_pending_error: - assert error_event is not None # guard-ok: retriable_pending_error implies this + assert error_event is not None _error_type, message, status_code = error_event raise MidStreamFallbackError( message=message, @@ -10951,10 +10951,8 @@ class Router: model_group_info.supports_fast_mode = model_group_info.supports_fast_mode and ( AnthropicModelInfo.supports_fast_mode(litellm_model, llm_provider) ) - deployment_reasoning_efforts = ( - resolve_supported_reasoning_efforts( # rebind-ok: recalculated per deployment - model_info, deployment_is_mapped=deployment_is_mapped - ) + deployment_reasoning_efforts = resolve_supported_reasoning_efforts( + model_info, deployment_is_mapped=deployment_is_mapped ) if deployment_reasoning_efforts is None: reasoning_efforts_unknown = True diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 1f4285a7960..0f252952a9d 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -1093,7 +1093,7 @@ def _with_classifier_forecast( if forecast is None: return decision verdict: Final = forecast.verdict - enriched: Final[StandardLoggingRoutingDecision] = { # mutable-ok: routing decisions are JSON TypedDict records + enriched: Final[StandardLoggingRoutingDecision] = { **decision, "classifier_crux": verdict.crux, "classifier_primary_rule": verdict.primary_rule, @@ -2484,7 +2484,7 @@ class ComplexityRouter(CustomLogger): {"role": "user", "content": opening_task}, # mutable-ok: SDK messages are dict-shaped ] if latest_follow_up is not None: - task_messages.append( # mutable-ok: the provider SDK requires a concrete message list + task_messages.append( {"role": "user", "content": latest_follow_up} # mutable-ok: SDK messages are dict-shaped ) diff --git a/litellm/router_strategy/tag_based_routing.py b/litellm/router_strategy/tag_based_routing.py index d4f46e94579..50dce250920 100644 --- a/litellm/router_strategy/tag_based_routing.py +++ b/litellm/router_strategy/tag_based_routing.py @@ -217,9 +217,7 @@ def _strip_routing_prefix(tags: Sequence[str], prefix: str) -> tuple[tuple[str, def _split_tags(tags: Sequence[str]) -> tuple[tuple[str, ...], list[str], tuple[str, ...]]: required: Final = tuple(tag[1:] for tag in tags if tag.startswith("&") and len(tag) > 1) - positive: Final = [ - t for t in tags if not t.startswith("!") and not t.startswith("&") - ] # mutable-ok: feeds _match_deployment's existing list[str]-typed request_tags param + positive: Final = [t for t in tags if not t.startswith("!") and not t.startswith("&")] excluded: Final = tuple(tag[1:] for tag in tags if tag.startswith("!") and len(tag) > 1) return required, positive, excluded diff --git a/litellm/router_utils/auto_router_tuning_baseline.py b/litellm/router_utils/auto_router_tuning_baseline.py index e87548bf6de..4707a51dfb8 100644 --- a/litellm/router_utils/auto_router_tuning_baseline.py +++ b/litellm/router_utils/auto_router_tuning_baseline.py @@ -120,7 +120,7 @@ def snapshot_tuning_baselines(deployments: Iterable[Mapping[str, object]]) -> Ma if (pair := heuristic_v1_router_fingerprint(deployment)) is not None for identity, fingerprint in (pair,) } - ) # mutable-ok: MappingProxyType owns the completed immutable snapshot + ) def is_mutable_tuned_candidate(candidate: Mapping[str, object], baselines: Mapping[str, str]) -> bool: diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index 4745e4094e6..60585b3cc38 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -655,7 +655,7 @@ async def run_async_fallback( # LOGGING kwargs = litellm_router.log_retry(kwargs=kwargs, e=original_exception) verbose_router_logger.info("Falling back to model_group = %s", mask_sensitive_structure(mg)) - kwargs.pop("_target_order", None) # rebind-ok: next hop must not inherit the previous order target + kwargs.pop("_target_order", None) if isinstance(mg, str): kwargs["model"] = mg elif isinstance(mg, dict): diff --git a/litellm/rust_bridge/lifecycle.py b/litellm/rust_bridge/lifecycle.py index 4096d386964..b3b7a1888c3 100644 --- a/litellm/rust_bridge/lifecycle.py +++ b/litellm/rust_bridge/lifecycle.py @@ -46,7 +46,7 @@ class StreamClosed(Exception): async def _settle(execution: Execution, step: Step) -> Settled: while isinstance(step, Await): try: - value = await step.awaitable # rebind-ok: each selected await produces the next protocol input + value = await step.awaitable except GeneratorExit: raise except BaseException as error: diff --git a/litellm/rust_bridge/logger.py b/litellm/rust_bridge/logger.py index bcd53852a2f..544e7195848 100644 --- a/litellm/rust_bridge/logger.py +++ b/litellm/rust_bridge/logger.py @@ -53,7 +53,7 @@ def emit( extra={ "rust_target": target, "rust_fields": dict(fields), - }, # mutable-ok: LogRecord requires JSON dict extras + }, ) _REDACTION.filter(record) _CORRELATION.filter(record) diff --git a/litellm/types/management_endpoints/auto_router_endpoints.py b/litellm/types/management_endpoints/auto_router_endpoints.py index 334ca0dfb08..e191470ec6e 100644 --- a/litellm/types/management_endpoints/auto_router_endpoints.py +++ b/litellm/types/management_endpoints/auto_router_endpoints.py @@ -156,7 +156,7 @@ class AutoRouterRoutingTestRequest(BaseModel): the serving path. """ return MappingProxyType( - { # mutable-ok: MappingProxyType needs a dict to wrap + { key: value for key, value in (("messages", self.messages), ("system", self.system), ("tools", self.tools)) if value is not None diff --git a/litellm/types/passthrough_endpoints/managed_id_rewriter.py b/litellm/types/passthrough_endpoints/managed_id_rewriter.py index 675aae96a5b..33749cc2ab8 100644 --- a/litellm/types/passthrough_endpoints/managed_id_rewriter.py +++ b/litellm/types/passthrough_endpoints/managed_id_rewriter.py @@ -55,9 +55,7 @@ class ManagedObjectRow(ManagedResourceRow, Protocol): unified_object_id: str -RowT = TypeVar( - "RowT", bound=ManagedResourceRow -) # rebind-ok: TypeVar declarations must stay bare assignments for pyright +RowT = TypeVar("RowT", bound=ManagedResourceRow) class ManagedTable(Protocol[RowT]): diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 3f8471cbde1..064e3040054 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -1354,9 +1354,7 @@ def add_provider_specific_fields(object: BaseModel, provider_specific_fields: di class Message(SafeAttributeModel, OpenAIObject): content: str | None role: Literal["assistant", "user", "system", "tool", "function"] - tool_calls: ( - list[ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall] | None - ) # mutable-ok: public pydantic response field; only the union member is new + tool_calls: list[ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall] | None function_call: FunctionCall | None audio: ChatCompletionAudioResponse | None = None images: list[ImageURLListItem] | None = None @@ -1479,9 +1477,7 @@ class Delta(SafeAttributeModel, OpenAIObject): content: str | None role: str | None function_call: FunctionCall | None - tool_calls: ( - list[ChatCompletionDeltaToolCall | ChatCompletionDeltaCustomToolCall] | None - ) # mutable-ok: public pydantic response field; only the union member is new + tool_calls: list[ChatCompletionDeltaToolCall | ChatCompletionDeltaCustomToolCall] | None audio: ChatCompletionAudioResponse | None images: list[ImageURLListItem] | None annotations: list[ChatCompletionAnnotation] | None diff --git a/litellm/utils.py b/litellm/utils.py index f3b9070ecff..64097021dff 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1927,9 +1927,7 @@ def client(original_function): is_completion_with_fallbacks: Final = kwargs.get("fallbacks") is not None kwargs.pop("_is_litellm_internal_call", None) # discard if injected _is_litellm_internal_call: Final = is_internal_call.get() - _deployment_call_end_time: datetime.datetime | None = ( - None # rebind-ok: set once, from inside the except below, only if the model call itself fails - ) + _deployment_call_end_time: datetime.datetime | None = None try: if logging_obj is None: @@ -2743,9 +2741,7 @@ def _supports_factory(model: str, custom_llm_provider: str | None, key: str) -> try: declared: Final = declared_authenticating_provider(model, custom_llm_provider) if declared is not None: - model = model.removeprefix( - f"{declared}/" - ) # rebind-ok: mirrors get_llm_provider's split without its OAuth flow + model = model.removeprefix(f"{declared}/") custom_llm_provider = declared # rebind-ok: same else: model, custom_llm_provider, _, _ = litellm.get_llm_provider( @@ -2846,9 +2842,7 @@ def is_explicitly_disabled_factory(model: str, custom_llm_provider: str | None, try: declared: Final = declared_authenticating_provider(model, custom_llm_provider) if declared is not None: - model = model.removeprefix( - f"{declared}/" - ) # rebind-ok: mirrors get_llm_provider's split without its OAuth flow + model = model.removeprefix(f"{declared}/") custom_llm_provider = declared # rebind-ok: same else: model, custom_llm_provider, _, _ = litellm.get_llm_provider( diff --git a/scripts/check_type_discipline.py b/scripts/check_type_discipline.py index a2ab4760c4f..2bb65072ad4 100644 --- a/scripts/check_type_discipline.py +++ b/scripts/check_type_discipline.py @@ -1,6 +1,6 @@ #!/usr/bin/env python3 """Type-discipline checker: the rules ruff can't enforce. - + Rules ----- LIT001 Mutable collection in a type annotation, anywhere it appears: function @@ -99,19 +99,23 @@ LIT012 TypedDict field without a `ReadOnly[...]` qualifier. A writable key lets the functional form (`X = TypedDict("X", {...})`) is checked too. A base imported from another module is out of reach without import resolution. Suppress with `# writable-ok: `. +LIT013 A `# -ok: ` suppression on a line where none of the rules + that token suppresses fires. Like ruff's RUF100: a marker that suppresses + nothing rots in place and hides real violations that land on the line + later. Delete it. LIT000 Setup failure: a target file could not be read, or contains a syntax error. Reported as a violation rather than crashing the run. - + Usage ----- python check_type_discipline.py litellm/ tests/ Exit code 1 if any violation is found. Stdlib only. """ - + from __future__ import annotations - + import ast import io import os @@ -122,28 +126,50 @@ from dataclasses import dataclass from multiprocessing import Pool from pathlib import Path from collections.abc import Iterable, Iterator, Mapping, Sequence +from types import MappingProxyType from typing import NamedTuple - + # Mutable collection types, banned in *every* annotation. Name-based, so `dict`, # `typing.Dict`, `collections.deque`, and `collections.abc.MutableMapping` all match # however they were imported. The read-only interfaces (Mapping, Sequence, the # immutable AbstractSet / `abc.Set`, Collection) and the immutable concretes (tuple, # frozenset) are the escape hatch and are deliberately absent -- as is the bare name # `Set`, which collides with the read-only `collections.abc.Set`. -MUTABLE_COLLECTIONS = frozenset(( - "dict", "list", "set", - "Dict", "List", "DefaultDict", "OrderedDict", "Counter", "Deque", "ChainMap", - "deque", "defaultdict", - "MutableMapping", "MutableSequence", "MutableSet", -)) +MUTABLE_COLLECTIONS = frozenset( + ( + "dict", + "list", + "set", + "Dict", + "List", + "DefaultDict", + "OrderedDict", + "Counter", + "Deque", + "ChainMap", + "deque", + "defaultdict", + "MutableMapping", + "MutableSequence", + "MutableSet", + ) +) # Callables whose result is a fresh *mutable* collection (LIT002). `tuple` and # `frozenset` are deliberately absent -- they are the wrappers you reach for, and # a generator expression fed to them is the blessed one-shot build. -MUTABLE_CONSTRUCTORS = frozenset(( - "dict", "list", "set", - "deque", "defaultdict", "OrderedDict", "Counter", "ChainMap", -)) +MUTABLE_CONSTRUCTORS = frozenset( + ( + "dict", + "list", + "set", + "deque", + "defaultdict", + "OrderedDict", + "Counter", + "ChainMap", + ) +) # A *qualified* call (`x.deque()`) counts as construction only for names that are rarely # method names; `dict`/`list`/`set` are dropped here because `.dict()` / `.set()` / `.list()` # are common methods (e.g. pydantic's `model.dict()`), not collection construction. A @@ -165,7 +191,7 @@ READONLY_QUALIFIER = "ReadOnly" FIELD_QUALIFIER_WRAPPERS = frozenset(("Required", "NotRequired", "Annotated")) TYPEDDICT_BASE = "TypedDict" MIN_REASON_LEN = 3 - + NOQA_RE = re.compile( r"#\s*noqa" r"(?P:\s*(?P[A-Z]+[0-9]+(?:\s*,\s*[A-Z]+[0-9]+)*))?" @@ -173,9 +199,7 @@ NOQA_RE = re.compile( re.IGNORECASE, ) TYPE_IGNORE_RE = re.compile(r"#\s*type:\s*ignore\b") -IGNORE_RE = re.compile( - r"#\s*(?:pyright|mypy):\s*ignore(?P\[[^\]]*\])?(?P.*)" -) +IGNORE_RE = re.compile(r"#\s*(?:pyright|mypy):\s*ignore(?P\[[^\]]*\])?(?P.*)") MUTABLE_OK_RE = re.compile(r"#\s*mutable-ok(?::\s*(?P.*))?") CAST_OK_RE = re.compile(r"#\s*cast-ok(?::\s*(?P.*))?") GUARD_OK_RE = re.compile(r"#\s*guard-ok(?::\s*(?P.*))?") @@ -183,48 +207,45 @@ KWARGS_OK_RE = re.compile(r"#\s*kwargs-ok(?::\s*(?P.*))?") REBIND_OK_RE = re.compile(r"#\s*rebind-ok(?::\s*(?P.*))?") WRITABLE_OK_RE = re.compile(r"#\s*writable-ok(?::\s*(?P.*))?") +@dataclass(frozen=True, slots=True) +class _OkToken: + """One `*-ok` suppression token: its comment pattern and the rule codes it suppresses.""" + + token: str + pattern: re.Pattern[str] + codes: frozenset[str] + + # Suppression tokens that must each carry a reason (LIT005). -OK_SUPPRESSIONS: tuple[tuple[str, re.Pattern[str]], ...] = ( - ("mutable-ok", MUTABLE_OK_RE), - ("cast-ok", CAST_OK_RE), - ("guard-ok", GUARD_OK_RE), - ("kwargs-ok", KWARGS_OK_RE), - ("rebind-ok", REBIND_OK_RE), - ("writable-ok", WRITABLE_OK_RE), +OK_SUPPRESSIONS: Final[tuple[_OkToken, ...]] = ( + _OkToken("mutable-ok", MUTABLE_OK_RE, frozenset(("LIT001", "LIT002"))), + _OkToken("cast-ok", CAST_OK_RE, frozenset(("LIT006",))), + _OkToken("guard-ok", GUARD_OK_RE, frozenset(("LIT007",))), + _OkToken("kwargs-ok", KWARGS_OK_RE, frozenset(("LIT008",))), + _OkToken("rebind-ok", REBIND_OK_RE, frozenset(("LIT010", "LIT011"))), + _OkToken("writable-ok", WRITABLE_OK_RE, frozenset(("LIT012",))), ) - - + + class Violation(NamedTuple): path: Path line: int code: str message: str - + def render(self) -> str: return f"{self.path}:{self.line}: {self.code} {self.message}" - - -@dataclass(frozen=True, slots=True) -class Comments: - """The lines carrying each valid `*-ok` suppression.""" - mutable_ok_lines: frozenset[int] - cast_ok_lines: frozenset[int] - guard_ok_lines: frozenset[int] - kwargs_ok_lines: frozenset[int] - rebind_ok_lines: frozenset[int] - writable_ok_lines: frozenset[int] - - + # --------------------------------------------------------------------------- # # Comment scanning (LIT003 / LIT004 / LIT005) # --------------------------------------------------------------------------- # - - + + def _reason_of(rest: str) -> str: return rest.strip().lstrip("#-").strip() - + def _valid_ok(regex: re.Pattern[str], text: str) -> bool: """True iff `text` carries this suppression with a reason of usable length.""" m = regex.search(text) @@ -233,35 +254,40 @@ def _valid_ok(regex: re.Pattern[str], text: str) -> bool: def _comment_violations(path: Path, line_no: int, text: str) -> Iterator[Violation]: """Pure: all LIT003/004/005 findings for one comment.""" - for token, regex in OK_SUPPRESSIONS: - m = regex.search(text) + for ok in OK_SUPPRESSIONS: + m = ok.pattern.search(text) if m and len((m.group("reason") or "").strip()) < MIN_REASON_LEN: - yield Violation(path, line_no, "LIT005", f"{token} requires a reason: `# {token}: `") - + yield Violation(path, line_no, "LIT005", f"{ok.token} requires a reason: `# {ok.token}: `") + m = NOQA_RE.search(text) if m: if not m.group("codes"): yield Violation(path, line_no, "LIT003", "noqa requires rule codes: `# noqa: XXX123 # `") elif len(_reason_of(m.group("rest"))) < MIN_REASON_LEN: yield Violation(path, line_no, "LIT003", "noqa requires a reason: `# noqa: XXX123 # `") - + if TYPE_IGNORE_RE.search(text): - yield Violation(path, line_no, "LIT009", - "`# type: ignore` is inert (enableTypeIgnoreComments is false, so " - "basedpyright never honors it); use `# pyright: ignore[ruleName] # `") + yield Violation( + path, + line_no, + "LIT009", + "`# type: ignore` is inert (enableTypeIgnoreComments is false, so " + "basedpyright never honors it); use `# pyright: ignore[ruleName] # `", + ) m = IGNORE_RE.search(text) if m: codes = m.group("codes") if not codes or codes == "[]": - yield Violation(path, line_no, "LIT004", - "ignore requires codes: `# pyright: ignore[ruleName] # `") + yield Violation(path, line_no, "LIT004", "ignore requires codes: `# pyright: ignore[ruleName] # `") elif len(_reason_of(m.group("rest"))) < MIN_REASON_LEN: - yield Violation(path, line_no, "LIT004", - "ignore requires a reason: `# pyright: ignore[ruleName] # `") - - -def scan_comments(path: Path, source: str) -> tuple[Comments, tuple[Violation, ...]]: + yield Violation( + path, line_no, "LIT004", "ignore requires a reason: `# pyright: ignore[ruleName] # `" + ) + + +def scan_comments(path: Path, source: str) -> tuple[Mapping[str, frozenset[int]], tuple[Violation, ...]]: + """Tokenize comments into (token -> lines with a valid reasoned marker, comment violations).""" try: tokens = tokenize.generate_tokens(io.StringIO(source).readline) comment_toks = tuple((t.start[0], t.string) for t in tokens if t.type == tokenize.COMMENT) @@ -269,27 +295,22 @@ def scan_comments(path: Path, source: str) -> tuple[Comments, tuple[Violation, . # tokenize raises TokenError (EOF mid-construct) or a SyntaxError subclass # (IndentationError / TabError) on malformed source; defer to ast.parse below, # which re-raises and is reported as LIT000 rather than crashing the run. - return Comments(frozenset(), frozenset(), frozenset(), frozenset(), frozenset(), frozenset()), () - - def _lines_with(regex: re.Pattern[str]) -> frozenset[int]: - return frozenset(line for line, text in comment_toks if _valid_ok(regex, text)) + return {ok.token: frozenset() for ok in OK_SUPPRESSIONS}, () return ( - Comments( - mutable_ok_lines=_lines_with(MUTABLE_OK_RE), - cast_ok_lines=_lines_with(CAST_OK_RE), - guard_ok_lines=_lines_with(GUARD_OK_RE), - kwargs_ok_lines=_lines_with(KWARGS_OK_RE), - rebind_ok_lines=_lines_with(REBIND_OK_RE), - writable_ok_lines=_lines_with(WRITABLE_OK_RE), + MappingProxyType( + { + ok.token: frozenset(line for line, text in comment_toks if _valid_ok(ok.pattern, text)) + for ok in OK_SUPPRESSIONS + } ), tuple(v for line, text in comment_toks for v in _comment_violations(path, line, text)), ) - - + + # --------------------------------------------------------------------------- # - - + + def _head_name(node: ast.expr) -> str | None: if isinstance(node, ast.Name): return node.id @@ -332,11 +353,13 @@ def mutable_names_in(annotation: ast.AST) -> Iterator[str]: yield from mutable_names_in(inner) for child in ast.iter_child_nodes(annotation): yield from mutable_names_in(child) - - + + def _mutable_ann(path: Path, line: int, name: str, where: str) -> Violation: return Violation( - path, line, "LIT001", + path, + line, + "LIT001", f"mutable `{name}` in {where}: a mutable collection can be grown or rewritten " f"by whoever holds it. Annotate a read-only view -- Mapping[...], Sequence[...], " f"AbstractSet[...], tuple[X, ...], frozenset[X], or a frozen dataclass / " @@ -345,37 +368,30 @@ def _mutable_ann(path: Path, line: int, name: str, where: str) -> Violation: ) -def _annotation_violations( - path: Path, annotation: ast.expr | None, line: int, where: str, ok_lines: frozenset[int] -) -> Iterator[Violation]: - if annotation is None or line in ok_lines: +def _annotation_violations(path: Path, annotation: ast.expr | None, line: int, where: str) -> Iterator[Violation]: + if annotation is None: return yield from (_mutable_ann(path, line, name, where) for name in mutable_names_in(annotation)) - - -def _function_violations( - path: Path, node: ast.FunctionDef | ast.AsyncFunctionDef, comments: Comments -) -> Iterator[Violation]: - mutable_ok = comments.mutable_ok_lines + + +def _function_violations(path: Path, node: ast.FunctionDef | ast.AsyncFunctionDef) -> Iterator[Violation]: args = node.args for arg in (*args.posonlyargs, *args.args, *args.kwonlyargs): - yield from _annotation_violations( - path, arg.annotation, arg.lineno, f"parameter `{arg.arg}` of `{node.name}`", mutable_ok - ) + yield from _annotation_violations(path, arg.annotation, arg.lineno, f"parameter `{arg.arg}` of `{node.name}`") # *args is allowed when typed (it's just a tuple); ruff ANN002 forces the # annotation, so here we only add the LIT001 mutable-collection check on the element type. if args.vararg is not None: - yield from _annotation_violations( - path, args.vararg.annotation, args.vararg.lineno, f"`*args` of `{node.name}`", mutable_ok - ) + yield from _annotation_violations(path, args.vararg.annotation, args.vararg.lineno, f"`*args` of `{node.name}`") # **kwargs is banned outright (LIT008): it erases the keyword contract and forces # Any-typing on everything it carries. ruff can require it be typed (ANN003) but # cannot ban the syntax, so this rule does. - if args.kwarg is not None and args.kwarg.lineno not in comments.kwargs_ok_lines: + if args.kwarg is not None: yield Violation( - path, args.kwarg.lineno, "LIT008", + path, + args.kwarg.lineno, + "LIT008", f"`**{args.kwarg.arg}` is banned: it erases the keyword contract and forces " f"Any-typing; declare explicit keyword parameters, or accept one frozen payload " f"(frozen dataclass / NamedTuple / ReadOnly TypedDict) " @@ -383,25 +399,20 @@ def _function_violations( ) if node.returns is not None: - yield from _annotation_violations( - path, node.returns, node.returns.lineno, f"return type of `{node.name}`", mutable_ok - ) - - -def iter_annotation_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]: + yield from _annotation_violations(path, node.returns, node.returns.lineno, f"return type of `{node.name}`") + + +def iter_annotation_violations(path: Path, tree: ast.AST) -> Iterator[Violation]: # Every annotation is in scope: signatures (params / *args / return) plus every # `x: T` -- class attribute, local, or module global. The latter three are all # ast.AnnAssign, so one walk covers them; only the signature annotations (which # are not AnnAssign) need the dedicated helper. for node in ast.walk(tree): if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): - yield from _function_violations(path, node, comments) + yield from _function_violations(path, node) elif isinstance(node, ast.AnnAssign): target = node.target.id if isinstance(node.target, ast.Name) else "" - yield from _annotation_violations( - path, node.annotation, node.lineno, - f"the type of `{target}`", comments.mutable_ok_lines, - ) + yield from _annotation_violations(path, node.annotation, node.lineno, f"the type of `{target}`") # --------------------------------------------------------------------------- # @@ -421,18 +432,20 @@ def _is_cast_call(node: ast.Call) -> bool: ) -def iter_cast_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]: +def iter_cast_violations(path: Path, tree: ast.AST) -> Iterator[Violation]: for node in ast.walk(tree): - if isinstance(node, ast.Call) and _is_cast_call(node) and node.lineno not in comments.cast_ok_lines: + if isinstance(node, ast.Call) and _is_cast_call(node): yield Violation( - path, node.lineno, "LIT006", + path, + node.lineno, + "LIT006", "cast() is an unchecked assertion (the type checker takes it on faith); " "validate into a frozen dataclass/NamedTuple/ReadOnly TypedDict at the " "boundary instead (suppress: `# cast-ok: `)", ) -def iter_guard_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]: +def iter_guard_violations(path: Path, tree: ast.AST) -> Iterator[Violation]: # TypeGuard/TypeIs are legal only as a function's return annotation (`-> TypeGuard[int]`), # so the walk is confined to `node.returns`; a runtime name that merely happens to read # `TypeGuard` is not a narrowing predicate. ruff bans the import; this flags the use. @@ -440,20 +453,18 @@ def iter_guard_violations(path: Path, tree: ast.AST, comments: Comments) -> Iter if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) or node.returns is None: continue for sub in ast.walk(node.returns): - name = ( - sub.id if isinstance(sub, ast.Name) - else sub.attr if isinstance(sub, ast.Attribute) - else None - ) - if name in UNSAFE_GUARDS and sub.lineno not in comments.guard_ok_lines: + name = sub.id if isinstance(sub, ast.Name) else sub.attr if isinstance(sub, ast.Attribute) else None + if name in UNSAFE_GUARDS: yield Violation( - path, sub.lineno, "LIT007", + path, + sub.lineno, + "LIT007", f"`{name}` narrowing predicate: the checker never verifies the body, so a " f"wrong guard silently corrupts types; parse into a concrete type instead " f"(suppress: `# guard-ok: `)", ) - - + + # --------------------------------------------------------------------------- # # Mutable-collection construction (LIT002) # --------------------------------------------------------------------------- # @@ -477,11 +488,7 @@ def _annotation_node_ids(tree: ast.AST) -> frozenset[int]: not construction, so the LIT002 walk must skip those subtrees. """ return frozenset( - id(sub) - for node in ast.walk(tree) - for ann in _annotations_of(node) - if ann is not None - for sub in ast.walk(ann) + id(sub) for node in ast.walk(tree) for ann in _annotations_of(node) if ann is not None for sub in ast.walk(ann) ) @@ -536,7 +543,9 @@ def _is_typeddict_annotation(annotation: ast.expr) -> bool: if head in TYPEDDICT_ANNOTATION_WRAPPERS: return _is_typeddict_annotation(annotation.slice) if head == "Annotated": - first = annotation.slice.elts[0] if isinstance(annotation.slice, ast.Tuple) and annotation.slice.elts else None + first = ( + annotation.slice.elts[0] if isinstance(annotation.slice, ast.Tuple) and annotation.slice.elts else None + ) return first is not None and _is_typeddict_annotation(first) return head is not None and head not in NON_TYPEDDICT_HEADS name = _head_name(annotation) @@ -591,7 +600,7 @@ def _construction_kind(node: ast.expr) -> str | None: return None -def iter_construction_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]: +def iter_construction_violations(path: Path, tree: ast.AST) -> Iterator[Violation]: in_annotation = _annotation_node_ids(tree) frozen_arguments = _frozen_argument_ids(tree) typeddict_builds = _typeddict_build_ids(tree) @@ -604,10 +613,12 @@ def iter_construction_violations(path: Path, tree: ast.AST, comments: Comments) ): continue kind = _construction_kind(node) - if kind is None or node.lineno in comments.mutable_ok_lines: + if kind is None: continue yield Violation( - path, node.lineno, "LIT002", + path, + node.lineno, + "LIT002", f"mutable {kind}: this builds a collection that can be grown or rewritten. " f"Build it in one shot and freeze it -- a tuple/frozenset wrapping a generator " f"(`tuple(f(x) for x in xs)`), a tuple literal, a frozen dataclass / NamedTuple, " @@ -615,8 +626,8 @@ def iter_construction_violations(path: Path, tree: ast.AST, comments: Comments) f"really must be dynamic) a MappingProxyType wrapping a dict literal or " f"comprehension (suppress: `# mutable-ok: `)", ) - - + + # --------------------------------------------------------------------------- # # Final-annotation discipline (LIT010) and argument immutability (LIT011) # --------------------------------------------------------------------------- # @@ -747,20 +758,11 @@ def _node_bindings(node: ast.AST, in_loop: bool) -> Iterator[Binding]: case ast.NamedExpr(target=ast.Name(id=name, lineno=line)): yield Binding(name, line, "walrus", in_loop) case ast.Import(names=aliases): - yield from ( - Binding((a.asname or a.name).partition(".")[0], node.lineno, "other", in_loop) - for a in aliases - ) + yield from (Binding((a.asname or a.name).partition(".")[0], node.lineno, "other", in_loop) for a in aliases) case ast.ImportFrom(names=aliases): - yield from ( - Binding(a.asname or a.name, node.lineno, "other", in_loop) - for a in aliases - if a.name != "*" - ) + yield from (Binding(a.asname or a.name, node.lineno, "other", in_loop) for a in aliases if a.name != "*") case ast.Delete(targets=targets): - yield from ( - Binding(t.id, t.lineno, "other", in_loop) for t in targets if isinstance(t, ast.Name) - ) + yield from (Binding(t.id, t.lineno, "other", in_loop) for t in targets if isinstance(t, ast.Name)) case ast.FunctionDef(name=name) | ast.AsyncFunctionDef(name=name) | ast.ClassDef(name=name): yield Binding(name, node.lineno, "other", in_loop) case ast.Global(names=names): @@ -792,9 +794,7 @@ def iter_scopes(tree: ast.AST) -> Iterator[ast.AST]: def _function_params(node: ast.FunctionDef | ast.AsyncFunctionDef | ast.Lambda) -> frozenset[str]: a = node.args - return frozenset( - p.arg for p in (*a.posonlyargs, *a.args, *a.kwonlyargs, a.vararg, a.kwarg) if p is not None - ) + return frozenset(p.arg for p in (*a.posonlyargs, *a.args, *a.kwonlyargs, a.vararg, a.kwarg) if p is not None) def _exempt_final_name(name: str) -> bool: @@ -812,26 +812,24 @@ def _is_config_surface(path: Path) -> bool: return path.parts[-2:] == CONFIG_SURFACE_PARTS -def iter_final_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]: +def iter_final_violations(path: Path, tree: ast.AST) -> Iterator[Violation]: for scope in iter_scopes(tree): if isinstance(scope, ast.Module) and _is_config_surface(path): continue - params = ( - _function_params(scope) - if isinstance(scope, (ast.FunctionDef, ast.AsyncFunctionDef)) - else frozenset() - ) + params = _function_params(scope) if isinstance(scope, (ast.FunctionDef, ast.AsyncFunctionDef)) else frozenset() bindings = scope_bindings(scope) declared = frozenset(b.name for b in bindings if b.form == "declared") first = _first_binding_index(bindings) for i, b in enumerate(bindings): if b.name in declared or b.name in params or b.in_loop: continue - if _exempt_final_name(b.name) or b.line in comments.rebind_ok_lines: + if _exempt_final_name(b.name): continue if b.form in ASSIGN_FORMS: yield Violation( - path, b.line, "LIT010", + path, + b.line, + "LIT010", f"`{b.name}` is assigned without a Final declaration, leaving it open to " f"rebinding: annotate `{b.name}: Final = ...` (or `Final[T]`, or a bare " f"`{b.name}: Final[T]` declaration with a single deferred assignment); " @@ -841,7 +839,9 @@ def iter_final_violations(path: Path, tree: ast.AST, comments: Comments) -> Iter ) elif b.form in IMPLICIT_FINAL_FORMS and i > first[b.name]: yield Violation( - path, b.line, "LIT010", + path, + b.line, + "LIT010", f"`{b.name}` is re-bound here after an earlier binding: unpacking and " f"walrus targets cannot carry Final, so their names are implicitly final; " f"bind a fresh name instead, or suppress with `# rebind-ok: `", @@ -895,9 +895,7 @@ def _iter_param_scopes( def _param_owners( scope: ast.AST, bindings: Sequence[Binding], enclosing: Sequence[_EnclosingFunction] ) -> Mapping[str, str]: - own_name = ( - scope.name if isinstance(scope, (ast.FunctionDef, ast.AsyncFunctionDef)) else "" - ) + own_name = scope.name if isinstance(scope, (ast.FunctionDef, ast.AsyncFunctionDef)) else "" nonlocal_params = { b.name: owner.name for b in bindings @@ -908,7 +906,7 @@ def _param_owners( return {**{p: own_name for p in _function_params(scope)}, **nonlocal_params} -def iter_param_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]: +def iter_param_violations(path: Path, tree: ast.AST) -> Iterator[Violation]: for scope, enclosing in _iter_param_scopes(tree): bindings = scope_bindings(scope) owners = _param_owners(scope, bindings, enclosing) @@ -917,19 +915,21 @@ def iter_param_violations(path: Path, tree: ast.AST, comments: Comments) -> Iter for b in bindings: if b.form in SCOPE_STATEMENT_FORMS or b.name not in owners: continue - if b.line in comments.rebind_ok_lines: - continue yield Violation( - path, b.line, "LIT011", + path, + b.line, + "LIT011", f"parameter `{b.name}` of `{owners[b.name]}` is re-bound: the name silently " f"detaches from what the caller passed; bind a new name instead " f"(suppress: `# rebind-ok: `)", ) for name, line in _mutation_sites(scope): - if name not in owners or name in SELF_PARAMS or line in comments.rebind_ok_lines: + if name not in owners or name in SELF_PARAMS: continue yield Violation( - path, line, "LIT011", + path, + line, + "LIT011", f"parameter `{name}` of `{owners[name]}` is mutated in place: the caller's " f"object is rewritten at a distance; build and return a new value instead " f"(suppress: `# rebind-ok: `)", @@ -1016,16 +1016,18 @@ def _functional_fields(tree: ast.AST) -> Iterator[_Field]: yield _Field(owner, key.value, value, value.lineno) -def iter_typeddict_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]: +def iter_typeddict_violations(path: Path, tree: ast.AST) -> Iterator[Violation]: fields = ( *(f for cls in _typeddict_classes(tree) for f in _class_fields(cls)), *_functional_fields(tree), ) for field in fields: - if _has_readonly_qualifier(field.annotation) or field.line in comments.writable_ok_lines: + if _has_readonly_qualifier(field.annotation): continue yield Violation( - path, field.line, "LIT012", + path, + field.line, + "LIT012", f"TypedDict field `{field.name}` of `{field.owner}` is writable: any holder " f"of the payload can rewrite the key after construction. Qualify it as " f"`ReadOnly[...]` (PEP 705; nests freely with Required/NotRequired/Annotated) " @@ -1033,36 +1035,76 @@ def iter_typeddict_violations(path: Path, tree: ast.AST, comments: Comments) -> ) +# --------------------------------------------------------------------------- # +# Suppression application and unused suppressions (LIT013) +# --------------------------------------------------------------------------- # + + +def apply_suppressions( + path: Path, + raw: Sequence[Violation], + suppressions: Mapping[str, frozenset[int]], +) -> tuple[Violation, ...]: + """Drop raw violations a valid `*-ok` marker suppresses; flag markers that suppress nothing.""" + kept = tuple( + v + for v in raw + if not any( + v.line in suppressions.get(ok.token, frozenset()) and v.code in ok.codes + for ok in OK_SUPPRESSIONS + ) + ) + unused = ( + Violation( + path, + line, + "LIT013", + f"`# {ok.token}` suppresses nothing: no " + f"{'/'.join(sorted(ok.codes))} violation on this line, so delete it", + ) + for ok in OK_SUPPRESSIONS + for line in sorted(suppressions.get(ok.token, frozenset())) + if not any(v.line == line and v.code in ok.codes for v in raw) + ) + return (*kept, *unused) + + # --------------------------------------------------------------------------- # # Driver # --------------------------------------------------------------------------- # - - + + def check_file(path: Path) -> tuple[Violation, ...]: try: source = path.read_text(encoding="utf-8") except (OSError, UnicodeDecodeError) as exc: return (Violation(path, 0, "LIT000", f"could not read file: {exc}"),) - - comments, violations = scan_comments(path, source) - + + suppressions, violations = scan_comments(path, source) + try: tree = ast.parse(source, filename=str(path)) except SyntaxError as exc: return (*violations, Violation(path, exc.lineno or 0, "LIT000", f"syntax error: {exc.msg}")) - + return ( *violations, - *iter_annotation_violations(path, tree, comments), - *iter_cast_violations(path, tree, comments), - *iter_guard_violations(path, tree, comments), - *iter_construction_violations(path, tree, comments), - *iter_final_violations(path, tree, comments), - *iter_param_violations(path, tree, comments), - *iter_typeddict_violations(path, tree, comments), + *apply_suppressions( + path, + ( + *iter_annotation_violations(path, tree), + *iter_cast_violations(path, tree), + *iter_guard_violations(path, tree), + *iter_construction_violations(path, tree), + *iter_final_violations(path, tree), + *iter_param_violations(path, tree), + *iter_typeddict_violations(path, tree), + ), + suppressions, + ), ) - - + + def collect_paths(raw: Iterable[str]) -> Iterator[Path]: for item in raw: p = Path(item) @@ -1070,8 +1112,8 @@ def collect_paths(raw: Iterable[str]) -> Iterator[Path]: yield from sorted(p.rglob("*.py")) elif p.suffix == ".py": yield p - - + + PARALLEL_MIN_PATHS = 200 MAX_WORKERS = 8 @@ -1099,18 +1141,17 @@ def main(argv: Sequence[str]) -> int: if not paths: print("usage: check_type_discipline.py ...", file=sys.stderr) return 2 - + targets = tuple(collect_paths(paths)) violations = sorted(scan_paths(targets)) for v in violations: print(v.render()) - + if violations: print(f"\n{len(violations)} violation(s).", file=sys.stderr) return 1 return 0 - - + + if __name__ == "__main__": raise SystemExit(main(sys.argv[1:])) - \ No newline at end of file diff --git a/scripts/type_discipline_gate.py b/scripts/type_discipline_gate.py index 40e61cf7265..4ba1a2ea393 100644 --- a/scripts/type_discipline_gate.py +++ b/scripts/type_discipline_gate.py @@ -17,7 +17,8 @@ without codes or reason), LIT006 (cast), LIT008 (`**kwargs`), LIT009 (inert LIT012 (TypedDict field without a `ReadOnly[...]` qualifier; suppress with `# writable-ok: `) carry limits at or above their current count to ratchet down; LIT005 (`*-ok` suppression without a reason) is frozen at limit 0 -so any net-new reasonless suppression trips the gate; and LIT007 +so any net-new reasonless suppression trips the gate; LIT013 (`*-ok` suppression +that suppresses nothing) is frozen at 0 for the same reason; and LIT007 (TypeGuard/TypeIs) is a hard zero. LIT010 and LIT011 were seeded at 1.5x the count left after the sweep that annotated every never-rebound name with Final, so that headroom is the hard @@ -129,7 +130,9 @@ def base_counts(ref: str) -> dict: # the body (or the `worktree add` itself) failed. rmtree is already best-effort. subprocess.run( ["git", "worktree", "remove", "--force", str(worktree)], - cwd=REPO_ROOT, capture_output=True, text=True, + cwd=REPO_ROOT, + capture_output=True, + text=True, ) shutil.rmtree(parent, ignore_errors=True) @@ -140,10 +143,7 @@ def over_ceiling(head: dict, budget: dict) -> frozenset: A rule can only breach when it is over its limit, so when none are the base comparison cannot change the verdict and the base worktree scan can be skipped. """ - return frozenset( - rule for rule, spec in budget.items() - if head.get(rule, 0) > spec["limit"] - ) + return frozenset(rule for rule, spec in budget.items() if head.get(rule, 0) > spec["limit"]) def evaluate(head: dict, base: dict, budget: dict) -> list: @@ -187,15 +187,11 @@ def cmd_check(base: str) -> None: return new = introduced( head, - parse_changed_lines( - _run(["git", "diff", base_point, "--unified=0", "--no-color", "--", TARGET]) - ), + parse_changed_lines(_run(["git", "diff", base_point, "--unified=0", "--no-color", "--", TARGET])), ) print(f"FAIL: LIT-rule totals exceed their limit (base {base}):") for breach in breaches: - print( - f" {breach.rule}: total {breach.total} over limit {breach.cap} (this change added {breach.added})" - ) + print(f" {breach.rule}: total {breach.total} over limit {breach.cap} (this change added {breach.added})") for violation in sorted(v for v in new if v.code == breach.rule): print(f" {violation.file}:{violation.line}") print( @@ -221,7 +217,8 @@ def ratcheted_budget(budget: dict, current: dict, base: dict, seeded: frozenset """ return { rule: { - "limit": spec["limit"] if rule in seeded + "limit": spec["limit"] + if rule in seeded else max(0, spec["limit"] - max(0, base.get(rule, 0) - current.get(rule, 0))) } for rule, spec in sorted(budget.items()) @@ -231,7 +228,9 @@ def ratcheted_budget(budget: dict, current: dict, base: dict, seeded: frozenset def _base_budget_rules(base_point: str) -> frozenset: proc = subprocess.run( ["git", "show", f"{base_point}:{BUDGET_PATH.name}"], - cwd=REPO_ROOT, capture_output=True, text=True, + cwd=REPO_ROOT, + capture_output=True, + text=True, ) if proc.returncode != 0: return frozenset() @@ -248,17 +247,12 @@ def cmd_update(base_ref: str) -> None: budget = json.loads(BUDGET_PATH.read_text()) base_point = resolve_base_point(base_ref) seeded = frozenset(budget) - _base_budget_rules(base_point) - updated = ratcheted_budget( - budget, count_by_rule(head_violations()), base_counts(base_point), seeded - ) + updated = ratcheted_budget(budget, count_by_rule(head_violations()), base_counts(base_point), seeded) BUDGET_PATH.write_text(json.dumps(updated, indent=2, sort_keys=True) + "\n") cleared = sum(budget[rule]["limit"] - updated[rule]["limit"] for rule in updated) print(f"Ratcheted LIT-rule limits down by {cleared} violations this branch fixed") if seeded: - print( - "Left untouched (seeded on this branch, absent from the base budget): " - + ", ".join(sorted(seeded)) - ) + print("Left untouched (seeded on this branch, absent from the base budget): " + ", ".join(sorted(seeded))) def main() -> None: diff --git a/tests/e2e/lifecycle.py b/tests/e2e/lifecycle.py index eb9704d4dcb..1ccf1bdbefa 100644 --- a/tests/e2e/lifecycle.py +++ b/tests/e2e/lifecycle.py @@ -54,9 +54,7 @@ class ResourceManager: client: ResourceClient strict_cleanup: bool = False - _cleanups: List[Callable[[], object]] = field( - default_factory=list - ) # mutable-ok: append-only teardown registry + _cleanups: List[Callable[[], object]] = field(default_factory=list) def init(self) -> None: """No global setup needed today; present for lifecycle symmetry.""" @@ -85,8 +83,7 @@ class ResourceManager: def teardown(self) -> None: failures: Final = tuple( - failure for cleanup in reversed(self._cleanups) - if (failure := _run_cleanup(cleanup)) is not None + failure for cleanup in reversed(self._cleanups) if (failure := _run_cleanup(cleanup)) is not None ) if failures and self.strict_cleanup: raise ExceptionGroup("Resource cleanup failed", failures) diff --git a/tests/e2e/load/proxy_usage.py b/tests/e2e/load/proxy_usage.py index 83463c078b8..b4e28478cdf 100644 --- a/tests/e2e/load/proxy_usage.py +++ b/tests/e2e/load/proxy_usage.py @@ -160,5 +160,5 @@ class ProxyUsageSampler: """ with self._lock: taken = tuple(self._samples) - self._samples = [taken[-1]] if taken else [] # rebind-ok: drains the buffer under the lock + self._samples = [taken[-1]] if taken else [] return UsageWindow(samples=taken) diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index 9c321269e38..4986f5ddcc0 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -51,7 +51,6 @@ def _owned(nodeid: str) -> bool: def pytest_collection_modifyitems(config: pytest.Config, items: list[pytest.Item]) -> None: order_seed: Final = config.getoption("integration_order_seed") if order_seed: - # rebind-ok: pytest requires this hook to reorder its shared collection list in place. items.sort(key=lambda item: hashlib.sha256(f"{order_seed}:{item.nodeid}".encode()).digest()) root: Final = Path(__file__).parent owned: Final = tuple( diff --git a/tests/test_litellm/proxy/db/test_exception_handler_reconnect_retry.py b/tests/test_litellm/proxy/db/test_exception_handler_reconnect_retry.py index 4286da23242..26ab6798d8e 100644 --- a/tests/test_litellm/proxy/db/test_exception_handler_reconnect_retry.py +++ b/tests/test_litellm/proxy/db/test_exception_handler_reconnect_retry.py @@ -27,9 +27,7 @@ def _make_client( `call_with_db_reconnect_retry` actually pokes at.""" client = MagicMock() if has_attempt_db_reconnect: - client.attempt_db_reconnect = AsyncMock( - return_value=attempt_db_reconnect_return - ) + client.attempt_db_reconnect = AsyncMock(return_value=attempt_db_reconnect_return) else: # `hasattr(client, "attempt_db_reconnect")` must return False — MagicMock # auto-creates attributes, so we wipe it out via `spec`. @@ -127,9 +125,7 @@ async def test_call_with_db_reconnect_retry_propagates_after_second_transport_er raise httpx.ReadError("still failing") with pytest.raises(httpx.ReadError): - await call_with_db_reconnect_retry( - client, _factory, reason="second_transport_error" - ) + await call_with_db_reconnect_retry(client, _factory, reason="second_transport_error") assert len(invocations) == 2 client.attempt_db_reconnect.assert_awaited_once() @@ -166,9 +162,7 @@ async def test_call_with_db_reconnect_retry_invokes_factory_twice_not_same_coro( raise httpx.ReadError("transport blip") return "ok" - result = await call_with_db_reconnect_retry( - client, _factory, reason="fresh_coro_on_retry" - ) + result = await call_with_db_reconnect_retry(client, _factory, reason="fresh_coro_on_retry") assert result == "ok" assert factory_call_count == 2 @@ -243,14 +237,13 @@ async def test_call_with_db_reconnect_retry_preserves_original_error_when_reconn raise original_exc with pytest.raises(httpx.ReadError) as exc_info: - await call_with_db_reconnect_retry( - client, _factory, reason="reconnect_itself_raises" - ) + await call_with_db_reconnect_retry(client, _factory, reason="reconnect_itself_raises") assert exc_info.value is original_exc assert exc_info.value.__cause__ is reconnect_exc client.attempt_db_reconnect.assert_awaited_once() + @pytest.mark.asyncio async def test_call_with_db_reconnect_retry_honors_narrowed_retry_safe_types(): """A non-idempotent write can pass `retry_safe_error_types` to opt out of @@ -259,7 +252,7 @@ async def test_call_with_db_reconnect_retry_honors_narrowed_retry_safe_types(): attempts = 0 async def _factory(): - nonlocal attempts # rebind-ok: attempt counter for a two-call helper + nonlocal attempts attempts += 1 raise httpx.ReadError("ambiguous") @@ -283,7 +276,7 @@ async def test_call_with_db_reconnect_retry_default_covers_every_transport_error attempts = 0 async def _factory(): - nonlocal attempts # rebind-ok: attempt counter for a two-call helper + nonlocal attempts attempts += 1 if attempts == 1: raise ClientNotConnectedError() diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_conduct.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_conduct.py index 323756f8fa0..a7c777248c0 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_conduct.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_conduct.py @@ -221,14 +221,10 @@ def test_missing_package_fails_at_config_load_with_install_hint() -> None: def test_plugin_that_swallows_unreachable_fallback_into_kwargs_is_rejected() -> None: class Swallowing: - def __init__( - self, *, fail_mode: str = "fail_closed", **kwargs: object - ) -> None: ... # kwargs-ok: models plugin 0.2.4 + def __init__(self, *, fail_mode: str = "fail_closed", **kwargs: object) -> None: ... class Binding: - def __init__( - self, *, unreachable_fallback: str | None = None, **kwargs: object - ) -> None: ... # kwargs-ok: plugin 0.2.5 + def __init__(self, *, unreachable_fallback: str | None = None, **kwargs: object) -> None: ... assert not binds_unreachable_fallback(Swallowing) assert binds_unreachable_fallback(Binding) diff --git a/tests/test_litellm/test_check_type_discipline.py b/tests/test_litellm/test_check_type_discipline.py index 2d49332e687..b0c8d2d5d56 100644 --- a/tests/test_litellm/test_check_type_discipline.py +++ b/tests/test_litellm/test_check_type_discipline.py @@ -40,17 +40,17 @@ def test_scan_comments_tokenizes_every_comment(): # was tokenized, and the valid cast-ok suppression line must be captured. A crash in the # readline path would leave both empty. source = "x = 1 # noqa\ny = 2 # cast-ok: validated upstream by the caller\n" - comments, violations = checker.scan_comments(Path("snippet.py"), source) + suppressions, violations = checker.scan_comments(Path("snippet.py"), source) assert [v.code for v in violations] == ["LIT003"] - assert comments.cast_ok_lines == frozenset({2}) + assert suppressions["cast-ok"] == frozenset({2}) def test_scan_comments_does_not_crash_on_malformed_source(): # A dedent mismatch makes tokenize raise IndentationError (a SyntaxError subclass); # scan_comments must swallow it, not propagate and crash the whole run. - comments, violations = checker.scan_comments(Path("x.py"), "if True:\n a = 1\n b = 2\n") + suppressions, violations = checker.scan_comments(Path("x.py"), "if True:\n a = 1\n b = 2\n") assert violations == () - assert comments.cast_ok_lines == frozenset() + assert suppressions["cast-ok"] == frozenset() def test_malformed_source_degrades_to_lit000(tmp_path): @@ -107,6 +107,38 @@ def test_ok_suppression_without_reason_is_flagged(tmp_path): assert "LIT002" in codes # and it does not suppress, so the construction still trips +def test_mutable_ok_on_a_real_violation_suppresses_and_is_not_lit013(tmp_path): + codes = _codes(tmp_path, "x: Final = [] # mutable-ok: seed\n") + assert "LIT002" not in codes + assert "LIT013" not in codes + + +def test_mutable_ok_on_a_clean_line_is_lit013(tmp_path): + f = tmp_path / "snippet.py" + f.write_text("x: Final = (1, 2) # mutable-ok: stale\n", encoding="utf-8") + found = checker.check_file(f) + assert [v.code for v in found] == ["LIT013"] + assert "mutable-ok" in found[0].message + + +def test_mutable_ok_does_not_suppress_rebind_codes(tmp_path): + codes = _codes(tmp_path, "x = 1 # mutable-ok: wrong token\n") + assert "LIT010" in codes + assert "LIT013" in codes + + +def test_rebind_ok_on_a_real_param_rebind_is_not_lit013(tmp_path): + codes = _codes(tmp_path, "def f(p: int) -> None:\n p = 2 # rebind-ok: reset\n") + assert "LIT011" not in codes + assert "LIT013" not in codes + + +def test_reasonless_ok_on_a_clean_line_is_lit005_not_lit013(tmp_path): + codes = _codes(tmp_path, "x: Final = (1, 2) # mutable-ok\n") + assert "LIT005" in codes + assert "LIT013" not in codes + + # --------------------------------------------------------------------------- # # Mutable annotations (LIT001) and construction (LIT002) # --------------------------------------------------------------------------- # @@ -213,15 +245,11 @@ def test_typeddict_annotated_dict_literal_is_exempt(tmp_path): def test_wrapped_typeddict_annotations_share_the_exemption(tmp_path): - assert "LIT002" not in _codes( - tmp_path, "from typing import Final, Optional\nx: Final[Optional[MyTD]] = {'a': 1}\n" - ) + assert "LIT002" not in _codes(tmp_path, "from typing import Final, Optional\nx: Final[Optional[MyTD]] = {'a': 1}\n") assert "LIT002" not in _codes( tmp_path, "from typing import Annotated, Final\nx: Final[Annotated[MyTD, 'meta']] = {'a': 1}\n" ) - assert "LIT002" not in _codes( - tmp_path, "from typing import ClassVar\nclass C:\n x: ClassVar[MyTD] = {'a': 1}\n" - ) + assert "LIT002" not in _codes(tmp_path, "from typing import ClassVar\nclass C:\n x: ClassVar[MyTD] = {'a': 1}\n") assert "LIT002" not in _codes(tmp_path, "from typing import Final\nx: Final[MyTD | None] = {'a': 1}\n") assert "LIT002" in _codes(tmp_path, "from typing import Final\nx: Final[dict[str, int] | None] = {'a': 1}\n") @@ -234,7 +262,8 @@ def test_bare_final_dict_literal_still_counts(tmp_path): def test_non_typeddict_annotations_do_not_exempt(tmp_path): assert "LIT002" in _codes(tmp_path, "from typing import Final\nx: Final[dict[str, int]] = {'a': 1}\n") assert "LIT002" in _codes( - tmp_path, "from collections.abc import Mapping\nfrom typing import Final\nx: Final[Mapping[str, int]] = {'a': 1}\n" + tmp_path, + "from collections.abc import Mapping\nfrom typing import Final\nx: Final[Mapping[str, int]] = {'a': 1}\n", ) assert "LIT002" in _codes(tmp_path, "from typing import Any, Final\nx: Final[Any] = {'a': 1}\n") assert "LIT002" in _codes(tmp_path, "from typing import Final\nx: Final[object] = {'a': 1}\n") @@ -372,10 +401,7 @@ def test_walrus_rebinding_is_flagged(tmp_path): def test_unpack_after_global_declaration_is_flagged(tmp_path): src = ( - "count = 0 # rebind-ok: seeded module counter\n" - "def f() -> None:\n" - " global count\n" - " count, other = (1, 2)\n" + "count = 0 # rebind-ok: seeded module counter\ndef f() -> None:\n global count\n count, other = (1, 2)\n" ) assert _codes(tmp_path, src).count("LIT010") == 1 @@ -411,14 +437,7 @@ def test_non_assignment_binding_forms_are_exempt(tmp_path): def test_dunder_underscore_class_body_and_type_alias_are_exempt(tmp_path): - src = ( - "from typing import TypeAlias\n" - "__all__ = ['C']\n" - "_ = 1\n" - "Alias: TypeAlias = str\n" - "class C:\n" - " field = 1\n" - ) + src = "from typing import TypeAlias\n__all__ = ['C']\n_ = 1\nAlias: TypeAlias = str\nclass C:\n field = 1\n" assert "LIT010" not in _codes(tmp_path, src) @@ -428,12 +447,7 @@ def test_comprehension_targets_are_exempt(tmp_path): def test_global_reassignment_inside_function_is_flagged(tmp_path): - src = ( - "count = 0 # rebind-ok: seeded module counter\n" - "def bump() -> None:\n" - " global count\n" - " count = 1\n" - ) + src = "count = 0 # rebind-ok: seeded module counter\ndef bump() -> None:\n global count\n count = 1\n" assert _codes(tmp_path, src).count("LIT010") == 1 @@ -585,11 +599,7 @@ def test_walrus_in_own_defaults_binds_in_enclosing_scope_not_the_parameter(tmp_p def test_walrus_in_nested_defaults_rebinds_the_enclosing_parameter(tmp_path): - src = ( - "def g(p: int) -> None:\n" - " def inner(q: int = (p := 2)) -> None:\n" - " return None\n" - ) + src = "def g(p: int) -> None:\n def inner(q: int = (p := 2)) -> None:\n return None\n" assert "LIT011" in _codes(tmp_path, src) @@ -604,11 +614,7 @@ def test_typeddict_writable_field_is_flagged(tmp_path): def test_typeddict_readonly_field_is_clean(tmp_path): - src = ( - "from typing_extensions import ReadOnly, TypedDict\n" - "class P(TypedDict):\n" - " a: ReadOnly[int]\n" - ) + src = "from typing_extensions import ReadOnly, TypedDict\nclass P(TypedDict):\n a: ReadOnly[int]\n" assert "LIT012" not in _codes(tmp_path, src) @@ -640,11 +646,7 @@ def test_readonly_in_annotated_metadata_position_does_not_qualify(tmp_path): def test_typeddict_subclass_in_same_module_is_flagged(tmp_path): src = ( - "from typing import TypedDict\n" - "class Base(TypedDict):\n" - " pass\n" - "class Child(Base, total=False):\n" - " a: int\n" + "from typing import TypedDict\nclass Base(TypedDict):\n pass\nclass Child(Base, total=False):\n a: int\n" ) assert "LIT012" in _codes(tmp_path, src) @@ -677,11 +679,7 @@ def test_writable_ok_with_reason_suppresses_lit012(tmp_path): def test_writable_ok_without_reason_is_lit005_and_does_not_suppress(tmp_path): - src = ( - "from typing import TypedDict\n" - "class P(TypedDict):\n" - " a: int # writable-ok\n" - ) + src = "from typing import TypedDict\nclass P(TypedDict):\n a: int # writable-ok\n" codes = _codes(tmp_path, src) assert "LIT005" in codes assert "LIT012" in codes @@ -717,7 +715,9 @@ def _corpus(tmp_path: Path, count: int) -> tuple[Path, ...]: def _run_checker(target: Path) -> list[str]: completed = subprocess.run( [sys.executable, str(_MODULE_PATH), str(target)], - capture_output=True, text=True, timeout=300, + capture_output=True, + text=True, + timeout=300, ) return completed.stdout.splitlines() @@ -727,9 +727,7 @@ def test_worker_count_stays_serial_below_the_threshold(): def test_worker_count_fans_out_at_the_threshold(): - assert checker._worker_count(checker.PARALLEL_MIN_PATHS) == max( - 1, min(os.cpu_count() or 1, checker.MAX_WORKERS) - ) + assert checker._worker_count(checker.PARALLEL_MIN_PATHS) == max(1, min(os.cpu_count() or 1, checker.MAX_WORKERS)) def test_worker_count_never_exceeds_the_cap(): diff --git a/tests/test_litellm_rust/support/isolation.py b/tests/test_litellm_rust/support/isolation.py index f98ce4843a8..26c7cd0f875 100644 --- a/tests/test_litellm_rust/support/isolation.py +++ b/tests/test_litellm_rust/support/isolation.py @@ -29,7 +29,7 @@ def _list_attribute(container: ModuleType, attribute: str) -> list[object]: def _isolated_list(container: ModuleType, attribute: str) -> Generator[None]: source: Final = _list_attribute(container, attribute) original: Final = list(source) - source.clear() # mutable-ok: test isolation mutates global registries by design + source.clear() try: yield finally: @@ -54,5 +54,5 @@ def isolated_callback_registries() -> Generator[None]: for attribute in CALLBACK_ATTRIBUTES: stack.enter_context(_isolated_list(litellm, attribute)) stack.enter_context(_isolated_list(litellm_logging, "_in_memory_loggers")) # pyright: ignore[reportPrivateUsage] # no public callback-cache accessor - stack.enter_context(rebound(utils, "callback_list", [])) # rebind-ok: isolate legacy callback registry + stack.enter_context(rebound(utils, "callback_list", [])) yield diff --git a/tests/unit/messages/test_dispatch.py b/tests/unit/messages/test_dispatch.py index 48eb1adbf51..88ef849f0e2 100644 --- a/tests/unit/messages/test_dispatch.py +++ b/tests/unit/messages/test_dispatch.py @@ -216,9 +216,7 @@ def test_binding_errors_delegate_to_python(args: tuple[object, ...], kwargs: Map captured: Final[list[tuple[tuple[object, ...], Mapping[str, object]]]] = [] expected: Final = response() - def python( - *call_args: object, **call_kwargs: object - ) -> AnthropicMessagesResponse: # kwargs-ok: records invalid call + def python(*call_args: object, **call_kwargs: object) -> AnthropicMessagesResponse: captured.append((call_args, call_kwargs)) return expected diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 0c0952289e2..beeb44474da 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -34,5 +34,8 @@ }, "LIT012": { "limit": 4486 + }, + "LIT013": { + "limit": 0 } } From ce582affaa2cddac0b61778c1852ac8913888fba Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 17:50:27 -0700 Subject: [PATCH 09/10] fix(mcp): reject duplicate MCP server names and aliases (#42791) * fix(mcp): reject duplicate MCP server names and aliases MCP server_name and alias were unchecked at write time, so two servers could share one tool prefix and tool routing resolved to an arbitrary winner. Writes now run inside an advisory-locked transaction that rejects a collision on either column case-insensitively with a 400 naming the colliding identifier, covering create, edit, connector import and restricted-admin submission. Server reload logs one warning per identifier already shared in the database. Co-Authored-By: bot_apk * fix(ui): block duplicate MCP server names and aliases before submit The create and edit forms now check the normalized name/alias against the loaded server list (case-insensitive, spaces to underscores, own row excluded on edit) and show a field error instead of submitting. Structured proxy error bodies are unwrapped so a 400 no longer renders as 'Error: [object Object]'. Co-Authored-By: bot_apk * fix(mcp): check identifier conflicts when an alias is cleared Clearing an alias drops the tool prefix to the stored server_name, so that name must go through the conflict check too; an explicit alias:null is now treated as an identifier write. Also narrows the new db tests to behavioral assertions instead of pinning prisma where shapes. Co-Authored-By: bot_apk * fix(mcp): treat an empty alias as a clear in conflict checks An empty-string alias was written unchecked even though the prefix falls back to server_name; the update path now treats any falsy alias like a clear. The edit form likewise compares a cleared alias as empty instead of re-checking the alias being removed. Co-Authored-By: bot_apk * test(mcp): cover clearing an alias to an empty string Co-Authored-By: bot_apk --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: bot_apk --- litellm/proxy/_experimental/mcp_server/db.py | 221 +++++++++++++++- .../mcp_server/discoverable_endpoints.py | 3 +- .../mcp_server/mcp_server_manager.py | 30 +++ .../mcp_management_endpoints.py | 44 +++- tests/integration/mcp/test_mcp_management.py | 60 ++++- .../mcp_server/test_mcp_partial_update.py | 180 ++++++++++++- .../mcp_server/test_mcp_server_manager.py | 59 +++++ .../test_mcp_management_endpoints.py | 236 +++++++++++++++++- .../_components/CreateMCPServer.tsx | 15 +- .../_components/duplicateServerCheck.test.ts | 49 ++++ .../_components/duplicateServerCheck.ts | 48 ++++ .../_components/mcp_server_edit.tsx | 17 +- .../_components/mcp_server_view.tsx | 3 + .../mcp-servers/_components/mcp_servers.tsx | 2 + 14 files changed, 938 insertions(+), 29 deletions(-) create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/duplicateServerCheck.test.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/duplicateServerCheck.ts diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 18723a76b2b..70c6e6f4bf3 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -3,6 +3,7 @@ import binascii import hashlib import json from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence +from dataclasses import dataclass from datetime import datetime, timedelta, timezone from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypedDict, cast @@ -64,6 +65,7 @@ if TYPE_CHECKING: class _UserEnvVarsTransactionClient(Protocol): litellm_mcpuserenvvars: "TableActions[prisma_db_models.LiteLLM_MCPUserEnvVars]" + litellm_mcpservertable: "TableActions[prisma_db_models.LiteLLM_MCPServerTable]" async def execute_raw(self, query: str, *args: object) -> int: ... @@ -74,6 +76,19 @@ class _UserEnvVarsTransaction(Protocol): async def __aexit__(self, exc_type: object, exc_value: object, traceback: object) -> bool | None: ... +@dataclass(frozen=True, slots=True) +class McpIdentifierConflict: + """An incoming ``server_name``/``alias`` already belongs to another MCP server row. + + ``field`` is the incoming identifier that collided, ``value`` the submitted + string, and ``server_id`` the existing row that owns it. + """ + + field: Literal["server_name", "alias"] + value: str + server_id: str + + _AUTH_FLOW_SCOPED_FIELDS: Final["frozenset[str]"] = frozenset( { "issuer", @@ -500,6 +515,121 @@ def _db_transaction_manager(prisma_client: PrismaClient) -> _UserEnvVarsTransact return manager +def _identifier_where(value: str, exclude_server_id: str | None) -> "prisma_db_types.LiteLLM_MCPServerTableWhereInput": + own_row_guard: Final = ( + ({"NOT": [{"server_id": exclude_server_id}]},) # mutable-ok: prisma where-inputs must be plain dicts + if exclude_server_id is not None + else () + ) + where: Final[prisma_db_types.LiteLLM_MCPServerTableWhereInput] = { + "AND": [ # mutable-ok: prisma where-inputs must be plain dicts + { + "OR": [ # mutable-ok: prisma where-inputs must be plain dicts + {"server_name": {"equals": value, "mode": "insensitive"}}, + {"alias": {"equals": value, "mode": "insensitive"}}, + ] + }, + { + "OR": [{"approval_status": None}, {"approval_status": {"not": MCPApprovalStatus.draft}}] + }, # mutable-ok: prisma where-inputs must be plain dicts + *own_row_guard, + ] + } + return where + + +def _identifier_field(data_dict: "Mapping[str, object]", field: str) -> str | None: + value: Final = data_dict.get(field) + return value if isinstance(value, str) else None + + +async def _find_mcp_server_identifier_conflict( + table: "TableActions[prisma_db_models.LiteLLM_MCPServerTable]", + *, + server_name: str | None, + alias: str | None, + exclude_server_id: str | None, +) -> McpIdentifierConflict | None: + """Return the collision between an incoming identifier and a stored row, else None. + + Each non-empty incoming identifier is compared case-insensitively against + BOTH the ``server_name`` and ``alias`` columns, because a value that matches + either column would still share the tool prefix another server answers to. + ``alias`` is checked first so the reported field is deterministic. Draft + rows back the transient OAuth session flow and never reach the registry, so + they cannot collide. NULL ``approval_status`` predates the approval + workflow and is kept via the inner OR, matching ``get_all_mcp_servers``. + """ + candidates: Final[tuple[tuple[Literal["alias", "server_name"], str | None], ...]] = ( + ("alias", alias), + ("server_name", server_name), + ) + for field_name, value in candidates: + if not value: + continue + if (row := await table.find_first(where=_identifier_where(value, exclude_server_id))) is not None: + return McpIdentifierConflict(field=field_name, value=value, server_id=row.server_id) + return None + + +async def find_mcp_server_identifier_conflict( + prisma_client: PrismaClient, + *, + server_name: str | None, + alias: str | None, + exclude_server_id: str | None, +) -> McpIdentifierConflict | None: + """Unlocked identifier-collision check, for callers outside a write path.""" + return await _find_mcp_server_identifier_conflict( + _mcp_server_table_actions(prisma_client), + server_name=server_name, + alias=alias, + exclude_server_id=exclude_server_id, + ) + + +def _mcp_identifier_lock_keys(*identifiers: str | None) -> tuple[int, ...]: + """Deterministic advisory-lock keys for the lowercased identifiers, sorted + so concurrent requests for the same pair always lock in the same order.""" + return tuple( + int.from_bytes( + hashlib.blake2b(f"mcp_identifier:{normalized}".encode(), digest_size=8).digest(), + "big", + signed=True, + ) + for normalized in sorted(frozenset(value.lower() for value in identifiers if value)) + ) + + +async def _mcp_server_write_if_identifier_free( + prisma_client: PrismaClient, + *, + server_name: str | None, + alias: str | None, + exclude_server_id: str | None, + write: "Callable[[TableActions[prisma_db_models.LiteLLM_MCPServerTable]], Awaitable[prisma_db_models.LiteLLM_MCPServerTable | None]]", +) -> "prisma_db_models.LiteLLM_MCPServerTable | McpIdentifierConflict | None": + """Run ``write`` only when no other live row owns ``server_name``/``alias``. + + The conflict check and the write share a transaction guarded by per-identifier + advisory locks, so two concurrent requests for the same name cannot both + pass the check and both insert. + """ + lock_keys: Final = _mcp_identifier_lock_keys(server_name, alias) + async with _db_transaction_manager(prisma_client) as tx: + for lock_key in lock_keys: + await tx.execute_raw("SELECT pg_advisory_xact_lock($1::bigint)", lock_key) + conflict: Final = await _find_mcp_server_identifier_conflict( + tx.litellm_mcpservertable, + server_name=server_name, + alias=alias, + exclude_server_id=exclude_server_id, + ) + if conflict is not None: + return conflict + return await write(tx.litellm_mcpservertable) + + async def _db_find_mcp_server_rows( prisma_client: PrismaClient, where: "prisma_db_types.LiteLLM_MCPServerTableWhereInput | None" = None, @@ -880,6 +1010,43 @@ async def create_mcp_server( return LiteLLM_MCPServerTable.model_validate(new_mcp_server.model_dump()) +async def create_mcp_server_if_identifier_free( + prisma_client: PrismaClient, data: NewMCPServerRequest, touched_by: str +) -> LiteLLM_MCPServerTable | McpIdentifierConflict: + """Create the row only when no other live server owns ``server_name``/``alias``. + + Returns the McpIdentifierConflict instead of inserting when the collision + check finds an existing row; the advisory-lock transaction keeps two + concurrent creates of the same identifier from both passing. + """ + if data.server_id is None: + data.server_id = str(uuid.uuid4()) + + data_dict: Final = _prepare_mcp_server_data(data) + data_dict["created_by"] = touched_by + data_dict["updated_by"] = touched_by + + async def _create( + table: "TableActions[prisma_db_models.LiteLLM_MCPServerTable]", + ) -> "prisma_db_models.LiteLLM_MCPServerTable | None": + return await table.create(data=data_dict) + + written: Final = await _mcp_server_write_if_identifier_free( + prisma_client, + server_name=_identifier_field(data_dict, "server_name"), + alias=_identifier_field(data_dict, "alias"), + exclude_server_id=None, + write=_create, + ) + if isinstance(written, McpIdentifierConflict): + return written + if written is None: + raise RuntimeError("inserted MCP server row missing") + + _decrypt_env_vars_on_returned_row(written) + return LiteLLM_MCPServerTable.model_validate(written.model_dump()) + + async def create_draft_mcp_server( prisma_client: PrismaClient, data: NewMCPServerRequest, @@ -970,14 +1137,57 @@ async def get_draft_mcp_server( return table +async def _update_mcp_server_row( + prisma_client: PrismaClient, + *, + server_id: str, + data_dict: Mapping[str, object], +) -> "prisma_db_models.LiteLLM_MCPServerTable | McpIdentifierConflict | None": + identifier_write: Final = any(field in data_dict for field in ("server_name", "alias")) + + async def _update( + table: "TableActions[prisma_db_models.LiteLLM_MCPServerTable]", + ) -> "prisma_db_models.LiteLLM_MCPServerTable | None": + return await table.update( + where={"server_id": server_id}, # mutable-ok: prisma where-inputs must be plain dicts + data=data_dict, + ) + + if not identifier_write: + return await _update(_mcp_server_table_actions(prisma_client)) + if "alias" in data_dict and not data_dict["alias"] and "server_name" not in data_dict: + # Clearing the alias drops the prefix to the stored server_name, which + # may already belong to another row, so that name needs the check too. + existing: Final = await _db_find_mcp_server_row(prisma_client, server_id) + if existing is None: + return await _update(_mcp_server_table_actions(prisma_client)) + return await _mcp_server_write_if_identifier_free( + prisma_client, + server_name=existing.server_name, + alias=None, + exclude_server_id=server_id, + write=_update, + ) + return await _mcp_server_write_if_identifier_free( + prisma_client, + server_name=_identifier_field(data_dict, "server_name"), + alias=_identifier_field(data_dict, "alias"), + exclude_server_id=server_id, + write=_update, + ) + + async def update_mcp_server( prisma_client: PrismaClient, data: UpdateMCPServerRequest, touched_by: str, fields_set: set[str] | None = None, -) -> LiteLLM_MCPServerTable | None: +) -> LiteLLM_MCPServerTable | McpIdentifierConflict | None: """ Update a new mcp server record in the db + + Returns McpIdentifierConflict instead of writing when the update would put + ``server_name``/``alias`` onto identifiers another live row already owns. """ from litellm.litellm_core_utils.safe_json_dumps import safe_dumps @@ -1086,11 +1296,14 @@ async def update_mcp_server( data_dict["credentials"] = Json(None) - updated_mcp_server: Final = await MCPServerRepository(prisma_client).table.update( - where={"server_id": data.server_id}, - data=data_dict, + updated_mcp_server: Final = await _update_mcp_server_row( + prisma_client, + server_id=data.server_id, + data_dict=data_dict, ) + if isinstance(updated_mcp_server, McpIdentifierConflict): + return updated_mcp_server _decrypt_env_vars_on_returned_row(updated_mcp_server) return LiteLLM_MCPServerTable.model_validate(updated_mcp_server.model_dump()) if updated_mcp_server else None diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 64bab0a7832..ade829a1b67 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -1570,6 +1570,7 @@ async def _persist_dcr_client_registration( } from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # avoids circular import + McpIdentifierConflict, update_mcp_server, upsert_mcp_server_oauth_client_credentials, ) @@ -1601,7 +1602,7 @@ async def _persist_dcr_client_registration( ), touched_by="mcp_oauth_dcr", ) - if updated_row is not None: + if updated_row is not None and not isinstance(updated_row, McpIdentifierConflict): await global_mcp_server_manager.update_server(updated_row) return "persisted" if global_mcp_server_manager.is_config_declared_server(mcp_server.server_id): diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 20c114a2f3e..0c520142fb3 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1453,6 +1453,35 @@ def _warn_on_server_name_fields( _warn("server_name", server_name) +def _warn_on_shared_identifier_prefixes(servers: Iterable[MCPServer]) -> None: + """Warn once per identifier that several servers share. + + ``get_server_prefix`` resolves alias first, so two servers sharing a + lowercased ``alias or server_name`` publish the same tool prefix and calls + routed by that prefix are ambiguous. A write-time uniqueness check keeps + new collisions out; this surfaces the ones already stored. + """ + pairs: Final = tuple( + ((server.alias or server.server_name or "").lower(), server.server_id) + for server in servers + if server.alias or server.server_name + ) + groups: Final = MappingProxyType( + { + identifier: tuple(sorted(server_id for key, server_id in pairs if key == identifier)) + for identifier in frozenset(key for key, _server_id in pairs) + } + ) + for identifier, server_ids in groups.items(): + if len(server_ids) > 1: + verbose_logger.warning( + "MCP servers %s share the identifier '%s'; tool routing for that prefix is ambiguous. " + "Rename or delete all but one.", + sorted(server_ids), + identifier, + ) + + def _warn_legacy_delegate_auth_if_applicable(server: MCPServer, *, source: str) -> None: """Direct legacy delegated OAuth configurations to the admitted replacement.""" if server.auth_type != MCPAuth.oauth2: @@ -6613,6 +6642,7 @@ class MCPServerManager: if previous_registry.get(server_id) != registered_registry.get(server_id): self._invalidate_discovery_lists(server_id) self.registry = registered_registry + _warn_on_shared_identifier_prefixes(registered_registry.values()) # A discovery task may have published into ``previous_registry`` while # this replacement was being staged. Reconcile every published entry # synchronously after the swap so a lost publication cannot also leave diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index c22081ebb3a..aa218f42023 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -27,6 +27,7 @@ from typing import ( Annotated, Final, Literal, + NoReturn, Protocol, cast, # noqa: TID251 # validated JSON values need explicit narrowing ) @@ -137,9 +138,10 @@ if MCP_AVAILABLE: return _ToolNameValidationResult() from litellm.proxy._experimental.mcp_server.db import ( + McpIdentifierConflict, approve_mcp_server, create_draft_mcp_server, - create_mcp_server, + create_mcp_server_if_identifier_free, delete_mcp_server, delete_user_credential, delete_user_env_vars, @@ -288,6 +290,21 @@ if MCP_AVAILABLE: _validate_mcp_server_name_fields(payload) _validate_upstream_token_header(payload) + def mcp_identifier_conflict_message(conflict: McpIdentifierConflict) -> str: + return ( + f"An MCP server with {conflict.field} '{conflict.value}' already exists " + f"(server_id={conflict.server_id}). " + "MCP server names and aliases must be unique, case-insensitive." + ) + + def raise_mcp_identifier_conflict(conflict: McpIdentifierConflict) -> NoReturn: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={ # mutable-ok: FastAPI HTTPException detail requires a plain dict + "error": mcp_identifier_conflict_message(conflict) + }, + ) + def warn_if_id_jag_server_outruns_sso(server_id: str | None, auth_type: MCPAuth | str | None) -> None: """Registering an ``oauth2_id_jag`` server under an SSO provider that captures no IdP identity assertion is a dead configuration: nothing here fails, and then every ID-JAG call @@ -1388,7 +1405,7 @@ if MCP_AVAILABLE: payload.submitted_at = datetime.now(timezone.utc) try: - new_mcp_server: Final = await create_mcp_server( + new_mcp_server: Final = await create_mcp_server_if_identifier_free( prisma_client, payload, touched_by=user_api_key_dict.user_id or user_api_key_dict.team_id, @@ -1399,6 +1416,8 @@ if MCP_AVAILABLE: status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail={"error": f"Error registering mcp server: {e}"}, ) + if isinstance(new_mcp_server, McpIdentifierConflict): + raise_mcp_identifier_conflict(new_mcp_server) # Do NOT add to runtime registry — pending servers are not active return _redact_mcp_credentials(new_mcp_server) @@ -1749,7 +1768,7 @@ if MCP_AVAILABLE: # The database write is the commit point: if it fails nothing was # persisted and the request is a genuine failure. try: - new_mcp_server: Final = await create_mcp_server( + new_mcp_server: Final = await create_mcp_server_if_identifier_free( prisma_client, payload, touched_by=user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME, @@ -1760,6 +1779,8 @@ if MCP_AVAILABLE: status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail={"error": f"Error creating mcp server: {e}"}, ) + if isinstance(new_mcp_server, McpIdentifierConflict): + raise_mcp_identifier_conflict(new_mcp_server) warn_if_id_jag_server_outruns_sso(new_mcp_server.server_id, new_mcp_server.auth_type) @@ -1808,7 +1829,7 @@ if MCP_AVAILABLE: conversions: Final = convert_connector_entries(payload) existing_servers: Final = await get_all_mcp_servers(prisma_client) existing_names: Final = frozenset( - name for server in existing_servers for name in (server.alias, server.server_name) if name + name.lower() for server in existing_servers for name in (server.alias, server.server_name) if name ) def _classify( @@ -1817,16 +1838,16 @@ if MCP_AVAILABLE: if isinstance(conversion, ConnectorConversionError): return conversion alias: Final = conversion.request.alias or "" - if alias in existing_names: + if alias.lower() in existing_names: return MCPConnectorImportSkipped( name=conversion.name, reason=f"An MCP server named '{alias}' already exists." ) earlier_aliases: Final = frozenset( - earlier.request.alias or "" + (earlier.request.alias or "").lower() for earlier in conversions[:index] if isinstance(earlier, ConvertedConnector) ) - if alias in earlier_aliases: + if alias.lower() in earlier_aliases: return MCPConnectorImportSkipped( name=conversion.name, reason=f"Duplicate connector name '{alias}' in the import payload." ) @@ -1834,7 +1855,7 @@ if MCP_AVAILABLE: async def _create( conversion: ConvertedConnector, - ) -> MCPConnectorImportResult | MCPConnectorImportFailure: + ) -> MCPConnectorImportResult | MCPConnectorImportFailure | MCPConnectorImportSkipped: try: validate_and_normalize_mcp_server_payload(conversion.request) except HTTPException as e: @@ -1843,7 +1864,7 @@ if MCP_AVAILABLE: ) return MCPConnectorImportFailure(name=conversion.name, error=error_text) try: - created: Final = await create_mcp_server( + created: Final = await create_mcp_server_if_identifier_free( prisma_client, conversion.request, touched_by=user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME, @@ -1851,6 +1872,8 @@ if MCP_AVAILABLE: except Exception as e: # noqa: BLE001 # any create failure must become a per-entry error, not a 500 verbose_proxy_logger.exception("Error importing mcp server %s: %s", conversion.name, e) return MCPConnectorImportFailure(name=conversion.name, error=str(e)) + if isinstance(created, McpIdentifierConflict): + return MCPConnectorImportSkipped(name=conversion.name, reason=mcp_identifier_conflict_message(created)) try: await global_mcp_server_manager.add_server(created) except Exception as e: # noqa: BLE001 # the row is committed; the reload after the loop retries registration @@ -2927,6 +2950,9 @@ if MCP_AVAILABLE: fields_set=payload_fields_set, ) + if isinstance(mcp_server_record_updated, McpIdentifierConflict): + raise_mcp_identifier_conflict(mcp_server_record_updated) + if mcp_server_record_updated is None: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, diff --git a/tests/integration/mcp/test_mcp_management.py b/tests/integration/mcp/test_mcp_management.py index 917acb9a1dc..bde18840d7d 100644 --- a/tests/integration/mcp/test_mcp_management.py +++ b/tests/integration/mcp/test_mcp_management.py @@ -2,7 +2,6 @@ import uuid from pathlib import Path from typing import Final -import pytest import yaml from integration._support.client import Gateway, eventually from integration._support.mcp import ( @@ -118,16 +117,67 @@ def test_delete_removes_listing_calls_and_database_row(gateway: Gateway) -> None def test_duplicate_alias_is_rejected_so_tool_prefixes_cannot_collide(gateway: Gateway) -> None: + import concurrent.futures + with mcp_peer() as peer, gateway.scenario() as scenario: alias: Final = "mgmt" + uuid.uuid4().hex[:8] - register_mcp(scenario, peer, alias) + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + duplicate: Final = gateway.request( "POST", "/v1/mcp/server", {"server_name": alias, "alias": alias, **peer.registration()} ) - if duplicate.status_code == 201: - scenario.cleanups.callback(forget_mcp, gateway, duplicate.json()["server_id"]) - pytest.skip("BUG: POST /v1/mcp/server accepts a duplicate alias, so two servers share one tool prefix") assert duplicate.status_code == 400, duplicate.text + assert alias in duplicate.json()["detail"]["error"], duplicate.text + + same_alias: Final = gateway.request( + "POST", "/v1/mcp/server", {"server_name": alias + "other", "alias": alias, **peer.registration()} + ) + assert same_alias.status_code == 400, same_alias.text + assert alias in same_alias.json()["detail"]["error"], same_alias.text + + case_variant: Final = gateway.request( + "POST", "/v1/mcp/server", {"server_name": alias.upper(), "alias": alias.upper(), **peer.registration()} + ) + assert case_variant.status_code == 400, case_variant.text + + same_name_no_alias: Final = gateway.request( + "POST", "/v1/mcp/server", {"server_name": alias, **peer.registration()} + ) + assert same_name_no_alias.status_code == 400, same_name_no_alias.text + + second_alias: Final = alias + "2" + second_identity: Final = register_mcp(scenario, peer, second_alias) + colliding_rename: Final = gateway.request( + "PUT", "/v1/mcp/server", {"server_id": second_identity, "alias": alias} + ) + assert colliding_rename.status_code == 400, colliding_rename.text + + cleared_alias: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": second_identity, "alias": None}) + assert cleared_alias.status_code == 202, cleared_alias.text + + name: Final = tool_names(gateway, key, identity)["add"] + response: Final = call_tool(gateway, key, identity, name, ADD) + assert response.status_code == 200, response.text + assert response.json()["content"][0]["text"] == "9", response.text + + racing_alias: Final = "race" + uuid.uuid4().hex[:8] + + def try_register() -> int: + response: Final = gateway.request( + "POST", "/v1/mcp/server", {"server_name": racing_alias, "alias": racing_alias, **peer.registration()} + ) + return response.status_code + + with concurrent.futures.ThreadPoolExecutor(max_workers=8) as pool: + statuses: Final = tuple(pool.map(lambda _i: try_register(), range(8))) + + assert statuses.count(201) == 1, statuses + assert statuses.count(400) == 7, statuses + winner: Final = next( + server["server_id"] for server in _servers(gateway).values() if server["alias"] == racing_alias + ) + scenario.cleanups.callback(forget_mcp, gateway, winner) def test_invalid_registrations_are_rejected(gateway: Gateway) -> None: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py index 669e094fee4..bd37a976286 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py @@ -29,11 +29,18 @@ def _credentials_cleared(value) -> bool: def _mock_prisma(): mock_prisma = MagicMock() mock_prisma.db.litellm_mcpservertable = AsyncMock() - row = models.LiteLLM_MCPServerTable.model_construct( - server_id="test-server", transport="http", env={}, env_vars=[] - ) + row = models.LiteLLM_MCPServerTable.model_construct(server_id="test-server", transport="http", env={}, env_vars=[]) mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=row) mock_prisma.db.litellm_mcpservertable.create = AsyncMock(return_value=row) + mock_prisma.db.litellm_mcpservertable.find_first = AsyncMock(return_value=None) + mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=None) + tx_client = MagicMock() + tx_client.execute_raw = AsyncMock() + tx_client.litellm_mcpservertable = mock_prisma.db.litellm_mcpservertable + tx = MagicMock() + tx.__aenter__ = AsyncMock(return_value=tx_client) + tx.__aexit__ = AsyncMock(return_value=False) + mock_prisma.db.tx = MagicMock(return_value=tx) return mock_prisma @@ -917,3 +924,170 @@ async def test_toolset_partial_update_ignores_a_null_name(): assert await _run_toolset_update({"toolset_id": "ts-1", "toolset_name": None, "description": "kept"}) == { "description": "kept" } + + +def _conflict_row(server_id: str = "other-server"): + return models.LiteLLM_MCPServerTable.model_construct( + server_id=server_id, server_name="taken", alias="taken", transport="http", env={}, env_vars=[] + ) + + +@pytest.mark.asyncio +async def test_find_identifier_conflict_reports_alias_hit(): + """A stored row matching the incoming alias yields a conflict naming it. + + Case-insensitive and cross-field matching is exercised end to end against + real Postgres by test_duplicate_alias_is_rejected_so_tool_prefixes_cannot_collide. + """ + from litellm.proxy._experimental.mcp_server.db import ( + find_mcp_server_identifier_conflict, + ) + + mock_prisma = _mock_prisma() + mock_prisma.db.litellm_mcpservertable.find_first = AsyncMock(return_value=_conflict_row()) + + conflict = await find_mcp_server_identifier_conflict( + mock_prisma, server_name="new-name", alias="taken", exclude_server_id="my-server" + ) + + assert conflict is not None + assert conflict.field == "alias" + assert conflict.value == "taken" + assert conflict.server_id == "other-server" + + +@pytest.mark.asyncio +async def test_find_identifier_conflict_reports_server_name_when_alias_is_free(): + """alias is checked first so the reported field is deterministic; a clean + alias does not mask a colliding server_name.""" + from litellm.proxy._experimental.mcp_server.db import ( + find_mcp_server_identifier_conflict, + ) + + mock_prisma = _mock_prisma() + mock_prisma.db.litellm_mcpservertable.find_first = AsyncMock(side_effect=[None, _conflict_row()]) + + conflict = await find_mcp_server_identifier_conflict( + mock_prisma, server_name="taken", alias="free", exclude_server_id=None + ) + + assert conflict is not None + assert conflict.field == "server_name" + + +@pytest.mark.asyncio +async def test_find_identifier_conflict_returns_none_when_free(): + from litellm.proxy._experimental.mcp_server.db import ( + find_mcp_server_identifier_conflict, + ) + + conflict = await find_mcp_server_identifier_conflict( + _mock_prisma(), server_name="fresh", alias="fresh", exclude_server_id=None + ) + + assert conflict is None + + +@pytest.mark.asyncio +async def test_update_writing_alias_returns_conflict_instead_of_row(): + from litellm.proxy._experimental.mcp_server.db import McpIdentifierConflict + + mock_prisma = _mock_prisma() + mock_prisma.db.litellm_mcpservertable.find_first = AsyncMock(return_value=_conflict_row()) + + result = await update_mcp_server( + mock_prisma, + UpdateMCPServerRequest(server_id="my-test-server", alias="taken"), + "test-user", + ) + + assert isinstance(result, McpIdentifierConflict) + + +@pytest.mark.asyncio +async def test_update_without_identifier_fields_returns_the_row(): + mock_prisma = _mock_prisma() + + result = await update_mcp_server( + mock_prisma, + UpdateMCPServerRequest(server_id="my-test-server", allowed_tools=["foo"]), + "test-user", + ) + + assert result is not None + + +@pytest.mark.asyncio +async def test_update_writing_free_alias_returns_the_row(): + mock_prisma = _mock_prisma() + + result = await update_mcp_server( + mock_prisma, + UpdateMCPServerRequest(server_id="my-test-server", alias="fresh-alias"), + "test-user", + ) + + assert result is not None + + +@pytest.mark.asyncio +async def test_clearing_alias_conflicts_on_the_fallback_server_name(): + """alias: null drops the tool prefix to the stored server_name, which may + already belong to another row, so that name goes through the conflict check.""" + from litellm.proxy._experimental.mcp_server.db import McpIdentifierConflict + + mock_prisma = _mock_prisma() + existing = MagicMock() + existing.server_name = "taken" + mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.db.litellm_mcpservertable.find_first = AsyncMock(return_value=_conflict_row()) + + result = await update_mcp_server( + mock_prisma, + UpdateMCPServerRequest(server_id="my-test-server", alias=None), + "test-user", + fields_set={"server_id", "alias"}, + ) + + assert isinstance(result, McpIdentifierConflict) + assert result.field == "server_name" + + +@pytest.mark.asyncio +async def test_clearing_alias_to_empty_string_conflicts_on_the_fallback_server_name(): + """alias: "" publishes the stored server_name as the tool prefix, just like + alias: null, so the fallback name must go through the conflict check too.""" + from litellm.proxy._experimental.mcp_server.db import McpIdentifierConflict + + mock_prisma = _mock_prisma() + existing = MagicMock() + existing.server_name = "taken" + mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.db.litellm_mcpservertable.find_first = AsyncMock(return_value=_conflict_row()) + + result = await update_mcp_server( + mock_prisma, + UpdateMCPServerRequest(server_id="my-test-server", alias=""), + "test-user", + fields_set={"server_id", "alias"}, + ) + + assert isinstance(result, McpIdentifierConflict) + assert result.field == "server_name" + + +@pytest.mark.asyncio +async def test_clearing_alias_with_free_server_name_returns_the_row(): + mock_prisma = _mock_prisma() + existing = MagicMock() + existing.server_name = "free-name" + mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + + result = await update_mcp_server( + mock_prisma, + UpdateMCPServerRequest(server_id="my-test-server", alias=None), + "test-user", + fields_set={"server_id", "alias"}, + ) + + assert result is not None diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 7725aca1948..cd5dae1269a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -14516,3 +14516,62 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie assert captured["client_ip"] is None finally: auth_context_var.reset(token) + + +class TestSharedIdentifierPrefixWarning: + """Two stored rows sharing lowercased alias-or-server_name publish one tool + prefix; reload must surface them once so the ambiguity is visible.""" + + @pytest.mark.asyncio + async def test_reload_warns_once_per_shared_identifier(self, caplog): + manager = MCPServerManager() + rows = [ + LiteLLM_MCPServerTable( + server_id="srv-a", server_name="alpha", alias="shared", url="https://a.example.com/mcp", + transport=MCPTransport.http, updated_at=datetime.now(), + ), + LiteLLM_MCPServerTable( + server_id="srv-b", server_name="beta", alias="Shared", url="https://b.example.com/mcp", + transport=MCPTransport.http, updated_at=datetime.now(), + ), + LiteLLM_MCPServerTable( + server_id="srv-c", server_name="gamma", alias="lonely", url="https://c.example.com/mcp", + transport=MCPTransport.http, updated_at=datetime.now(), + ), + ] + raw_rows = [MagicMock(model_dump=lambda row=row: row.model_dump()) for row in rows] + repository = MagicMock() + repository.table.find_many = AsyncMock(return_value=raw_rows) + + async def build_from_table(table, **_kwargs): + return MCPServer( + server_id=table.server_id, + name=table.alias or table.server_name, + alias=table.alias, + server_name=table.server_name, + url=table.url, + transport=table.transport, + ) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", + return_value=repository, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + patch.object(manager, "build_mcp_server_from_table", new=build_from_table), + patch.object(manager, "_maybe_register_openapi_tools", new=AsyncMock()), + patch.object(manager, "_prime_oauth_metadata_discovery_for_servers"), + caplog.at_level(logging.WARNING, logger="LiteLLM"), + ): + await manager.reload_servers_from_database() + + shared_warnings = [m for m in caplog.messages if "share the identifier" in m] + assert len(shared_warnings) == 1 + assert "srv-a" in shared_warnings[0] + assert "srv-b" in shared_warnings[0] + assert "srv-c" not in shared_warnings[0] + assert "'shared'" in shared_warnings[0] diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 53645e62034..557e753a76f 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -3936,7 +3936,7 @@ class TestAddMCPServerAtomicity: MagicMock(), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server", + "litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server_if_identifier_free", AsyncMock(return_value=created_server), ) as create_mock, patch( @@ -3977,7 +3977,7 @@ class TestAddMCPServerAtomicity: MagicMock(), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server", + "litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server_if_identifier_free", AsyncMock(side_effect=Exception("db down")), ), patch( @@ -4043,7 +4043,7 @@ class TestIdJagRegistrationWarnsAboutTheSSOGap: return_value=MagicMock(), ), patch( # test-quality-ok: endpoint test stubs MCP server creation - "litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server", + "litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server_if_identifier_free", AsyncMock(return_value=self._server_record(auth_type)), ), patch( # test-quality-ok: endpoint reads the global MCP manager @@ -4592,7 +4592,7 @@ class TestMCPApprovalWorkflow: MagicMock(), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server", + "litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server_if_identifier_free", AsyncMock(return_value=created_record), ) as mock_create, ): @@ -7532,7 +7532,7 @@ class TestImportMCPServers: AsyncMock(return_value=existing_servers), ), patch( # test-quality-ok: endpoint takes collaborators from module scope, matching the suite's pattern - "litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server", + "litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server_if_identifier_free", create_mock, ), patch( # test-quality-ok: endpoint takes collaborators from module scope, matching the suite's pattern @@ -7940,3 +7940,229 @@ async def test_config_server_edit_preserves_api_contract_without_creating_rows(r prisma.tx.assert_not_called() assert server.model_dump() == original assert manager.registry == {} + + +class TestDuplicateIdentifierRejection: + """server_name/alias must be unique across live servers, case-insensitive. + + The DB layer returns McpIdentifierConflict instead of writing; every write + path maps it to a 400 naming the colliding identifier, so a second server + can never share another server's tool prefix. + """ + + @staticmethod + def _conflict(field: str, value: str, server_id: str = "existing-1"): + from litellm.proxy._experimental.mcp_server.db import McpIdentifierConflict + + return McpIdentifierConflict(field=field, value=value, server_id=server_id) + + @pytest.mark.asyncio + async def test_create_conflict_returns_400_naming_the_alias(self): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + add_mcp_server, + ) + + payload = NewMCPServerRequest( + alias="echo", + url="https://echo.example.com/mcp", + transport=MCPTransport.http, + ) + admin = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user") + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.validate_and_normalize_mcp_server_payload", + MagicMock(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server_if_identifier_free", + AsyncMock(return_value=self._conflict("alias", "echo")), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + MagicMock(), + ), + ): + with pytest.raises(HTTPException) as exc_info: + await add_mcp_server(payload=payload, user_api_key_dict=admin) + + assert exc_info.value.status_code == 400 + assert "echo" in exc_info.value.detail["error"] + assert "existing-1" in exc_info.value.detail["error"] + + @pytest.mark.asyncio + async def test_submission_conflict_returns_400(self): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + register_mcp_server, + ) + + payload = NewMCPServerRequest( + alias="echo", + url="https://echo.example.com/mcp", + transport=MCPTransport.http, + ) + team_member = generate_mock_user_api_key_auth( + user_role=LitellmUserRoles.INTERNAL_USER, user_id="member", team_id="team-1" + ) + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.validate_and_normalize_mcp_server_payload", + MagicMock(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server_if_identifier_free", + AsyncMock(return_value=self._conflict("server_name", "echo")), + ), + ): + with pytest.raises(HTTPException) as exc_info: + await register_mcp_server(payload=payload, user_api_key_dict=team_member) + + assert exc_info.value.status_code == 400 + assert "echo" in exc_info.value.detail["error"] + + @pytest.mark.asyncio + async def test_edit_conflict_returns_400(self): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + edit_mcp_server, + ) + + admin = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user") + existing = generate_mock_mcp_server_db_record(server_id="edit-1", alias="first") + + mock_manager = MagicMock() + mock_manager.update_server = AsyncMock() + mock_manager.reload_servers_from_database = AsyncMock() + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.validate_and_normalize_mcp_server_payload", + MagicMock(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server", + AsyncMock(return_value=existing), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.update_mcp_server", + AsyncMock(return_value=self._conflict("alias", "taken", server_id="other-1")), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + mock_manager, + ), + ): + with pytest.raises(HTTPException) as exc_info: + await edit_mcp_server( + payload=UpdateMCPServerRequest(server_id="edit-1", alias="taken"), + user_api_key_dict=admin, + ) + + assert exc_info.value.status_code == 400 + assert "taken" in exc_info.value.detail["error"] + assert "other-1" in exc_info.value.detail["error"] + mock_manager.update_server.assert_not_awaited() + + @pytest.mark.asyncio + async def test_edit_rename_to_free_alias_succeeds(self): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + edit_mcp_server, + ) + + admin = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user") + existing = generate_mock_mcp_server_db_record(server_id="edit-1", alias="first") + updated = generate_mock_mcp_server_db_record(server_id="edit-1", alias="renamed") + + mock_manager = MagicMock() + mock_manager.update_server = AsyncMock() + mock_manager.reload_servers_from_database = AsyncMock() + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.validate_and_normalize_mcp_server_payload", + MagicMock(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server", + AsyncMock(return_value=existing), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.update_mcp_server", + AsyncMock(return_value=updated), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + mock_manager, + ), + ): + result = await edit_mcp_server( + payload=UpdateMCPServerRequest(server_id="edit-1", alias="renamed"), + user_api_key_dict=admin, + ) + + assert result.alias == "renamed" + mock_manager.update_server.assert_awaited_once_with(updated) + + @pytest.mark.asyncio + async def test_import_skips_case_variant_duplicate(self): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + MCPConnectorImportRequest, + import_mcp_servers, + ) + + payload = MCPConnectorImportRequest.model_validate( + {"mcpServers": {"EXISTING": {"url": "https://dup.example/mcp"}}} + ) + admin = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user") + existing = generate_mock_mcp_server_db_record(server_id="existing-1", alias="existing") + create_mock = AsyncMock() + mock_manager = MagicMock() + + with ExitStack() as stack: + for p in TestImportMCPServers._import_patches([existing], create_mock, mock_manager): + stack.enter_context(p) + result = await import_mcp_servers(payload=payload, user_api_key_dict=admin) + + assert [entry.name for entry in result.skipped] == ["EXISTING"] + assert "already exists" in result.skipped[0].reason + create_mock.assert_not_awaited() + + @pytest.mark.asyncio + async def test_import_skips_db_reported_identifier_conflict(self): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + MCPConnectorImportRequest, + import_mcp_servers, + ) + + payload = MCPConnectorImportRequest.model_validate( + {"mcpServers": {"fresh": {"url": "https://dup.example/mcp"}}} + ) + admin = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user") + existing = generate_mock_mcp_server_db_record(server_id="existing-1", alias="existing") + create_mock = AsyncMock(return_value=self._conflict("alias", "fresh", server_id="other-9")) + mock_manager = MagicMock() + + with ExitStack() as stack: + for p in TestImportMCPServers._import_patches([existing], create_mock, mock_manager): + stack.enter_context(p) + result = await import_mcp_servers(payload=payload, user_api_key_dict=admin) + + assert [entry.name for entry in result.skipped] == ["fresh"] + assert "fresh" in result.skipped[0].reason + assert result.imported == () diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx index ebca780c766..f656bd2fc60 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx @@ -35,6 +35,7 @@ import { buildCreateServerPayload, reduceStaticHeaders, } from "./createServerPayload"; +import { DUPLICATE_IDENTIFIER_MESSAGE, findDuplicateMcpServer, mcpSubmitErrorReason } from "./duplicateServerCheck"; import { readCreateUiSnapshot, writeCreateUiSnapshot } from "./createOAuthUiState"; import AwsSigV4Fields from "./AwsSigV4Fields"; import OpenApiByokFields from "./OpenApiByokFields"; @@ -78,6 +79,7 @@ interface CreateMCPServerProps { isModalVisible: boolean; setModalVisible: (visible: boolean) => void; availableAccessGroups: string[]; + existingServers?: MCPServer[]; prefillData?: DiscoverableMCPServer | null; onBackToDiscovery?: () => void; } @@ -108,6 +110,7 @@ const CreateMCPServer: React.FC = ({ isModalVisible, setModalVisible, availableAccessGroups, + existingServers, prefillData, onBackToDiscovery, }) => { @@ -418,6 +421,16 @@ const CreateMCPServer: React.FC = ({ }; const handleCreate = async (values: Record) => { + const duplicate = findDuplicateMcpServer( + existingServers, + typeof values.server_name === "string" ? values.server_name : undefined, + typeof values.alias === "string" ? values.alias : undefined, + ); + if (duplicate) { + form.setError(duplicate.field, { type: "duplicate", message: DUPLICATE_IDENTIFIER_MESSAGE }); + toast.fromError(DUPLICATE_IDENTIFIER_MESSAGE); + return; + } const built = buildCreateServerPayload(values, { transportType, costConfig, @@ -488,7 +501,7 @@ const CreateMCPServer: React.FC = ({ onCreateSuccess(response); } } catch (error) { - const reason = error instanceof Error ? error.message : String(error); + const reason = mcpSubmitErrorReason(error); toast.fromError(isAdmin ? `Error creating MCP Server: ${reason}` : `Error submitting MCP Server: ${reason}`); } finally { setIsLoading(false); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/duplicateServerCheck.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/duplicateServerCheck.test.ts new file mode 100644 index 00000000000..8590e40229c --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/duplicateServerCheck.test.ts @@ -0,0 +1,49 @@ +import { describe, expect, it } from "vitest"; +import { ApiError } from "@/lib/http/client"; +import { findDuplicateMcpServer, mcpSubmitErrorReason } from "./duplicateServerCheck"; + +const servers = [ + { server_id: "s1", server_name: "GitHub_MCP", alias: "github" }, + { server_id: "s2", server_name: "Email Service", alias: "email_service" }, +]; + +describe("findDuplicateMcpServer", () => { + it("flags an incoming server_name that matches an existing alias", () => { + expect(findDuplicateMcpServer(servers, "github", "other")?.field).toBe("server_name"); + }); + + it("flags an incoming alias that matches an existing server_name", () => { + expect(findDuplicateMcpServer(servers, "new", "GitHub_MCP")?.serverId).toBe("s1"); + }); + + it("matches case-insensitively", () => { + expect(findDuplicateMcpServer(servers, "GITHUB", "new")?.serverId).toBe("s1"); + }); + + it("normalizes spaces to underscores like the backend does", () => { + expect(findDuplicateMcpServer(servers, "new", "email service")?.serverId).toBe("s2"); + }); + + it("does not flag the server's own identifiers while editing", () => { + expect(findDuplicateMcpServer(servers, "GitHub_MCP", "github", "s1")).toBeNull(); + }); + + it("flags the same alias on a different server while editing", () => { + expect(findDuplicateMcpServer(servers, "other", "github", "s2")?.serverId).toBe("s1"); + }); + + it("returns null when nothing matches", () => { + expect(findDuplicateMcpServer(servers, "brand_new", "brand_new")).toBeNull(); + }); +}); + +describe("mcpSubmitErrorReason", () => { + it("unwraps the FastAPI detail.error envelope into readable toast text", () => { + const error = new ApiError("boom", 400, { detail: { error: "An MCP server with alias 'x' already exists" } }); + expect(mcpSubmitErrorReason(error)).toContain("An MCP server with alias 'x' already exists"); + }); + + it("never produces [object Object] for a non-Error rejection", () => { + expect(mcpSubmitErrorReason({ detail: { error: "structured 400" } })).toBe("structured 400"); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/duplicateServerCheck.ts b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/duplicateServerCheck.ts new file mode 100644 index 00000000000..0e24e13e7f1 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/duplicateServerCheck.ts @@ -0,0 +1,48 @@ +import { MCPServer } from "@/components/mcp_tools/types"; +import { ApiError, deriveErrorMessage, unwrapProxyErrorMessage } from "@/lib/http/client"; + +export type McpIdentifierField = "server_name" | "alias"; + +export interface McpIdentifierDuplicate { + field: McpIdentifierField; + serverId: string; +} + +export const normalizeMcpIdentifier = (value: string | null | undefined): string => + (value ?? "").trim().replace(/\s+/g, "_").toLowerCase(); + +export function findDuplicateMcpServer( + servers: readonly Pick[] | undefined, + serverName: string | null | undefined, + alias: string | null | undefined, + excludeServerId?: string, +): McpIdentifierDuplicate | null { + const candidates: ReadonlyArray = [ + ["alias", alias], + ["server_name", serverName], + ]; + for (const [field, value] of candidates) { + const normalized = normalizeMcpIdentifier(value); + if (!normalized) { + continue; + } + const hit = (servers ?? []).find( + (server) => + server.server_id !== excludeServerId && + [server.server_name, server.alias].some((existing) => normalizeMcpIdentifier(existing) === normalized), + ); + if (hit) { + return { field, serverId: hit.server_id }; + } + } + return null; +} + +export const DUPLICATE_IDENTIFIER_MESSAGE = "An MCP server with this name/alias already exists."; + +export const mcpSubmitErrorReason = (error: unknown): string => { + if (error instanceof ApiError) { + return deriveErrorMessage(error.body); + } + return error instanceof Error ? unwrapProxyErrorMessage(error.message) : deriveErrorMessage(error); +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx index 2a37029a2c4..d45909445e9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx @@ -51,6 +51,7 @@ import MCPLogoSelector from "./MCPLogoSelector"; import EnvVarsSection from "./EnvVarsSection"; import { validateMCPServerUrl, validateMCPServerName, normalizeToolOverrideMap } from "./utils"; import { EditServerFormValues, buildEditServerPayload, editPayloadErrorMessage } from "./editServerPayload"; +import { DUPLICATE_IDENTIFIER_MESSAGE, findDuplicateMcpServer, mcpSubmitErrorReason } from "./duplicateServerCheck"; import { toast } from "@/lib/toast"; import { getEditToolPreview } from "./editToolPreview"; import { useMcpOAuthFlow } from "@/hooks/useMcpOAuthFlow"; @@ -88,6 +89,7 @@ interface MCPServerEditProps { onCancel: () => void; onSuccess: (server: MCPServer) => void; availableAccessGroups: string[]; + existingServers?: MCPServer[]; } const AUTH_TYPES_REQUIRING_AUTH_VALUE = [AUTH_TYPE.API_KEY, AUTH_TYPE.BEARER_TOKEN, AUTH_TYPE.TOKEN, AUTH_TYPE.BASIC]; @@ -100,6 +102,7 @@ const MCPServerEdit: React.FC = ({ onCancel, onSuccess, availableAccessGroups, + existingServers, }) => { const initialStaticHeaders = React.useMemo(() => { if (!mcpServer.static_headers) { @@ -724,6 +727,17 @@ const MCPServerEdit: React.FC = ({ const handleSave = async (values: EditServerFormValues) => { if (!accessToken) return; + const duplicate = findDuplicateMcpServer( + existingServers, + values.server_name || mcpServer.server_name, + (values.alias ?? mcpServer.alias) || null, + mcpServer.server_id, + ); + if (duplicate) { + form.setError(duplicate.field, { type: "duplicate", message: DUPLICATE_IDENTIFIER_MESSAGE }); + toast.fromError(DUPLICATE_IDENTIFIER_MESSAGE); + return; + } try { const built = buildEditServerPayload(values, { mcpServer, @@ -783,7 +797,8 @@ const MCPServerEdit: React.FC = ({ setAppMayNotMatchUpstream(false); onSuccess(updated); } catch (error: any) { - toast.fromError("Failed to update MCP Server" + (error?.message ? `: ${error.message}` : "")); + const reason = mcpSubmitErrorReason(error); + toast.fromError("Failed to update MCP Server" + (reason ? `: ${reason}` : "")); } }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx index 6045be4607d..a7ff34301a0 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx @@ -27,6 +27,7 @@ interface MCPServerViewProps { userID: string | null; isViewOnly?: boolean; availableAccessGroups: string[]; + existingServers?: MCPServer[]; initialTabIndex?: number; } @@ -58,6 +59,7 @@ export const MCPServerView: React.FC = ({ userID, isViewOnly = false, availableAccessGroups, + existingServers, initialTabIndex = 0, }) => { // Open the editing Settings tab on first render when returning from the edit OAuth @@ -244,6 +246,7 @@ export const MCPServerView: React.FC = ({ onCancel={() => setEditing(false)} onSuccess={handleSuccess} availableAccessGroups={availableAccessGroups} + existingServers={existingServers} /> ) : (
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx index 738409e28c2..b4b7ab6b3c8 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx @@ -497,6 +497,7 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID, i isModalVisible={isModalVisible} setModalVisible={setModalVisible} availableAccessGroups={uniqueMcpAccessGroups} + existingServers={mcpServers} prefillData={prefillData} onBackToDiscovery={() => { setModalVisible(false); @@ -610,6 +611,7 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID, i userRole={userRole} isViewOnly={isViewOnly} availableAccessGroups={uniqueMcpAccessGroups} + existingServers={mcpServers} initialTabIndex={selectedServerId === toolsTabServerId ? 1 : 0} /> ) : ( From 1e5403288ce0731c8bd1a2222a87f40692218915 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 17:51:16 -0700 Subject: [PATCH 10/10] feat(proxy): honor model_info.discoverable on the model listing endpoints (#42825) * feat(proxy): honor model_info.discoverable on the model listing endpoints A model_list entry marked model_info: {discoverable: false} is left out of GET /v1/models (OpenAI and Anthropic shapes, scope=expand and wildcard routes included), the list path of GET /v1/model/info and GET /model_group/info for every caller without the admin view, while direct requests naming the model keep routing to it. The field defaults to None so an absent flag reads as discoverable and nothing is persisted or echoed for configs that never set it. * fix(proxy): hide flagged team models under their public name and cover the scope=expand filter The discoverability lookup now resolves a listed name with the caller's team context, so a team-scoped deployment marked discoverable: false drops out for that team's keys under its public name instead of failing open. The scope=expand branch is now exercised by a team admin caller, and the OCI secrets test builds a real UserAPIKeyAuth instead of a spec mock that has no pydantic fields. * perf(proxy): resolve only candidate names in the discoverable filter --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/proxy/_types.py | 1 + .../common_utils/discoverable_model_filter.py | 118 +++++++++ litellm/proxy/proxy_server.py | 29 ++- litellm/types/router.py | 1 + .../test_discoverable_model_filter.py | 156 ++++++++++++ .../proxy/test_model_list_discoverable.py | 229 ++++++++++++++++++ tests/test_litellm/proxy/test_proxy_server.py | 9 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 4 + 8 files changed, 533 insertions(+), 14 deletions(-) create mode 100644 litellm/proxy/common_utils/discoverable_model_filter.py create mode 100644 tests/test_litellm/proxy/common_utils/test_discoverable_model_filter.py create mode 100644 tests/test_litellm/proxy/test_model_list_discoverable.py diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 54574ed64e3..4affa55f903 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1133,6 +1133,7 @@ class ModelInfo(LiteLLMPydanticObjectBase): ] | None ) + discoverable: bool | None = None model_config = ConfigDict(protected_namespaces=(), extra="allow") diff --git a/litellm/proxy/common_utils/discoverable_model_filter.py b/litellm/proxy/common_utils/discoverable_model_filter.py new file mode 100644 index 00000000000..d22f1a9f6dc --- /dev/null +++ b/litellm/proxy/common_utils/discoverable_model_filter.py @@ -0,0 +1,118 @@ +from __future__ import annotations + +import re +from collections.abc import Iterable, Mapping +from typing import TYPE_CHECKING, Final + +from pydantic import TypeAdapter + +from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider, get_llm_provider +from litellm.proxy._types import UserAPIKeyAuth, user_api_key_has_admin_view + +if TYPE_CHECKING: + from litellm.router import Router + from litellm.types.router import RouterModelGroupAliasItem + +_PATTERN_DEPLOYMENTS: Final = TypeAdapter(Mapping[str, tuple[Mapping[str, object], ...]]) + + +def is_undiscoverable_deployment(deployment: Mapping[str, object]) -> bool: + model_info: Final = deployment.get("model_info") + if not isinstance(model_info, Mapping): + return False + return "discoverable" in model_info and model_info["discoverable"] is False + + +def is_undiscoverable_model_name(model_name: str, llm_router: Router | None, team_id: str | None) -> bool: + if llm_router is None: + return False + deployments: Final = llm_router.get_model_list(model_name=model_name, team_id=team_id) + if not deployments: + return False + return all(is_undiscoverable_deployment(deployment) for deployment in deployments) + + +def _team_public_model_name(deployment: Mapping[str, object]) -> object: + model_info: Final = deployment.get("model_info") + return model_info.get("team_public_model_name") if isinstance(model_info, Mapping) else None + + +def _alias_target(alias: str | RouterModelGroupAliasItem) -> str: + return alias if isinstance(alias, str) else alias["model"] + + +def _undiscoverable_served_names( + undiscoverable_rows: Iterable[Mapping[str, object]], + model_group_alias: Mapping[str, str | RouterModelGroupAliasItem], +) -> frozenset[str]: + served: Final = frozenset( + name + for row in undiscoverable_rows + for name in (row.get("model_name"), _team_public_model_name(row)) + if isinstance(name, str) + ) + aliases: Final = frozenset(alias for alias, target in model_group_alias.items() if _alias_target(target) in served) + return served | aliases + + +def _undiscoverable_patterns(llm_router: Router, team_id: str | None) -> tuple[re.Pattern[str], ...]: + team_pattern_router: Final = llm_router.team_pattern_routers.get(team_id) if team_id is not None else None + pattern_routers: Final = ( + (llm_router.pattern_router,) + if team_pattern_router is None + else (llm_router.pattern_router, team_pattern_router) + ) + return tuple( + re.compile(regex) + for pattern_router in pattern_routers + for regex, deployments in _PATTERN_DEPLOYMENTS.validate_python(pattern_router.patterns).items() + if any(is_undiscoverable_deployment(deployment) for deployment in deployments) + ) + + +def _resolved_provider(model_name: str) -> str | None: + try: + return get_llm_provider(model=model_name)[1] + except Exception: # noqa: BLE001 # get_llm_provider raises when the provider is unknown; the name then routes as-is + return None + + +def _matches_undiscoverable_pattern(model_name: str, patterns: tuple[re.Pattern[str], ...]) -> bool: + if not patterns: + return False + if any(pattern.match(model_name) for pattern in patterns): + return True + provider: Final = declared_authenticating_provider(model_name) or _resolved_provider(model_name) + return any(pattern.match(f"{provider}/{model_name}") for pattern in patterns) + + +def undiscoverable_model_names( + model_names: Iterable[str], + llm_router: Router | None, + user_api_key_dict: UserAPIKeyAuth, + team_id: str | None, +) -> frozenset[str]: + if llm_router is None or user_api_key_has_admin_view(user_api_key_dict): + return frozenset() + undiscoverable_rows: Final = tuple( + row for row in llm_router.get_model_list() or () if is_undiscoverable_deployment(row) + ) + if not undiscoverable_rows: + return frozenset() + served_names: Final = _undiscoverable_served_names(undiscoverable_rows, llm_router.model_group_alias) + patterns: Final = _undiscoverable_patterns(llm_router, team_id) + return frozenset( + name + for name in model_names + if (name in served_names or _matches_undiscoverable_pattern(name, patterns)) + and is_undiscoverable_model_name(name, llm_router, team_id) + ) + + +def discoverable_rows( + rows: Iterable[Mapping[str, object]], + user_api_key_dict: UserAPIKeyAuth, +) -> tuple[Mapping[str, object], ...]: + if user_api_key_has_admin_view(user_api_key_dict): + return tuple(rows) + return tuple(row for row in rows if not is_undiscoverable_deployment(row)) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 26a423d7162..7e18db742cd 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -399,6 +399,7 @@ from litellm.proxy.common_utils.config_includes import resolve_include_file_path from litellm.proxy.common_utils.config_sync_pubsub import ConfigSyncSubscriber from litellm.proxy.common_utils.debug_utils import init_verbose_loggers from litellm.proxy.common_utils.debug_utils import router as debugging_endpoints_router +from litellm.proxy.common_utils.discoverable_model_filter import discoverable_rows, undiscoverable_model_names from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, @@ -11241,9 +11242,11 @@ async def model_list( only_model_access_groups=only_model_access_groups or False, ) - # Hide paused/unhealthy models from the public listing - if hidden_names: - all_models = [m for m in all_models if m not in hidden_names] + expanded_undiscoverable_names: Final = undiscoverable_model_names( + all_models, llm_router, user_api_key_dict, team_id or user_api_key_dict.team_id + ) + if hidden_names or expanded_undiscoverable_names: + all_models = [m for m in all_models if m not in hidden_names and m not in expanded_undiscoverable_names] # Surface the public team name by default; legacy internal keys via flag. # The internal routing key drives the metadata/fallback lookup, while the @@ -11294,9 +11297,11 @@ async def model_list( user_api_key_cache=user_api_key_cache, ) - # Hide paused/unhealthy models from the public listing - if hidden_names: - all_models = [m for m in all_models if m not in hidden_names] + undiscoverable_names: Final = undiscoverable_model_names( + all_models, llm_router, user_api_key_dict, team_id or user_api_key_dict.team_id + ) + if hidden_names or undiscoverable_names: + all_models = [m for m in all_models if m not in hidden_names and m not in undiscoverable_names] # Surface the public team name by default; legacy internal keys via flag. # The internal routing key drives the metadata/fallback lookup, while the @@ -15793,7 +15798,10 @@ async def model_info_v1( general_settings=general_settings, llm_router=llm_router, ) - visible_models: Final = [model for model in all_models if model.get("model_name") not in hidden_names] + visible_models: Final = discoverable_rows( + (model for model in all_models if model.get("model_name") not in hidden_names), + user_api_key_dict, + ) verbose_proxy_logger.debug("all_models: %s", visible_models) return _model_info_json_response(visible_models) @@ -16072,8 +16080,13 @@ async def model_group_info( user_api_key_cache=user_api_key_cache, ) ) + undiscoverable_group_names: Final = undiscoverable_model_names( + all_models_str, llm_router, user_api_key_dict, user_api_key_dict.team_id + ) model_groups: list[ModelGroupInfoProxy] = _get_model_group_info( - llm_router=llm_router, all_models_str=all_models_str, model_group=model_group + llm_router=llm_router, + all_models_str=[name for name in all_models_str if name not in undiscoverable_group_names], + model_group=model_group, ) # Append A2A agents to model groups diff --git a/litellm/types/router.py b/litellm/types/router.py index 57bd4263894..c0f724584fd 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -246,6 +246,7 @@ class ModelInfo(MirroredPricingParams): # admin-toggled pause flag; mirrors LiteLLM_ProxyModelTable.blocked blocked: bool | None = None + discoverable: bool | None = None access_windows: tuple[ModelAccessWindow, ...] | None = None diff --git a/tests/test_litellm/proxy/common_utils/test_discoverable_model_filter.py b/tests/test_litellm/proxy/common_utils/test_discoverable_model_filter.py new file mode 100644 index 00000000000..17619afbb07 --- /dev/null +++ b/tests/test_litellm/proxy/common_utils/test_discoverable_model_filter.py @@ -0,0 +1,156 @@ +""" +Tests for the operator-declared discoverability filter shared by the model +listing endpoints: a deployment marked `model_info: {discoverable: false}` is +hidden from listings for callers without the admin view while it still routes. +""" + +import pytest + +from litellm import Router +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.common_utils.discoverable_model_filter import ( + discoverable_rows, + undiscoverable_model_names, +) + + +def _deployment(model_name: str, model: str = "openai/gpt-4o", **model_info): + return { + "model_name": model_name, + "litellm_params": {"model": model, "api_key": "sk-fake"}, + "model_info": {"id": f"{model_name}-id", **model_info}, + } + + +def _router(*deployments, **router_kwargs) -> Router: + return Router(model_list=list(deployments), **router_kwargs) + + +def _non_admin() -> UserAPIKeyAuth: + return UserAPIKeyAuth(api_key="sk-test", user_role=LitellmUserRoles.INTERNAL_USER) + + +def _admin(role: LitellmUserRoles = LitellmUserRoles.PROXY_ADMIN) -> UserAPIKeyAuth: + return UserAPIKeyAuth(api_key="sk-test", user_role=role) + + +def test_flagged_model_is_undiscoverable_for_non_admin(): + router = _router(_deployment("gpt-4"), _deployment("internal-evaluator", discoverable=False)) + + assert undiscoverable_model_names(["gpt-4", "internal-evaluator"], router, _non_admin(), None) == { + "internal-evaluator" + } + + +def test_missing_flag_and_explicit_true_are_discoverable(): + router = _router(_deployment("gpt-4"), _deployment("public-eval", discoverable=True)) + + assert undiscoverable_model_names(["gpt-4", "public-eval"], router, _non_admin(), None) == frozenset() + + +@pytest.mark.parametrize("role", [LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY]) +def test_admin_view_sees_flagged_models(role): + router = _router(_deployment("internal-evaluator", discoverable=False)) + + assert undiscoverable_model_names(["internal-evaluator"], router, _admin(role), None) == frozenset() + + +def test_group_with_one_discoverable_deployment_stays_listed(): + router = _router( + _deployment("shared", discoverable=False), + { + "model_name": "shared", + "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-fake"}, + "model_info": {"id": "shared-public"}, + }, + ) + + assert undiscoverable_model_names(["shared"], router, _non_admin(), None) == frozenset() + + +def test_unknown_name_and_missing_router_fail_open(): + router = _router(_deployment("internal-evaluator", discoverable=False)) + + assert undiscoverable_model_names(["not-configured"], router, _non_admin(), None) == frozenset() + assert undiscoverable_model_names(["internal-evaluator"], None, _non_admin(), None) == frozenset() + + +def test_alias_follows_its_target_deployments(): + router = _router( + _deployment("gpt-4"), + _deployment("internal-evaluator", discoverable=False), + model_group_alias={"eval": "internal-evaluator", "chat": "gpt-4"}, + ) + + assert undiscoverable_model_names(["eval", "chat"], router, _non_admin(), None) == {"eval"} + + +def test_wildcard_expansions_follow_the_wildcard_entry(): + router = _router(_deployment("gpt-4"), _deployment("anthropic/*", model="anthropic/*", discoverable=False)) + + hidden = undiscoverable_model_names( + ["gpt-4", "anthropic/*", "anthropic/claude-opus-5"], router, _non_admin(), None + ) + + assert hidden == {"anthropic/*", "anthropic/claude-opus-5"} + + +def test_flagged_team_model_is_undiscoverable_for_its_team_member(): + router = _router( + _deployment("gpt-4"), + _deployment( + "model_name_team1_abc", team_id="team1", team_public_model_name="team-gpt", discoverable=False + ), + ) + member = UserAPIKeyAuth( + api_key="sk-test", user_role=LitellmUserRoles.INTERNAL_USER, team_id="team1", team_models=["team-gpt"] + ) + + assert undiscoverable_model_names(["gpt-4", "team-gpt"], router, member, "team1") == {"team-gpt"} + + +def test_hidden_model_still_routes_for_direct_requests(): + router = _router(_deployment("gpt-4"), _deployment("internal-evaluator", discoverable=False)) + + assert "internal-evaluator" in undiscoverable_model_names(["internal-evaluator"], router, _non_admin(), None) + deployment = router.get_available_deployment( + model="internal-evaluator", messages=[{"role": "user", "content": "hi"}] + ) + assert deployment["model_name"] == "internal-evaluator" + + +def test_discoverable_rows_drops_flagged_rows_only_for_non_admin(): + rows = [ + {"model_name": "gpt-4", "model_info": {"id": "a"}}, + {"model_name": "internal-evaluator", "model_info": {"id": "b", "discoverable": False}}, + {"model_name": "no-model-info"}, + ] + + assert [row["model_name"] for row in discoverable_rows(rows, _non_admin())] == ["gpt-4", "no-model-info"] + assert [row["model_name"] for row in discoverable_rows(rows, _admin())] == [ + "gpt-4", + "internal-evaluator", + "no-model-info", + ] + + +def test_expanded_name_served_by_a_discoverable_wildcard_too_stays_listed(): + router = _router( + _deployment("anthropic/*", model="anthropic/*", discoverable=False), + _deployment("anthropic/claude-*", model="anthropic/claude-*"), + ) + + hidden = undiscoverable_model_names( + ["anthropic/claude-opus-5", "anthropic/other-model"], router, _non_admin(), None + ) + + assert hidden == {"anthropic/other-model"} + + +def test_hidden_alias_of_a_flagged_model_is_undiscoverable(): + router = _router( + _deployment("internal-evaluator", discoverable=False), + model_group_alias={"eval": {"model": "internal-evaluator", "hidden": True}}, + ) + + assert undiscoverable_model_names(["eval"], router, _non_admin(), None) == {"eval"} diff --git a/tests/test_litellm/proxy/test_model_list_discoverable.py b/tests/test_litellm/proxy/test_model_list_discoverable.py new file mode 100644 index 00000000000..bcd52479f2c --- /dev/null +++ b/tests/test_litellm/proxy/test_model_list_discoverable.py @@ -0,0 +1,229 @@ +""" +Tests for `model_info.discoverable: false` on the model listing endpoints: +GET /v1/models (`model_list`, OpenAI and Anthropic shapes), GET /v1/models/{id} +(`model_info`), GET /v1/model/info (`model_info_v1`) and GET /model_group/info +(`model_group_info`). Flagged models drop out of the listings for callers without +the admin view and stay reachable by name. +""" + +import json + +import pytest +from starlette.requests import Request + +from litellm import Router +from litellm.proxy import proxy_server +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + + +def _deployment(model_name: str, model: str = "openai/gpt-4o", **model_info): + return { + "model_name": model_name, + "litellm_params": {"model": model, "api_key": "sk-fake"}, + "model_info": {"id": f"{model_name}-id", **model_info}, + } + + +def _install_router(monkeypatch, *deployments) -> Router: + router = Router(model_list=list(deployments)) + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(proxy_server, "llm_model_list", router.model_list) + monkeypatch.setattr(proxy_server, "prisma_client", None) + monkeypatch.setattr(proxy_server, "general_settings", {}) + monkeypatch.setattr(proxy_server, "user_model", None) + return router + + +@pytest.fixture +def flagged_router(monkeypatch) -> Router: + return _install_router( + monkeypatch, + _deployment("gpt-4"), + _deployment("internal-evaluator", discoverable=False), + ) + + +@pytest.fixture +def flagged_wildcard_router(monkeypatch) -> Router: + return _install_router( + monkeypatch, + _deployment("gpt-4"), + _deployment("anthropic/*", model="anthropic/*", discoverable=False), + ) + + +@pytest.fixture +def flagged_team_router(monkeypatch) -> Router: + return _install_router( + monkeypatch, + _deployment("gpt-4"), + _deployment( + "model_name_team1_abc", team_id="team1", team_public_model_name="team-gpt", discoverable=False + ), + _deployment("model_name_team1_def", team_id="team1", team_public_model_name="team-chat"), + ) + + +@pytest.fixture +def team_admin_privileges(monkeypatch) -> None: + from litellm.proxy.management_endpoints import common_utils + + async def _is_team_admin(**kwargs) -> bool: + return True + + monkeypatch.setattr(common_utils, "_user_has_admin_privileges", _is_team_admin) + + +def _non_admin() -> UserAPIKeyAuth: + return UserAPIKeyAuth(api_key="sk-test", user_role=LitellmUserRoles.INTERNAL_USER) + + +def _team_member(role: LitellmUserRoles = LitellmUserRoles.INTERNAL_USER) -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key="sk-test", user_id="u", user_role=role, team_id="team1", team_models=["team-gpt", "team-chat"] + ) + + +def _admin() -> UserAPIKeyAuth: + return UserAPIKeyAuth(api_key="sk-test", user_id="u", user_role=LitellmUserRoles.PROXY_ADMIN, team_models=[]) + + +def _anthropic_request() -> Request: + return Request( + scope={ + "type": "http", + "method": "GET", + "path": "/v1/models", + "query_string": b"", + "headers": [(b"anthropic-version", b"2023-06-01")], + } + ) + + +async def _v1_models(user_api_key_dict: UserAPIKeyAuth, **kwargs) -> list[str]: + response = await proxy_server.model_list(user_api_key_dict=user_api_key_dict, **kwargs) + return [m["id"] for m in response["data"]] + + +async def _v1_model_info_names(user_api_key_dict: UserAPIKeyAuth, **kwargs) -> list[str]: + response = await proxy_server.model_info_v1(user_api_key_dict=user_api_key_dict, **kwargs) + return [row["model_name"] for row in json.loads(response.body)["data"]] + + +async def _model_groups(user_api_key_dict: UserAPIKeyAuth) -> list[str]: + response = await proxy_server.model_group_info(user_api_key_dict=user_api_key_dict) + return [group.model_group for group in response["data"]] + + +@pytest.mark.asyncio +async def test_v1_models_openai_shape_hides_flagged_model_from_non_admin_only(flagged_router): + assert await _v1_models(_non_admin()) == ["gpt-4"] + assert await _v1_models(_admin()) == ["gpt-4", "internal-evaluator"] + + +@pytest.mark.asyncio +async def test_v1_models_anthropic_shape_hides_flagged_model_from_non_admin_only(flagged_router): + assert await _v1_models(_non_admin(), request=_anthropic_request()) == ["gpt-4"] + assert await _v1_models(_admin(), request=_anthropic_request()) == ["gpt-4", "internal-evaluator"] + + +@pytest.mark.asyncio +async def test_v1_models_scope_expand_hides_flagged_model_from_team_admin_only(flagged_router, team_admin_privileges): + assert await _v1_models(_non_admin(), scope="expand") == ["gpt-4"] + assert await _v1_models(_admin(), scope="expand") == ["gpt-4", "internal-evaluator"] + + +@pytest.mark.asyncio +async def test_v1_models_by_id_still_serves_the_hidden_model_to_non_admin(flagged_router): + assert "internal-evaluator" not in await _v1_models(_non_admin()) + + response = await proxy_server.model_info(model_id="internal-evaluator", user_api_key_dict=_non_admin()) + assert response["id"] == "internal-evaluator" + + +@pytest.mark.asyncio +async def test_v1_models_group_with_one_discoverable_deployment_stays_listed(monkeypatch): + _install_router( + monkeypatch, + _deployment("shared", discoverable=False), + { + "model_name": "shared", + "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-fake"}, + "model_info": {"id": "shared-public"}, + }, + _deployment("internal-evaluator", discoverable=False), + ) + + assert await _v1_models(_non_admin()) == ["shared"] + + +@pytest.mark.asyncio +async def test_v1_models_only_an_explicit_false_hides_a_model(monkeypatch): + _install_router( + monkeypatch, + _deployment("gpt-4"), + _deployment("public-eval", discoverable=True), + _deployment("internal-evaluator", discoverable=False), + ) + + assert await _v1_models(_non_admin()) == ["gpt-4", "public-eval"] + + +@pytest.mark.asyncio +async def test_v1_model_info_hides_flagged_rows_from_non_admin_only(flagged_router): + assert await _v1_model_info_names(_non_admin()) == ["gpt-4"] + assert await _v1_model_info_names(_admin()) == ["gpt-4", "internal-evaluator"] + + +@pytest.mark.asyncio +async def test_v1_model_info_by_id_still_serves_the_hidden_row_to_non_admin(flagged_router): + assert "internal-evaluator" not in await _v1_model_info_names(_non_admin()) + + assert await _v1_model_info_names(_non_admin(), litellm_model_id="internal-evaluator-id") == [ + "internal-evaluator" + ] + + +@pytest.mark.asyncio +async def test_model_group_info_hides_flagged_group_from_non_admin_only(flagged_router): + assert await _model_groups(_non_admin()) == ["gpt-4"] + assert await _model_groups(_admin()) == ["gpt-4", "internal-evaluator"] + + +@pytest.mark.asyncio +async def test_v1_models_hides_flagged_team_model_from_its_team_member_only(flagged_team_router): + assert await _v1_models(_team_member()) == ["team-chat"] + assert set(await _v1_models(_team_member(LitellmUserRoles.PROXY_ADMIN))) >= {"team-gpt", "team-chat"} + + +@pytest.mark.asyncio +async def test_model_group_info_hides_flagged_team_model_from_its_team_member(flagged_team_router): + assert await _model_groups(_team_member()) == ["team-chat"] + + +@pytest.mark.asyncio +async def test_v1_models_hides_flagged_wildcard_expansions_from_non_admin(flagged_wildcard_router): + assert await _v1_models(_non_admin(), return_wildcard_routes=True) == ["gpt-4"] + + admin_ids = await _v1_models(_admin(), return_wildcard_routes=True) + assert "gpt-4" in admin_ids + assert any(model_id.startswith("anthropic/") for model_id in admin_ids) + + +@pytest.mark.asyncio +async def test_v1_model_info_hides_flagged_wildcard_expanded_rows_from_non_admin(flagged_wildcard_router): + assert await _v1_model_info_names(_non_admin()) == ["gpt-4"] + + admin_names = await _v1_model_info_names(_admin()) + assert "gpt-4" in admin_names + assert any(name.startswith("anthropic/") for name in admin_names) + + +@pytest.mark.asyncio +async def test_hidden_model_still_routes_for_direct_requests(flagged_router): + assert "internal-evaluator" not in await _v1_models(_non_admin()) + + deployment = flagged_router.get_available_deployment( + model="internal-evaluator", messages=[{"role": "user", "content": "hi"}] + ) + assert deployment["model_name"] == "internal-evaluator" diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 89fd9c5c9d4..62ff08230d7 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -5787,12 +5787,9 @@ async def test_model_info_v1_oci_secrets_not_leaked(): from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.proxy_server import model_info_v1 - # Mock user authentication - mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) - mock_user_api_key_dict.user_id = "test-user" - mock_user_api_key_dict.api_key = "test-key" - mock_user_api_key_dict.team_models = [] - mock_user_api_key_dict.models = ["oci-grok-test"] + mock_user_api_key_dict = UserAPIKeyAuth( + user_id="test-user", api_key="test-key", team_models=[], models=["oci-grok-test"] + ) # Mock model data with OCI sensitive information mock_model_data = { diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index bdfd4aec316..37cfd752936 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -42027,6 +42027,8 @@ export interface components { litellm__proxy___types__ModelInfo: { /** Base Model */ base_model: ("gpt-4-1106-preview" | "gpt-4-32k" | "gpt-4" | "gpt-3.5-turbo-16k" | "gpt-3.5-turbo" | "text-embedding-ada-002") | null; + /** Discoverable */ + discoverable?: boolean | null; /** Id */ id: string | null; /** @@ -42074,6 +42076,8 @@ export interface components { * @default false */ db_model: boolean; + /** Discoverable */ + discoverable?: boolean | null; /** Enable Tag Filtering */ enable_tag_filtering?: boolean | null; /** Id */