mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
Merge branch 'main' into litellm_mcp_persistent_upstream_session
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
commit
e2af6f7bd0
159 changed files with 4451 additions and 932 deletions
161
.github/pull_request_template.md
vendored
161
.github/pull_request_template.md
vendored
|
|
@ -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 -->
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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`.
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
# ====================================================================== #
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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 ("}", "]")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1133,6 +1133,7 @@ class ModelInfo(LiteLLMPydanticObjectBase):
|
|||
]
|
||||
| None
|
||||
)
|
||||
discoverable: bool | None = None
|
||||
|
||||
model_config = ConfigDict(protected_namespaces=(), extra="allow")
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
118
litellm/proxy/common_utils/discoverable_model_filter.py
Normal file
118
litellm/proxy/common_utils/discoverable_model_filter.py
Normal 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))
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
]
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"))),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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] = {}
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue