Merge branch 'main' into litellm_mcp_persistent_upstream_session
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-24 00:56:27 +00:00
commit e2af6f7bd0
159 changed files with 4451 additions and 932 deletions

View file

@ -1,27 +1,162 @@
<!-- Plain English please. Describe the change the way you would explain it to a teammate who has not seen the code: what it does and why, not which functions, files, or tables it touches -->
<!-- The whole description's target audience is humans, not AI agents: write it in plain, simple,
everyday engineering language, extremely parsable and readable at a glance. This goes double for
the TLDR, User Flow, and Caveats sections -->
## What's the problem?
## TLDR
## What's the solution?
<!-- Fill in the bullets below and keep each one short and concrete: one line per bullet, roughly 10 words max -->
<!-- The approach, in a sentence or two -->
Problem this solves:
## How does it fix it?
- <blah>
- ...
<!-- What actually changes so the problem can't happen anymore -->
How it solves it:
## How does the product experience change?
- <blah>
- ...
<!-- What a user could do or see before, and what they can do or see after. If nothing user-facing changes, say so -->
## User Flow
## What caveats are there, if any?
<!-- Two ordered lists, Before and After, walking the same end user through the same task, written strictly from that user's seat
Read the linked issue, ticket, or customer thread first so the flow reflects the real application and the routes its users actually hit; don't invent a generic scenario
Lead each list with one plain sentence saying where the flow fails (Before) or succeeds (After), then number the steps
Every step is something the user does or observes: the HTTP method and full URL they hit, what they sent, and what visibly came back (status code, error text, the shape of an ID). UI steps name the page URL and what is on screen
No LiteLLM internals: never name functions, files, DB tables, config classes, hooks, callbacks, or code paths. "The upload hands back an ID that looks like OpenAI's own `file-abc123` instead of the scrambled one the gateway returned" is right, "no managed-file row was registered" is wrong
Keep the two lists step-for-step identical until they diverge, so the changed step is obvious
If the bug had a security or authorization consequence, end each list with what another user could or could no longer do
Regenerate this section whenever new commits change the PR's behavior, so it never describes an older revision
<!-- Major ones only: behavior that breaks on purpose, migrations that lock tables, auth changes, known gaps you did not fix. Write "None" if there are none -->
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
<!-- e.g., "Fixes #000" -->
## Affected release
<!-- Only for a fix to a regression in a released or rc version (perf, memory, crash, or behavior): name the version it regressed in, e.g. "regression in v1.100.0" or "since v1.101.0-rc.1", and add the `backport-stable` label so the fix is cherry-picked onto the rc line before the stable is tagged. Leave the section blank otherwise -->
## Linear ticket
<!-- Internal contributors: Resolves LIT-1234. Otherwise leave blank -->
<!-- if you are an internal contributor, add "Resolves " followed by the Linear ticket e.g., "Resolves LIT-1234" to link the Linear ticket to the GitHub PR. If you don't have one, leave the section blank rather than guessing -->
## 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/<your_test_file>.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
<!-- Include screenshots, screen recordings, or command (e.g., curl) + output demonstrating that your changes work as expected
The proof must be completely e2e with no mocks, using actual LLM calls costing real $$$ if applicable. `pytest` commands are not enough
Show ONLY the latest run: capture Before at the merge base and After at the PR's current tip, and when new commits change behavior, replace this whole section with the fresh run instead of stacking it on top of older ones. The run must be up to date. As soon as a new commit is made and it makes this PR description's after sha stale (it's no longer tip of PR), you must re-run the QA
Structure the section exactly as below: Before and After one heading level below this section, each naming the commit hash it was captured at, one lower-level heading per case inside each, the same case names in the same order on both sides, and numbered steps (command, observed output) under every case, never loose prose; shared setup (config, payloads) goes above Before, and with a single case, drop the case headings and number the steps directly
### Before (<hash>)
#### <case 1>
1. ...
2. ...
#### <case 2>
1. ...
### After (<hash>)
#### <case 1>
1. ...
2. ...
#### <case 2>
1. ...
For bug fixes: Before shows the reproduction, After shows the same steps passing
For new features: Before shows the capability missing, After shows it working end-to-end
If the change applies to all three LLM endpoints (/v1/responses, /v1/chat/completions, /v1/messages), make each endpoint its own case, not just one
For UI changes: before/after screenshots under the same headings
If the main use case runs through a coding tool like Claude Code or Codex, drive that tool interactively the way the user does (never `claude -p`, `codex exec`, or curl on its own) and embed before/after screenshots of its pane under the same headings; curl replays and headless runs can follow as extra cases, never as the only proof -->
## Type
<!-- Select the type of Pull Request -->
<!-- Keep only the necessary ones -->
🆕 New Feature
🐛 Bug Fix
🧹 Refactoring
📖 Documentation
🚄 Infrastructure
✅ Test
## Caveats (if any)
<!-- Group caveats under severity subheadings (### Severe, ### High, ### Medium, ### Low), with
short bullet points inside each, just like the TLDR: one line per bullet, roughly 10 words max
Call out known limitations, follow-up work, or anything a reviewer should watch out for
Include only the tiers that have caveats; drop the empty ones
- Severe: inherent to what the PR deliberately ships, there even when the code works as intended:
it can degrade or take down a running deployment (e.g. a slow or table-locking boot migration),
rewrite data by design, break an existing workflow on purpose, or change auth behavior. An
operator must plan around it before rollout
- High: an unintended hole: a correctness, security, data-loss, or backward-compatibility bug,
unsafe to ship as is
- Medium: a real gap someone can hit, but with a workaround or a narrow blast radius
- Low: anything else worth noting: naming, cleanup, an edge case nobody hits
Nest bullets as deep as helps: hierarchy beats one long line when it makes things clearer to a
human reader
If you assumed something instead of testing it, e.g. "only reproduces with X on" or "no
user-observable behavior difference", list it here too with what breaks if it is wrong
Leave this section empty if there are none -->
## QA runbook
<!-- Only needed when your PR edits tests/e2e; delete this section otherwise
For each e2e test you added or changed, list the manual steps a reviewer can follow to reproduce it by hand against a live proxy, mapping 1:1 to what the test asserts: one top-level bullet per test giving its pytest node id followed by what it proves in plain words, then a nested "- [ ]" checklist where each item is a concrete action (route, request body, expected response) and the final item is the sanity-check step shown in the examples. Note environment prerequisites (provider credentials, config flags) and any nuances a manual run will hit. See PRs #32914 and #32963 for full examples
Example checklists:
- tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py::TestKeyRateLimits::test_rpm_limit_blocks_over_limit - a key allowed 2 requests a minute serves exactly 2 and refuses the 3rd
- [ ] Generate a limited key: curl -X POST http://localhost:4000/key/generate -H "Authorization: Bearer sk-1234" -d '{"rpm_limit": 2}'
- [ ] Send three /v1/chat/completions requests with that key inside one minute
- [ ] Expect the first two to return 200 and the third to return 429 naming the rpm limit
- [ ] Sanity check: this test makes sense to add and is not hand-wavey (e.g., assert actual expected spend instead of just spend > 0) or potentially flaky
- tests/e2e/management/test_management_e2e.py::TestModelRoutes::test_model_create_appears_in_ui - a deployment created through the API shows up on the Admin UI models page
- [ ] POST /model/new with the master key, a bedrock model, and aws_region_name (needs STORE_MODEL_IN_DB=True and AWS credentials)
- [ ] Open http://localhost:4000/ui/?page=models and expect a deployment row showing the returned model id
- [ ] Sanity check: this test makes sense to add and is not hand-wavey (e.g., assert actual expected spend instead of just spend > 0) or potentially flaky
-->
## 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
<!-- What you ran against a live proxy and what came back. Commands with output or before/after screenshots work well. Unit tests alone are not enough -->

View file

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

View file

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

View file

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

View file

@ -767,7 +767,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
@ -930,9 +930,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

View file

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

View file

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

View file

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

View file

@ -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,
)
# ====================================================================== #

View file

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

View file

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

View file

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

View file

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

View file

@ -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})}
),
}
)
@ -808,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
)
@ -836,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)
@ -881,6 +945,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 +980,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 +998,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,
)
@ -961,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
@ -1275,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(

View file

@ -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},
)

View file

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

View file

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

View file

@ -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 ("}", "]")

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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]:
"""

View file

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

View file

@ -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}},
)

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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 = []

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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,
@ -636,8 +766,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)
@ -882,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,
@ -972,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
@ -1088,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

View file

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

View file

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

View file

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

View file

@ -1460,6 +1460,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:
@ -6676,6 +6705,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

View file

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

View file

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

View file

@ -1133,6 +1133,7 @@ class ModelInfo(LiteLLMPydanticObjectBase):
]
| None
)
discoverable: bool | None = None
model_config = ConfigDict(protected_namespaces=(), extra="allow")

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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"]:

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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:
@ -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(
@ -3126,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

View file

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

View file

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

View file

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

View file

@ -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
@ -706,9 +723,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
@ -1390,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,
@ -1401,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)
@ -1751,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,
@ -1762,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)
@ -1810,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(
@ -1819,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."
)
@ -1836,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:
@ -1845,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,
@ -1853,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
@ -1865,9 +1886,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))
@ -2931,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,

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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,
@ -5261,9 +5262,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 +10136,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 +10162,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 +10198,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
@ -11243,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
@ -11296,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
@ -15795,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)
@ -16074,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

View file

@ -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] = {}

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

Some files were not shown because too many files have changed in this diff Show more