feat: litellm oss 110626 (#30202)

* Add gpt-realtime-whisper Realtime transcription support (OpenAI + Azure) (#29775)

* Add gpt-realtime-whisper Realtime transcription support (OpenAI + Azure)

Adds first-class support for the gpt-realtime-whisper streaming speech-to-text
model, which uses the Realtime transcription session API rather than the
file-based /audio/transcriptions path.

Model registration: registers gpt-realtime-whisper and azure/gpt-realtime-whisper
with audio-duration pricing (input_cost_per_second = 0.017/60, matching the
published $0.017/minute input audio rate).

REST endpoint: implements POST /v1/realtime/transcription_sessions (plus /realtime
and /openai/v1 aliases) to mint an ephemeral transcription session for the
WebRTC flow. Adds request/response types, OpenAI and Azure URL builders, a shared
base handler (refactored from the client_secrets handler), the
acreate_realtime_transcription_session SDK function, and route registration. The
proxy encrypts the ephemeral key returned under client_secret.value and records
the session type in the token so the follow-up /realtime/calls replays
type=transcription rather than type=realtime.

WebSocket: forwards intent=transcription through to the Azure handler (OpenAI
already received it) with URL-encoding, so gpt-realtime-whisper opens a
transcription session. Transcription-only sessions no longer trigger an
erroneous response.create.

Cost tracking: transcription sessions emit no response.done events; their usage
arrives on conversation.item.input_audio_transcription.completed as
{type: duration, seconds}. That usage is captured out-of-band (usage only, no
transcript duplication) and billed by input_cost_per_second, with a token-billed
fallback for token-priced transcription models.

Adds tests for pricing math, URL builders, request/response types, the proxy
route and SDK function, WebSocket intent forwarding, transcription-session
streaming behavior, and the /realtime/calls session-type replay.

* Address PR review: URL-encode all Azure WS query params; forward query_params through provider_config branch

* Address PR review: session_type validation, model auth fix, cost perf, billing fallback, detail/docs cleanup

* Improve test coverage: detection from backend, error paths, unknown usage type, resolved_model None

* Backport realtime transcription websocket fixes

* Enforce authorized realtime transcription model

* Enforce realtime transcription model access

* Enforce realtime resolved model scopes

* Enforce WebRTC transcription model scope

* Lazy evaluate debug log in pass-through endpoint (#30177)

* Pass through debug lazy logging

* fix(proxy): convert remaining eager pass-through debug logs to lazy formatting

* fix(parallel_ai): migrate search integration from v1beta to v1 endpoint (#30157)

* fix(parallel_ai): migrate search integration from v1beta to v1 endpoint

The Parallel Search API moved from /v1beta/search (processor: base/pro,
parallel-beta header) to /v1/search (mode: turbo/basic/advanced, no beta
header). Request fields moved too: max_results, source_policy, and excerpt
settings are now nested under advanced_settings, and source_policy uses
include_domains/exclude_domains. The v1 response returns publish_date per
result, which now maps to SearchResult.date instead of being hardcoded to
None. The legacy processor param is mapped to the equivalent mode so
existing callers keep working.

* fix(parallel_ai): default mode to basic and simplify param handling

The v1 API defaults to advanced mode when mode is omitted, while v1beta
defaulted to the base processor. Without an explicit default, callers who
pass no mode would be silently upgraded to a tier costing 2.25x more while
litellm's cost map reports the basic-tier price. Sending mode=basic
preserves the v1beta default and keeps cost tracking accurate.

Also replaces the handled_params set with pop-as-consumed param handling so
mapped params no longer need to be tracked in two places, and extends the
tests to pin the default mode, processor=base mapping, mode-over-processor
precedence, and top-level v1 param passthrough.

* fix(parallel_ai): avoid double /v1 when api_base is already versioned

A PARALLEL_AI_API_BASE like https://api.parallel.ai/v1 previously produced
.../v1/v1/search. Strip a trailing /v1 before appending the search path and
cover the api_base variants with a parametrized test.

---------

Co-authored-by: shin-berri <shin-laptop@berri.ai>
Co-authored-by: yuneng-jiang <yuneng@berri.ai>

* feat(focus): add Mavvrik destination for FOCUS export (#29935)

* fix: preserve responses streaming flag (#30189)

* fix: preserve responses streaming flag

* test: cover async responses streaming flag

* fix(spend/daily-activity): stable offset pagination via id tiebreaker (#30164) (#30167)

date alone is not a unique sort key for LiteLLM_DailyUserSpend or
LiteLLM_DailyTeamSpend (many rows per date: api_key x model x
model_group x provider x endpoint). Offset pagination over a
non-unique sort landed on arbitrary boundaries, so a client paging
through all results and summing per-page metrics (the Usage dashboard)
got non-deterministic totals - sometimes inflated, sometimes deflated,
different at different page_size values.

Adding the row's UUID id (present on both tables) as a secondary sort
gives every page a stable cursor. order=[{date desc}, {id asc}].

Fixes #30164

* fix(oci): inject a default maxTokens so omitted max_tokens doesn't truncate responses (#30018)

* fix(oci): inject default maxTokens so omitted max_tokens doesn't truncate

OCI GenAI applies a tiny server-side maxTokens default (~20 tokens) when the
request omits it, so any call that doesn't send max_tokens comes back cut off
mid-string with finishReason "length". MLflow judges never send max_tokens, so
their JSON responses arrived as unterminated strings and json.loads failed in
MLflow's gateway adapter.

When no maxTokens/maxCompletionTokens target is set, inject
DEFAULT_OCI_CHAT_MAX_TOKENS (env-overridable, defaults 4096), mirroring the
Anthropic config's default-max-tokens behaviour. An explicit max_tokens still
wins, and reasoning models still route to maxCompletionTokens. Used a fixed
default rather than the catalog max_output_tokens because the catalog value is
unreliable for some models (grok-4 reports max_output_tokens equal to its
context window, not a real output cap, which would risk 400s).

Adds TestOCIDefaultMaxTokens covering Cohere and generic injection, the
explicit-override case, and the reasoning maxCompletionTokens branch.

* test(oci): e2e regression that omitted max_tokens isn't truncated

Real-proxy integration test asserting a chat completion that omits max_tokens
completes with finish_reason "stop" instead of being cut off at OCI's ~20-token
server default. Fails before the maxTokens-default injection (finish_reason
"length", ~19 tokens), passes after.

* test(oci): update cohere default-params test for injected maxTokens

test_cohere_default_parameters asserted no maxTokens was injected, encoding the
old behaviour where OCI's ~20-token server default truncated responses. Now
that transform_request injects DEFAULT_OCI_CHAT_MAX_TOKENS, assert maxTokens
equals that default while the other params (topK/topP/frequencyPenalty) stay
pass-through with no hardcoded default.

* fix(oci): make DEFAULT_OCI_CHAT_MAX_TOKENS a plain constant

Drop the os.getenv override. The env knob was not requested and introducing a
new env var forced a cross-repo dependency on litellm-docs (test_env_keys.py
validates every referenced env var against the docs table there). A plain 4096
constant keeps the PR self-contained; callers who want a different limit pass
max_tokens explicitly per request.

* fix(oci): route all OpenAI commercial models to maxCompletionTokens

OCI serves OpenAI models (gpt-4.1, gpt-5.1 through 5.5, o-series) that
the litellm catalog doesn't track, so the supports_reasoning lookup
returned False for them and the provider sent maxTokens, which the
reasoning families reject with HTTP 400. With the injected default
maxTokens this broke every request to those models, not just ones with
an explicit max_tokens. Route the whole openai.* vendor prefix to
maxCompletionTokens since OpenAI accepts max_completion_tokens on every
chat model; the openai.gpt-oss-* open weights are served by OCI's own
stack and keep maxTokens. Verified live against gpt-5.2, gpt-5, gpt-4o,
gpt-4.1, gpt-oss-120b, llama-3.3, command-a and grok-3-mini

* test(oci): hoist transformation imports and drop unused ones

Makes the generic-chat test file ruff-clean: the per-test local imports
of OCIChatConfig/OCIVendors shadowed the module-level import (F811) and
left it unused (F401), and json plus three OCI type imports were never
referenced

* fix(oci): translate response_format json_schema to OCI's accepted shape (#29691)

* fix(oci): translate response_format json_schema to OCI's accepted shape

OCI GenAI rejected every json_schema response_format with HTTP 400
"Please pass in correct format of request", which broke structured-output
callers such as MLflow LLM judges (they always send a json_schema).

The provider forwarded OpenAI's raw json_schema body unchanged. For GENERIC
models OCI's ResponseJsonSchema accepts only name/description/schema/isStrict,
so OpenAI's `strict` key (and any other extra) 400s the request; the key must
be renamed to isStrict and the body whitelisted. For Cohere models there is no
JSON_SCHEMA type at all; the schema has to ride on JSON_OBJECT as
{"type": "JSON_OBJECT", "schema": ...}. Cohere type values must also be the
canonical uppercase TEXT/JSON_OBJECT.

_normalize_response_format now branches by vendor and emits the exact shape
each one accepts (verified live against OCI GenAI for Cohere, Meta, Gemini and
Grok). Drops the unused, incorrect Cohere response-format pydantic models.

Two existing tests asserted the broken behavior (lowercase type, raw
jsonSchema on Cohere); they are rewritten to assert the corrected shape, and
generic/Cohere json_schema regression tests are added.

* fix(oci): raise early on json_schema response_format with no body

A GENERIC model request with {"type": "json_schema"} and no json_schema
object fell through to the JSON_OBJECT branch and emitted a bodyless
{"type": "JSON_SCHEMA"}, which OCI rejects with an opaque HTTP 400. Raise a
descriptive 400 at translation time instead. Cohere is unaffected since it
always maps to JSON_OBJECT.

* test(oci): gateway integration test for response_format json_schema

Added to tests/integration/ (the real-network integration suite) reusing the
existing OCI proxy harness, not tests/llm_translation/ which is mock-only.

---------

Co-authored-by: Sameer Kankute <sameer@berri.ai>

* fix(oci): accept default n=1 on Cohere instead of hard-failing (#29705)

* fix(oci): accept default n=1 on Cohere instead of hard-failing

Cohere on OCI has no numGenerations field, so n was mapped to False and
map_openai_params raised "param `n` is not supported on OCI" whenever a client
sent n. But n=1 (and None) is the OpenAI default single-generation request,
which every OCI model produces anyway, so standard clients that always send
n=1 (such as the MLflow gateway) were rejected with a 500.

Drop n=1/None silently for Cohere; only n>1 is genuinely unsupported and still
raises (or drops under drop_params). Generic models are unaffected and keep
numGenerations, including n>1.

* docs(oci): explain why n is not advertised for Cohere despite tolerating n=1

* test(oci): gateway integration test for Cohere default n=1

Added to tests/integration/ (the real-network integration suite) reusing the
existing OCI proxy harness, not tests/llm_translation/ which is mock-only.

---------

Co-authored-by: Sameer Kankute <sameer@berri.ai>

* fix(oci): drop max_retries instead of hard-failing on OCI (#29727)

max_retries is a litellm-level control param (litellm applies retries itself),
not a generation param OCI accepts. The provider mapped it to False and raised
"param `max_retries` is not supported on OCI" whenever it was present. The
litellm proxy injects max_retries on every request, so any OCI call through the
proxy 500'd unless drop_params was set.

Drop max_retries silently in map_openai_params. Adds a unit test (Cohere and
generic) and a gateway integration test that a plain request succeeds through a
proxy without drop_params.

Co-authored-by: Sameer Kankute <sameer@berri.ai>

* fix(spend-logs): rehydrate metadata JSONB text on ui_view_spend_logs (#29682)

Fixes #29674.

`/spend/logs/ui` raw-SQL path returns the JSONB metadata column as a
string — prisma's query_raw skips the ORM-layer hydration. The UI reads
metadata.status / metadata.error_information as object fields, so
provider-failure rows look like successes.

Fix: json.loads the metadata field right after query_raw, fall back to
{} on malformed JSON.

3 existing error-code/error-message tests called json.loads on
response.data[0]["metadata"] — they were leaning on the bug. Updated
to read the dict directly. Plus 2 new regression tests (failure metadata
roundtrip + invalid-json fallback). Reverting the fix makes both new
tests fail with AssertionError: metadata should be dict, got <class 'str'>.

* fix(proxy): release max_parallel_requests slot when a stream is cancelled mid-flight (#27955) (#30020)

* fix(proxy): release max_parallel_requests slot when a stream is cancelled mid-flight (#27955)

* fix: refund max_parallel_requests on disconnect from outer streaming generators

The cancellation refund previously lived in async_post_call_streaming_iterator_hook,
but that hook is nested inside the outer streaming generators and a nested async
generator only receives GeneratorExit on garbage collection (non-deterministic).
With only the v3 limiter enabled, /chat/completions also bypasses the hook entirely
(needs_iterator_wrap() is false). Move the release into async_data_generator and
async_streaming_data_generator, the generators Starlette closes on client disconnect,
so the refund fires deterministically on every streaming route. Warn when no event
loop is running, and document the window TTL refresh on the decrement

* fix(mcp): propagate model into model_call_details for passthrough tool calls (#30122)

* fix(mcp): propagate model into model_call_details for passthrough tool calls

The @client decorator on call_mcp_tool creates the logging object via
function_setup without a model kwarg, so model_call_details["model"]
starts as None. execute_mcp_tool only set logging_obj.model as an
instance attribute, which the spend-log writer never reads (it reads
kwargs["model"] from model_call_details). MCP passthrough tools/call
rows therefore persisted with model="" while list_tools rows showed
"MCP: list_tools", degrading the Logs UI display and bucketing all MCP
tool spend under an empty model in DailyUserSpend.

Propagate the model into model_call_details alongside the existing
attribute assignment so the StandardLoggingPayload and SpendLogs writer
pick it up. Covers the /mcp passthrough, REST /mcp-rest/tools/call, and
orchestrated paths (the latter already passed model into function_setup,
so this is a no-op there).

* test(mcp): trim regression test docstring

* fix(mcp): surface upstream challenges for delegated OAuth (#30124)

* fix(mcp): surface upstream challenges for delegated OAuth

* docs(mcp): clarify delegated upstream auth comments

* perf(benchmarks): add CPU timing metrics to streaming benchmark (#29980)

* Add CPU timing metrics to streaming benchmark

* Fix spacing around timing sample dataclass

* fix(gemini): don't emit empty choices on metadata-only stream chunks (#29167)

web_search + reasoning makes Gemini stream mid-chunks that carry only
grounding/thought metadata — no content part, no finishReason.
_process_candidates skips content-less candidates and the existing
fallback only ran when finishReason was set, so choices stayed empty
and the downstream streaming handler raised IndexError on choices[0].
Emit an empty-delta choice for content-less chunks regardless of
finishReason.

Fixes #28884

* fix(key): allow /key/update to clear budget_limits with [] or null (#30085)

* Fix /key/update rejecting budget_limits clear requests with HTTP 400

Sending budget_limits: [] or null to /key/update returned HTTP 400, so
once a key had budget windows the last one could never be removed.

prepare_key_update_data only json.dumps'd budget_limits when the value
was truthy, so [] and None passed through raw to the Prisma Json?
column; jsonify_object only serializes dicts, and prisma-client-py has
no DbNull sentinel for Json? writes, so Prisma rejected both shapes.

Serialize the clear case explicitly as the JSON literal null, matching
how memory_endpoints encodes metadata for the same column type. Truthy
values keep the existing reset_at window initialization path.

Fixes #30067.

* Require admin access for budget_limits changes on /key/update

Clearing budget_limits via [] or null is a budget mutation, but
_validate_update_key_data only counted max_budget and spend as budget
changes before deciding whether to skip _check_key_admin_access. A
non-admin key owner or a team member with /key/update could therefore
remove a key's per-window spend caps without admin authorization.

Treat any explicit budget_limits value in the request (set, change, or
clear) as a budget change so it gates through the same admin check as
max_budget. model_fields_set is used because an explicit null is
indistinguishable from an omitted field by value alone.

* fix(proxy): persist guardrail info in spend logs for /v1/responses (#30092)

Pre-call guardrail blocks on /v1/responses wrote guardrail_information
as null in LiteLLM_SpendLogs because _handle_logging_proxy_only_error
splits request_data by LoggedLiteLLMParams keys and litellm_metadata,
where the Responses API stores request metadata including
standard_logging_guardrail_information, was not among them. It fell
into optional_params, so merge_litellm_metadata never saw it. Add
litellm_metadata to LoggedLiteLLMParams so it routes into
litellm_params the same way metadata does on the chat completions path

Fixes #28971.

* fix(proxy): handle non-standard SSE frames in Anthropic passthrough logging (#26000)

Some third-party Anthropic-compatible providers emit non-standard SSE
frames (OpenAI-style [DONE] sentinels, non-JSON keep-alive lines) in
streaming responses. These caused json.JSONDecodeError in
_build_complete_streaming_response, breaking the passthrough logging
pipeline so the request was never logged or billed.

Skip whole-line 'data: [DONE]' sentinels and catch JSONDecodeError per
event. Matching the full line (not a substring) keeps a valid chunk
whose text payload contains '[DONE]' from being dropped.

Co-authored-by: shin-berri <shin-laptop@berri.ai>
Co-authored-by: yuneng-jiang <yuneng@berri.ai>
Co-authored-by: Sameer Kankute <sameer@berri.ai>

* feat(newrelic): Add New Relic extension  (#26989)

* initial New Relic integration.

* Minor fixes for basic observability.

* Implemented basic support for the success path. Generates New Relic
custom events needed by the AI Monitorin interface.

* Supportability metric is sent on first request.

* Emit supportability metric every hour instead of once a day.

* Add the start/end times to the messages before sending them so that the
start time and end time reflect the correct time and both are not set
to 'now'.

* Make use of `turn_off_message_logging` configuration that is available
by default from CustomLogger.

* Enabling New Relic agent to be wired when docker container starts if an environment variable
is set.

* If we cannot find trace information, send the AI events without the
trace ID attached.

* Use a fake trace_id if we cannot find one.

* Implementing a configuration so that users can use litellm configuration
to disable sending LLM messages to New Relic. There is a second method
to do this via New Relic env var.

* Mised file.

* Cleaning up logic to turn off recording content via either the
LiteLLM configuration or an env var.

* Removing debugging.
Fixed logic / comments around how often to send supportability metric.

* Initial version of public doc for New Relic.

* Use a proper name for the doc file.

* Updating newrelic.md document.

* Updating LiteLLM documentation for New Relic extension.

* Moving New Relic imports into the methods to support unit tests.

* Adding unit tests for the New Relic extension.

* Updating linting and the unit tests that are not running in the CI environment.

* Address reviewer feedback on New Relic integration.

- Fix _record_error_metric to use app.record_custom_metric() instead of
  module-level newrelic.agent.record_custom_metric() so the call works
  outside of an active transaction context
- Remove unreachable except ImportError block in _get_trace_context
- Update stale "23 hours" comment to "27 hours" (matches 97200s threshold)
- Remove commented-out debug code from _process_success
- Fix docs typo: NEW_RELIC_CUSTOM_INSIGHTS_EVENTS_MAX_SAMPLES_STOREDA ->
  NEW_RELIC_CUSTOM_INSIGHTS_EVENTS_MAX_SAMPLES_STORED
- Update TestRecordErrorMetric to verify app.record_custom_metric call

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

* Reformating for the linter.

* Addressing additional automated feedback.

- Removed a legacy comment about the New Relic header
- Reordered imports in one file
- Switched another file to use the import at the top of the file instead of inline when used
- Added unit tests for untested methods that were identified

* Addressing new feedback.

- Proper handling of time to floats. Created a util method and updated code to use it.
- added the missing guard to ensure the app is enabled

* Addressing feedback.

- When an error occurs, still check if the periodic supportability metric should be emitted
- Added a check to ensure the extension is ready in the error handler to match _process_success

* Updating the NR event timestamps to more accurately reflect when
the messages were generated.

* Addressing feedback for potential better practice.

* Addressing feedback on accessing default values. Added tests for most of
these cases.

* Adding a new catch exception block based on feedback.

* Addressing feedback about a potential issue around a timestamp for the
supportability metric.

* Addressing minor feedback on length of generated, fallback traceId.

* Addressing feedback.

- A few more cases were found where the dictionary access might not return the correct value.
- Handling cases where `traceparent` is not lower cased

* Addressed feedback where the newrelic options might not apply correctly.

* Addressing some feedback.

* Addressing feedback.

* Validating testing / formatting for our changes.

* Updating linting, adding tests, defining data type for UI.

* Configuration for the logging callback definition.

* Adding a newrelic image for the UI to use.

* Putting the New Relic callback in proper alphabetic order.

* Copying the logo to a committed output directory so it shows up in a locally
built container.

* Adding missing definition of new env vars that were causing a build failure.

* Addressing automated feedback from greptile.

* Adding a few more unit tests to increase the code coverage just a bit more.

* Additional unit tests to push coverage to almost 90%.

* Adding a custom newrelic docker image build process. This removes the need to add the newrelic agent
to the core litellm container or dependencies.

* Clarifying message when the New Relic agent is not installed and someone
is trying to use the newrelic extension. Either use the proper image
when using docker, or install the agent manually when running from source.

* Ensuring pip is available to install the New Relic agent.

* Updating the definition and handling of traceId (no spanId).
Clarifying behavior of env vars vs UI configuration for
the newrelic extension.

* Removing entries from the New Relic logger configuraiton UI as these
values must be set as part of running the image.

* Removing a stale doc file that has moved to the litellm-docs repo.
Cleanup of Dockerfile to remove a LABEL that was incorrect.

* Updating container image name to be the best guess for the new name.

* Addressing feedback from greptile.

- Added a comment around token_count=0
- Updated the boolean parser to allow a wider set of options which matches existing patterns in other parts of LiteLLM.

* Removing option for a separate New Relic container image. The agreement
is to handle this in the New Relic integration docs.

* Updating error message when New Relic agent is not available.

* Wiring in the test message from the LiteLLM callback UX.

* Missed saving one of the file conflicts.

* Fixed a lint error I introduced. Somehow, I dropped another string
and now added it back.

* Adding newrelic to the schema definition.

* Added an admin check on the call before sending test message
as mentioned by the AI code review.

* Updating to use should_redact_message_logging(kwargs) as part of the
logic to determine if message content should be sent to New Relic
or not. This still uses the `record_content` property as well, but
both have to be true in order for content to be included.

---------

Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com>

* Add Azure AI Foundry DeepSeek V3.1 and V4 Pro/Flash global pricing to cost map (#30134)

Co-authored-by: Cursor <cursoragent@cursor.com>

* fix(logging): translate Responses bridge result to ModelResponse for spend logs (#28985)

PR #29394 fixed the AnthropicResponse.model_validate crash for the streaming
anthropic_messages -> OpenAI Responses bridge by unwrapping terminal events
and returning the inner ResponsesAPIResponse. The spend_logs row lands and
usage/cost are correct, but the row's response field stores the Responses
API shape (output[...].content[...].text). The proxy UI Logs tab reads
response.choices[0].message via parseMessages in prettyMessagesUtils.ts
with no fallback for the Responses shape, so the OutputCard renders "No
response data available" for every cross-routed call. The same shape
mismatch affects every downstream consumer of spend_logs that assumes the
canonical chat-completion shape

This change keeps the unwrap from #29394 but routes the resulting
ResponsesAPIResponse (and the bare-response non-streaming path) through
LiteLLMResponsesTransformationHandler.transform_response, which is the
same conversion already used by the chat-completion Responses bridge.
Spend_logs now stores a ModelResponse with choices[0].message.content, so
the UI and other consumers see the assistant text. On a translation
failure (eg. empty output on an incomplete response) the handler falls
back to a minimal ModelResponse carrying model and usage so the row still
lands rather than being dropped as a Non-Blocking error

Also corrects a stale comment in the Responses adapter that implied the
call type was reclassified to acompletion; the code preserves
anthropic_messages and the success handler translates back to
ModelResponse for the row

Fixes #28595

* fix(anthropic-adapter): re-emit first delta on streaming content-block transitions (#30024)

* fix(anthropic-adapter): re-emit first delta on streaming content-block transitions

The `/v1/messages` -> `/v1/chat/completions` streaming adapter
(`AnthropicStreamWrapper`) silently dropped the first non-empty delta of
every content block that started via a *transition* (e.g. text -> tool_use ->
text, text -> thinking).

When an upstream chunk both triggers a new content block (its type differs
from the active block) and carries that block's first delta, the wrapper
emitted `content_block_stop` -> `content_block_start` and then only re-queued
the trigger chunk when it was an `input_json_delta` (bundled tool args). The
synthesized `content_block_start` always carries an empty body, so the first
`text_delta` / `thinking_delta` was lost — the client output started from the
second token (e.g. "Hi, how can I help you?" rendered as ", how can I help
you?", or text resuming after a tool call lost its first sentence). This is
especially visible with Claude Code-style clients that consume Anthropic
Messages streaming events strictly.

Fix: re-queue the trigger chunk's translated delta whenever it carries
non-empty content (text/thinking/signature/tool args), via a shared
`_trigger_delta_has_content` helper used by both the sync and async paths.
Empty trigger deltas are still suppressed so no spurious empty
`content_block_delta` is introduced.

Fixes #30014

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

* test(anthropic-adapter): cover all _trigger_delta_has_content branches

Add a direct parametrized unit test for the re-emit predicate so every delta
type (text/input_json/thinking/signature), the empty-payload guards, and the
malformed/non-delta cases are exercised independently of upstream chunk
translation. Raises patch coverage for the new helper.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

---------

Co-authored-by: shin-berri <shin-laptop@berri.ai>
Co-authored-by: yuneng-jiang <yuneng@berri.ai>
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>

* feat: add opt-in healthy_only filter to GET /v1/models (#30130)

* feat: add opt-in healthy_only filter to GET /v1/models

Adds an opt-in `healthy_only=true` query parameter to GET /v1/models and
GET /models that hides models whose backing deployments are all marked
unhealthy by background health checks.

- Add Router.async_get_fully_unhealthy_model_names(), mirroring the
  semantics of get_fully_blocked_model_names(): a model is hidden only
  when every backing deployment is unhealthy and the health state is
  not stale (fail open otherwise).
- Reuses the existing DeploymentHealthCache populated by
  _run_background_health_check(), so no new health state is introduced.
- No-op when allowed_fails_policy is set, mirroring
  _async_filter_health_check_unhealthy_deployments semantics.
- team_public_model_name aliases are aggregated alongside model_name.
- Hiding is presentation-only; default behavior is unchanged.

Fixes #30128

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>

* docs: address Greptile review notes

- Note team-alias asymmetry vs get_fully_blocked_model_names
- Debug-log when healthy_only is set but no health state is available

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>

---------

Co-authored-by: shin-berri <shin-laptop@berri.ai>
Co-authored-by: yuneng-jiang <yuneng@berri.ai>
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>

* Dedupe team soft budget alerts by team_id instead of token (#30097)

_team_soft_budget_check sends type="soft_budget" alerts with
event_group=TEAM, but SoftBudgetAlert.get_id always returned the
request token. The alert cache key was therefore scoped per virtual
key, so every active key in a team over its soft budget fired its own
alert within budget_alert_ttl. Branch on event_group so team-level
alerts dedupe by team_id, matching TeamBudgetAlert, while key and
project level alerts keep per-token dedupe.

Fixes #27398.

* feat(bedrock guardrails): support contextual grounding qualifiers (request-side) (#30057)

* test: add failing tests for Bedrock contextual grounding (request-side)

Drive the request-side of Bedrock contextual grounding: callers tag message
content blocks as grounding_source/query, the post_call hook assembles an
ApplyGuardrail(OUTPUT) call carrying source + query + response(guard_content),
and the bedrock converse transform must render the tags as prompt text instead
of silently dropping them. Non-grounding payloads must stay byte-identical.

* feat(bedrock guardrails): support contextual grounding qualifiers

Bedrock contextual grounding scores a model response against a reference
source and the user query, expressed via a per-content-block `qualifiers`
array on ApplyGuardrail. The guardrail hook previously sent plain text only,
so grounding could not be driven through it even though the response-side
contextualGroundingPolicy parsing already existed.

Callers now tag message content blocks `{"type":"grounding_source"}` /
`{"type":"query"}` (mirroring the existing `guarded_text` marker). On the
generate path the bedrock converse transform renders them as plain text; at
post_call the hook harvests them from the request and assembles one
ApplyGuardrail(OUTPUT) call carrying grounding_source + query + the response
(as guard_content). Requests without these tags produce a byte-identical
payload, so existing behaviour is unchanged.

* Feat(guardrail): Adding support for custom Ovalix guardrail (#21887)

* Feat(guardrail): Adding support for custom Ovalix guardrail

* Internal CR comments fixes

* greptileai comments fixes

* fix conflict

* fixes

* fix sha256

* clarify Ovalix actor-id hash is for normalization, not PII protection

* fix(github_copilot): normalize per-event item_id in /responses streaming (#30072)

GitHub Copilot's native /v1/responses stream assigns a different item_id to
every event of a single output item (output_item.added, the part.added /
delta / done events, and output_item.done). Spec-strict clients like the
Vercel AI SDK key streaming parts by item_id and abort with
"reasoning part <id> not found" / "text part <id> not found" when a delta
references an unregistered id.

Override transform_streaming_response in GithubCopilotResponsesAPIConfig to
anchor every event of an output item to the id from its output_item.added.
Copilot accepts that id paired with the final encrypted_content on the next
turn, so multi-turn replay is unaffected.

Fixes #30071

* feat: add /model/block and /model/unblock endpoints (#30125)

* feat: add /model/block and /model/unblock endpoints

Add dedicated proxy-admin POST /model/block and /model/unblock endpoints
over the existing blocked flag on LiteLLM_ProxyModelTable, mirroring the
/key/block and /key/unblock pattern. Calling a model whose deployments are
all blocked now returns a clear 403 "Model is blocked" instead of a generic
no-deployment error, including direct-dispatch route types (e.g. eval) via a
pre-route guard. Includes audit-log entries for block/unblock and unit tests.

Closes #29742

Signed-off-by: AgentGymLeader <264910004+AgentGymLeader@users.noreply.github.com>

* chore: regenerate dashboard API types for model block/unblock endpoints

Regenerate ui/litellm-dashboard/src/lib/http/schema.d.ts from the proxy
OpenAPI spec (npm run gen:api) so it includes the new endpoints.

Signed-off-by: AgentGymLeader <264910004+AgentGymLeader@users.noreply.github.com>

* fix: widen router block-helper param type and add direct unit tests

Type the _are_all_deployments_blocked deployments parameter to match its
callers (DeploymentTypedDict) so mypy passes, and add
tests/test_litellm/test_router_block_helpers.py with direct unit tests for
the three block helper methods so router_code_coverage recognizes them.

Signed-off-by: AgentGymLeader <264910004+AgentGymLeader@users.noreply.github.com>

* fix: restore type-ignore on messages arg after black reflow

Signed-off-by: AgentGymLeader <264910004+AgentGymLeader@users.noreply.github.com>

* refactor: raise model-block 403 in proxy layer, not SDK Router

Keep the SDK Router's documented behavior for blocked deployments (filtered ->
"no healthy deployment") and move the 403 PermissionDeniedError into the proxy
layer (route_llm_request), where model blocking is an admin concept. This avoids
a backwards-incompatible 403 for SDK users who set blocked=True on their own
deployments, per maintainer review.

Signed-off-by: FugoP <264910004+AgentGymLeader@users.noreply.github.com>

---------

Signed-off-by: AgentGymLeader <264910004+AgentGymLeader@users.noreply.github.com>
Signed-off-by: FugoP <264910004+AgentGymLeader@users.noreply.github.com>
Co-authored-by: AgentGymLeader <264910004+AgentGymLeader@users.noreply.github.com>
Co-authored-by: Sameer Kankute <sameer@berri.ai>

* fix: add week unit support to get_next_standardized_reset_time (#30100)

* fix: add week unit support to get_next_standardized_reset_time

The function handled d/h/m/s/mo units but silently fell through to
the default next-midnight branch for the w (week) unit. This was
inconsistent: _extract_from_regex already accepted w in its character
class, and duration_in_seconds already returned value * 604800 for it.

Add the missing elif unit == 'w' branch that delegates to
_handle_day_reset with value * 7, which reuses the existing Monday-
alignment logic for 1w and the generic N-day-from-midnight path for
larger multiples.

Add test_week_based_resets covering 1w from a Wednesday (expects next
Monday) and 2w from a Monday (expects 14 days forward at midnight).

Signed-off-by: FugoP <264910004+AgentGymLeader@users.noreply.github.com>

* test: exercise relative week semantics with non-Monday base dates + add docstring

Signed-off-by: FugoP <264910004+AgentGymLeader@users.noreply.github.com>

---------

Signed-off-by: FugoP <264910004+AgentGymLeader@users.noreply.github.com>
Co-authored-by: FugoP <264910004+AgentGymLeader@users.noreply.github.com>

* fix: black formatting and remove undocumented MAVVRIK_FOCUS_FREQUENCY env var

* fix: black formatting with correct version and sync schema.d.ts for healthy_only param

* fix: resolve mypy errors and add transcription_sessions to JSON schema endpoint enum

* fix: restore MAVVRIK_FOCUS_FREQUENCY guard and exclude it from docs key scan

* fix: address Greptile P2 comments - move constant, use UTC datetime, skip redundant team lookup

* revert: restore original team lookup logic in can_key_call_resolved_model

---------

Signed-off-by: AgentGymLeader <264910004+AgentGymLeader@users.noreply.github.com>
Signed-off-by: FugoP <264910004+AgentGymLeader@users.noreply.github.com>
Co-authored-by: Emerson Gomes <emerson.gomes@thalesgroup.com>
Co-authored-by: nina-hu <nina.huuu@gmail.com>
Co-authored-by: Sahith Jagarlamudi <104647530+s-jag@users.noreply.github.com>
Co-authored-by: shin-berri <shin-laptop@berri.ai>
Co-authored-by: yuneng-jiang <yuneng@berri.ai>
Co-authored-by: Praveen Ghuge <95286176+pghuge-cloudwiz@users.noreply.github.com>
Co-authored-by: alex107ivanov <30668368+alex107ivanov@users.noreply.github.com>
Co-authored-by: hcl <chenglunhu@gmail.com>
Co-authored-by: Fede Kamelhar <federico.kamelhar@oracle.com>
Co-authored-by: Armaan Sandhu <74664101+Ar-maan05@users.noreply.github.com>
Co-authored-by: Teo Xian Zhong Augustine <35527068+auggie246@users.noreply.github.com>
Co-authored-by: King Star <mcxin.y@gmail.com>
Co-authored-by: Saksham Maggo <122939011+SakshamMaggo@users.noreply.github.com>
Co-authored-by: Filippo Menghi <113345637+Cyberfilo@users.noreply.github.com>
Co-authored-by: Kelvin <leikaiwei@outlook.com>
Co-authored-by: Josh Bonczkowski <josh.bonczkowski@gmail.com>
Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com>
Co-authored-by: M. Dennis Turp <mdturp@pm.me>
Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: Piotr Minkina <piotrminkina@users.noreply.github.com>
Co-authored-by: Martín Alcalá Rubí <martin@tryolabs.com>
Co-authored-by: T. Kobayashi <13004314+nix-tkobayashi@users.noreply.github.com>
Co-authored-by: João Costa <13508071+jpv-costa@users.noreply.github.com>
Co-authored-by: Shalom <shalom@ovalix.io>
Co-authored-by: codgician <15964984+codgician@users.noreply.github.com>
Co-authored-by: FugoP <kim@pomsora.com>
Co-authored-by: AgentGymLeader <264910004+AgentGymLeader@users.noreply.github.com>
Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
Sameer Kankute 2026-06-12 11:00:26 +05:30 • committed by GitHub
parent 9ddf0535b1
commit cfcdf8714a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
112 changed files with 12140 additions and 446 deletions

View file

@ -43,6 +43,7 @@ from typing import (
Type,
)
from litellm.types.integrations.datadog import DatadogInitParams
from litellm.types.integrations.newrelic import NewRelicInitParams
from litellm._logging import (
set_verbose,
_turn_on_debug,
@ -154,10 +155,12 @@ _custom_logger_compatible_callbacks_literal = Literal[
"gitlab",
"cloudzero",
"focus",
"mavvrik",
"vantage",
"posthog",
"levo",
"compression_interception",
"newrelic",
]
cold_storage_custom_logger: Optional[_custom_logger_compatible_callbacks_literal] = None
logged_real_time_event_types: Optional[Union[List[str], Literal["*"]]] = None
@ -412,6 +415,7 @@ s3_callback_params: Optional[Dict] = None
s3_audit_callback_params: Optional[Dict] = None
datadog_llm_observability_params: Optional[Union[DatadogLLMObsInitParams, Dict]] = None
datadog_params: Optional[Union[DatadogInitParams, Dict]] = None
newrelic_params: Optional[Union[NewRelicInitParams, Dict]] = None
aws_sqs_callback_params: Optional[Dict] = None
generic_logger_headers: Optional[Dict] = None
default_key_generate_params: Optional[Dict] = None
@ -1373,6 +1377,7 @@ from .search.main import *
from .realtime_api.main import (
_arealtime,
acreate_realtime_client_secret,
acreate_realtime_transcription_session,
arealtime_calls,
)
from .responses.main import _aresponses_websocket

View file

@ -171,6 +171,8 @@ class ResponsesToCompletionBridgeHandler:
model_response = validated_kwargs["model_response"]
logging_obj = validated_kwargs["logging_obj"]
custom_llm_provider = validated_kwargs["custom_llm_provider"]
if kwargs.get("stream") is True and "stream" not in optional_params:
optional_params = {**optional_params, "stream": True}
request_data = self.transformation_handler.transform_request(
model=model,
@ -263,6 +265,8 @@ class ResponsesToCompletionBridgeHandler:
model_response = validated_kwargs["model_response"]
logging_obj = validated_kwargs["logging_obj"]
custom_llm_provider = validated_kwargs["custom_llm_provider"]
if kwargs.get("stream") is True and "stream" not in optional_params:
optional_params = {**optional_params, "stream": True}
try:
request_data = self.transformation_handler.transform_request(

View file

@ -421,6 +421,7 @@ REPLICATE_POLLING_DELAY_SECONDS = float(
DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS = int(
os.getenv("DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS", 4096)
)
DEFAULT_OCI_CHAT_MAX_TOKENS = 4096
TOGETHER_AI_4_B = int(os.getenv("TOGETHER_AI_4_B", 4))
TOGETHER_AI_8_B = int(os.getenv("TOGETHER_AI_8_B", 8))
TOGETHER_AI_21_B = int(os.getenv("TOGETHER_AI_21_B", 21))
@ -1483,6 +1484,7 @@ DB_SPEND_UPDATE_JOB_NAME = "db_spend_update_job"
DB_DAILY_TAG_SPEND_UPDATE_JOB_NAME = "db_daily_tag_spend_update_job"
PROMETHEUS_EMIT_BUDGET_METRICS_JOB_NAME = "prometheus_emit_budget_metrics"
CLOUDZERO_EXPORT_USAGE_DATA_JOB_NAME = "cloudzero_export_usage_data"
MAVVRIK_FOCUS_EXPORT_JOB_NAME = "mavvrik_focus_export_usage_data"
CLOUDZERO_MAX_FETCHED_DATA_RECORDS = int(
os.getenv("CLOUDZERO_MAX_FETCHED_DATA_RECORDS", 50000)
)

View file

@ -2488,6 +2488,11 @@ class RealtimeAPITokenUsageProcessor(BaseTokenUsageProcessor):
)
_TRANSCRIPTION_COMPLETED_EVENT_TYPE = (
"conversation.item.input_audio_transcription.completed"
)
def handle_realtime_stream_cost_calculation(
results: OpenAIRealtimeStreamList,
combined_usage_object: Usage,
@ -2533,4 +2538,99 @@ def handle_realtime_stream_cost_calculation(
break # exit if we find a valid model
total_cost = input_cost_per_token + output_cost_per_token
if any(r.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE for r in results):
total_cost += handle_realtime_transcription_cost_calculation(
results=results,
custom_llm_provider=custom_llm_provider,
litellm_model_name=litellm_model_name,
)
return total_cost
def handle_realtime_transcription_cost_calculation(
results: OpenAIRealtimeStreamList,
custom_llm_provider: str,
litellm_model_name: str,
) -> float:
"""
Cost for realtime transcription sessions (e.g. gpt-realtime-whisper).
Transcription sessions emit no `response.done` events; instead each
`conversation.item.input_audio_transcription.completed` event carries a
`usage` object billed by the ASR model. The usage is one of:
- {"type": "duration", "seconds": <float>} → priced via input_cost_per_second
- {"type": "tokens", "input_tokens": ...} → priced via input/audio token cost
"""
completed_events = [
cast(dict, result)
for result in results
if result.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE
]
if not completed_events:
return 0.0
model_name = (
_get_transcription_model_name_from_results(results) or litellm_model_name
)
try:
model_info = litellm.get_model_info(
model=model_name, custom_llm_provider=custom_llm_provider
)
except Exception:
model_info = None
total_cost = 0.0
for event in completed_events:
usage = event.get("usage") or {}
total_cost += _transcription_usage_cost(usage, model_info)
return total_cost
def _get_transcription_model_name_from_results(
results: OpenAIRealtimeStreamList,
) -> Optional[str]:
"""Resolve the ASR model from a transcription_session.* / session.* event."""
for result in results:
if result.get("type") in (
"transcription_session.created",
"transcription_session.updated",
"session.created",
"session.updated",
):
session = cast(dict, result).get("session", {}) or {}
transcription = (
(session.get("audio", {}) or {}).get("input", {}) or {}
).get("transcription", {}) or session.get("input_audio_transcription", {})
model = (transcription or {}).get("model") or session.get("model")
if model:
return model
return None
def _transcription_usage_cost(usage: dict, model_info: Optional[ModelInfo]) -> float:
if model_info is None:
return 0.0
usage_type = usage.get("type")
if usage_type == "duration":
seconds = usage.get("seconds") or 0.0
per_second = model_info.get("input_cost_per_second") or 0.0
return float(seconds) * float(per_second)
if usage_type == "tokens":
input_token_details = usage.get("input_token_details") or {}
audio_tokens = input_token_details.get("audio_tokens") or 0
text_tokens = input_token_details.get("text_tokens") or 0
output_tokens = usage.get("output_tokens") or 0
audio_cost = float(audio_tokens) * float(
model_info.get("input_cost_per_audio_token")
or model_info.get("input_cost_per_token")
or 0.0
)
text_cost = float(text_tokens) * float(
model_info.get("input_cost_per_token") or 0.0
)
output_cost = float(output_tokens) * float(
model_info.get("output_cost_per_token") or 0.0
)
return audio_cost + text_cost + output_cost
return 0.0

View file

@ -1,7 +1,7 @@
from abc import ABC, abstractmethod
from typing import Literal
from litellm.proxy._types import CallInfo
from litellm.proxy._types import CallInfo, Litellm_EntityType
class BaseBudgetAlertType(ABC):
@ -31,6 +31,8 @@ class SoftBudgetAlert(BaseBudgetAlertType):
return "Soft Budget Crossed: "
def get_id(self, user_info: CallInfo) -> str:
if user_info.event_group == Litellm_EntityType.TEAM:
return user_info.team_id or "default_id"
return user_info.token or "default_id"

View file

@ -290,6 +290,21 @@
},
"description": "Langsmith Logging Integration"
},
{
"id": "newrelic",
"displayName": "New Relic",
"logo": "newrelic.png",
"supports_key_team_logging": false,
"dynamic_params": {
"NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED": {
"type": "text",
"ui_name": "Record AI Content (default: true)",
"description": "Whether to record AI message content. Set to false to disable.",
"required": false
}
},
"description": "New Relic AI Monitoring Integration"
},
{
"id": "openmeter",
"displayName": "OpenMeter",

View file

@ -4,6 +4,7 @@ from .base import FocusDestination, FocusTimeWindow
from .factory import FocusDestinationFactory
from .gcs_destination import FocusGCSDestination
from .s3_destination import FocusS3Destination
from .mavvrik_destination import FocusMavvrikDestination
from .vantage_destination import FocusVantageDestination
__all__ = [
@ -12,5 +13,6 @@ __all__ = [
"FocusGCSDestination",
"FocusTimeWindow",
"FocusS3Destination",
"FocusMavvrikDestination",
"FocusVantageDestination",
]

View file

@ -8,6 +8,7 @@ from typing import Any, Dict, Optional
from .base import FocusDestination
from .gcs_destination import FocusGCSDestination
from .s3_destination import FocusS3Destination
from .mavvrik_destination import FocusMavvrikDestination
from .vantage_destination import FocusVantageDestination
@ -32,6 +33,8 @@ class FocusDestinationFactory:
return FocusVantageDestination(prefix=prefix, config=normalized_config)
if provider_lower == "gcs":
return FocusGCSDestination(prefix=prefix, config=normalized_config)
if provider_lower == "mavvrik":
return FocusMavvrikDestination(prefix=prefix, config=normalized_config)
raise NotImplementedError(
f"Provider '{provider}' not supported for Focus export"
)
@ -87,6 +90,15 @@ class FocusDestinationFactory:
"FOCUS_GCS_BUCKET_NAME must be provided for GCS exports"
)
return {k: v for k, v in resolved.items() if v is not None}
if provider == "mavvrik":
resolved = {
"api_key": overrides.get("api_key") or os.getenv("MAVVRIK_API_KEY"),
"api_endpoint": overrides.get("api_endpoint")
or os.getenv("MAVVRIK_API_ENDPOINT"),
"connection_id": overrides.get("connection_id")
or os.getenv("MAVVRIK_CONNECTION_ID"),
}
return {k: v for k, v in resolved.items() if v is not None}
raise NotImplementedError(
f"Provider '{provider}' not supported for Focus export configuration"
)

View file

@ -0,0 +1,345 @@
"""Mavvrik GCS destination for FOCUS export.
Flow:
1. GET /metrics/agent/ai/{connection_id}/upload-url → GCS signed URL
2. PUT <signed_url> with CSV content
"""
from __future__ import annotations
import gzip
from typing import Any, Optional
from urllib.parse import urlparse
from litellm._logging import verbose_logger
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
get_async_httpx_client,
httpxSpecialProvider,
)
from .base import FocusDestination, FocusTimeWindow
_MAVVRIK_ALLOWED_SUFFIXES = (".mavvrik.dev", ".mavvrik.ai", ".mavvrik.app")
# GCS requires intermediate chunks to be a multiple of 256 KB.
# 8 MB gives a good balance between round-trips and memory pressure.
_GCS_CHUNK_SIZE = 8 * 1024 * 1024 # 8 MB
def _validate_api_endpoint(api_endpoint: str) -> None:
if not api_endpoint.startswith("https://"):
raise ValueError("MAVVRIK_API_ENDPOINT must be an HTTPS URL")
hostname = (urlparse(api_endpoint).hostname or "").lower()
if not any(hostname.endswith(suffix) for suffix in _MAVVRIK_ALLOWED_SUFFIXES):
raise ValueError(
"MAVVRIK_API_ENDPOINT host must be a Mavvrik domain "
"(e.g. https://api.mavvrik.dev/<tenant_id>)"
)
def _validate_gcs_url(url: str, label: str) -> None:
parsed = urlparse(url)
if parsed.scheme != "https":
raise ValueError(
f"Mavvrik FOCUS destination: {label} must be HTTPS, got scheme '{parsed.scheme}'"
)
hostname = (parsed.hostname or "").lower()
if not (
hostname == "storage.googleapis.com"
or hostname.endswith(".storage.googleapis.com")
):
raise ValueError(
f"Mavvrik FOCUS destination: {label} must be a GCS endpoint "
f"(storage.googleapis.com), got '{hostname}'"
)
class FocusMavvrikDestination(FocusDestination):
"""Upload FOCUS CSV exports to Mavvrik via GCS signed URL."""
def __init__(
self,
*,
prefix: str,
config: Optional[dict[str, Any]] = None,
) -> None:
config = config or {}
api_key = config.get("api_key")
api_endpoint = config.get("api_endpoint")
connection_id = config.get("connection_id")
if not api_key:
raise ValueError(
"MAVVRIK_API_KEY must be provided for Mavvrik FOCUS destination "
"(set MAVVRIK_API_KEY env var or pass in destination_config)"
)
if not api_endpoint:
raise ValueError(
"MAVVRIK_API_ENDPOINT must be provided for Mavvrik FOCUS destination "
"(set MAVVRIK_API_ENDPOINT env var or pass in destination_config)"
)
if not connection_id:
raise ValueError(
"MAVVRIK_CONNECTION_ID must be provided for Mavvrik FOCUS destination "
"(set MAVVRIK_CONNECTION_ID env var or pass in destination_config)"
)
_validate_api_endpoint(api_endpoint)
self.api_key = api_key
self.api_endpoint = api_endpoint.rstrip("/")
self.connection_id = connection_id
self.prefix = prefix
self._http: AsyncHTTPHandler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.LoggingCallback
)
self._registered = False
@property
def _agent_url(self) -> str:
return f"{self.api_endpoint}/metrics/agent/ai/{self.connection_id}"
@property
def _upload_url_endpoint(self) -> str:
return f"{self.api_endpoint}/metrics/agent/ai/{self.connection_id}/upload-url"
@property
def _auth_headers(self) -> dict[str, str]:
return {"Content-Type": "application/json", "x-api-key": self.api_key}
async def _ensure_registered(self) -> Optional[int]:
"""POST agent endpoint to register/initialize the connector (once per instance).
Returns metricsMarker from the Mavvrik response — the last date index
Mavvrik has successfully processed. Used by the logger to catch up any
dates that were missed due to previous export failures.
Returns None if the connector was already registered (cached).
"""
if self._registered:
return None
resp = await self._http.client.request(
method="POST",
url=self._agent_url,
headers=self._auth_headers,
json={"name": self.connection_id},
timeout=30.0,
)
if resp.status_code == 410:
# Connector has been disconnected in Mavvrik — reset flag so next
# delivery attempt re-registers after it becomes active again.
self._registered = False
raise RuntimeError(
"Mavvrik FOCUS destination: connector is disconnected (410). "
"Re-enable the connection in the Mavvrik dashboard."
)
if resp.status_code >= 400:
raise RuntimeError(
f"Mavvrik FOCUS destination: register failed "
f"({resp.status_code}): {resp.text[:200]}"
)
self._registered = True
metrics_marker = resp.json().get("metricsMarker", 0)
verbose_logger.debug(
"Mavvrik FOCUS destination: connector registered (metricsMarker=%s)",
metrics_marker,
)
return metrics_marker
async def _get_signed_url(self, date_str: str) -> str:
"""GET upload-url endpoint → GCS signed URL for the given date."""
params = {"name": date_str, "type": "metrics", "datetime": date_str}
resp = await self._http.client.request(
method="GET",
url=self._upload_url_endpoint,
headers=self._auth_headers,
params=params,
timeout=30.0,
)
if resp.status_code >= 400:
raise RuntimeError(
f"Mavvrik FOCUS destination: failed to get signed URL "
f"({resp.status_code}): {resp.text[:200]}"
)
signed_url = resp.json().get("url")
if not signed_url:
raise RuntimeError(
f"Mavvrik FOCUS destination: response missing 'url' field: {resp.json()}"
)
_validate_gcs_url(signed_url, "signed URL")
verbose_logger.debug(
"Mavvrik FOCUS destination: got signed URL for date %s", date_str
)
return signed_url
async def _upload_to_gcs(self, signed_url: str, content: bytes) -> None:
"""Upload gzip-compressed CSV to GCS via chunked resumable upload.
The full CSV is gzip-compressed first, then uploaded in _GCS_CHUNK_SIZE
chunks using the GCS resumable upload protocol. GCS assembles the chunks
server-side into a single complete object — the bucket receives one file
regardless of how many chunks were sent.
Intermediate chunks: Content-Range: bytes X-Y/* → expect 308
Final chunk: Content-Range: bytes X-Y/T → expect 200/201
This handles exports larger than available memory for a single PUT while
keeping the destination code self-contained (no changes to the FOCUS
pipeline upstream).
"""
gzip_bytes = gzip.compress(content)
total = len(gzip_bytes)
# Step 1: initiate resumable upload session
metadata = b'{"contentEncoding":"gzip","contentDisposition":"attachment"}'
init_resp = await self._http.client.request(
method="POST",
url=signed_url,
headers={
"Content-Type": "application/gzip",
"x-goog-resumable": "start",
},
content=metadata,
timeout=30.0,
)
if init_resp.status_code not in (200, 201):
raise RuntimeError(
f"Mavvrik FOCUS destination: GCS session init failed "
f"({init_resp.status_code}): {init_resp.text[:400]}"
)
session_uri = init_resp.headers.get("Location")
if not session_uri:
raise RuntimeError(
"Mavvrik FOCUS destination: GCS session init missing Location header"
)
_validate_gcs_url(session_uri, "session URI")
verbose_logger.debug(
"Mavvrik FOCUS destination: GCS session started, uploading %d gzip bytes "
"in %d chunk(s)",
total,
max(1, -(-total // _GCS_CHUNK_SIZE)), # ceiling division
)
# Step 2: upload in chunks; cancel session on any failure to avoid
# lingering GCS sessions (they stay open for ~1 week otherwise).
offset = 0
try:
while offset < total:
chunk = gzip_bytes[offset : offset + _GCS_CHUNK_SIZE]
chunk_end = offset + len(chunk) - 1
is_final = (offset + len(chunk)) >= total
content_range = (
f"bytes {offset}-{chunk_end}/{total}"
if is_final
else f"bytes {offset}-{chunk_end}/*"
)
expected_statuses = {200, 201} if is_final else {308}
resp = await self._http.client.request(
method="PUT",
url=session_uri,
headers={
"Content-Type": "application/gzip",
"Content-Range": content_range,
},
content=chunk,
timeout=120.0,
)
if resp.status_code not in expected_statuses:
raise RuntimeError(
f"Mavvrik FOCUS destination: GCS chunk upload failed "
f"(chunk offset={offset}, expected={expected_statuses}, "
f"got={resp.status_code}): {resp.text[:400]}"
)
offset += len(chunk)
verbose_logger.debug(
"Mavvrik FOCUS destination: uploaded chunk offset=%d/%d",
offset,
total,
)
except Exception:
# Cancel the open GCS session so it doesn't linger for up to 1 week.
try:
await self._http.client.request(
method="DELETE", url=session_uri, timeout=10.0
)
verbose_logger.debug(
"Mavvrik FOCUS destination: cancelled GCS session after error"
)
except Exception:
pass
raise
async def get_metrics_marker(self) -> Optional[int]:
"""Register with Mavvrik and return the current metricsMarker.
The metricsMarker is a Unix timestamp (seconds) representing the last
date Mavvrik has successfully ingested. Called on every scheduled run
so the logger can detect and catch up any dates missed due to previous
export failures.
Always calls the Mavvrik register API — unlike deliver() which skips
registration once _registered is True, catch-up requires a fresh
marker value on every run.
"""
resp = await self._http.client.request(
method="POST",
url=self._agent_url,
headers=self._auth_headers,
json={"name": self.connection_id},
timeout=30.0,
)
if resp.status_code == 410:
self._registered = False
raise RuntimeError(
"Mavvrik FOCUS destination: connector is disconnected (410). "
"Re-enable the connection in the Mavvrik dashboard."
)
if resp.status_code >= 400:
raise RuntimeError(
f"Mavvrik FOCUS destination: register failed "
f"({resp.status_code}): {resp.text[:200]}"
)
self._registered = True
metrics_marker = resp.json().get("metricsMarker", 0)
verbose_logger.debug(
"Mavvrik FOCUS destination: got metricsMarker=%s", metrics_marker
)
return metrics_marker
async def deliver(
self,
*,
content: bytes,
time_window: FocusTimeWindow,
filename: str,
) -> None:
"""Upload FOCUS CSV to Mavvrik via GCS signed URL.
Uses the start date of the time window as the object date key.
"""
if not content:
verbose_logger.debug(
"Mavvrik FOCUS destination: empty content, skipping upload"
)
return
date_str = time_window.start_time.strftime("%Y-%m-%d")
verbose_logger.debug(
"Mavvrik FOCUS destination: uploading %d bytes for date=%s (%s)",
len(content),
date_str,
filename,
)
await self._ensure_registered()
signed_url = await self._get_signed_url(date_str)
await self._upload_to_gcs(signed_url, content)
verbose_logger.debug(
"Mavvrik FOCUS destination: upload complete for date=%s", date_str
)

View file

@ -0,0 +1,272 @@
"""MavvrikFocusLogger — FOCUS-based Mavvrik export logger.
Usage in config.yaml:
litellm_settings:
callbacks: ["mavvrik"]
Required env vars:
MAVVRIK_API_KEY
MAVVRIK_API_ENDPOINT
MAVVRIK_CONNECTION_ID
Optional env vars:
MAVVRIK_FOCUS_MAX_ROWS — row cap per export window (default: 500000)
Only daily frequency is supported. The Mavvrik ingestion protocol stores one
file per calendar date (metrics/YYYY-MM-DD). Hourly or interval exports would
overwrite each other within the same day, producing incomplete data.
"""
from __future__ import annotations
import os
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, List, Optional
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import MAVVRIK_FOCUS_EXPORT_JOB_NAME
from litellm.integrations.focus.destinations.base import FocusTimeWindow
from litellm.integrations.focus.focus_logger import FocusLogger
if TYPE_CHECKING:
from apscheduler.schedulers.asyncio import AsyncIOScheduler
else:
AsyncIOScheduler = Any
def _parse_metrics_marker(
marker: Optional[object],
) -> Optional[datetime]:
"""Parse metricsMarker from Mavvrik register response into a UTC datetime.
Handles both formats Mavvrik may return:
- Unix timestamp (int/float): e.g. 1749340800
- ISO date string: e.g. "2026-06-09" or "2026-06-09T00:00:00Z"
Returns None for falsy values (0, None, empty string) which indicate
no data has been ingested yet.
"""
if not marker:
return None
try:
if isinstance(marker, (int, float)):
return datetime.fromtimestamp(float(marker), tz=timezone.utc).replace(
hour=0, minute=0, second=0, microsecond=0
)
if isinstance(marker, str):
marker = marker.strip()
if not marker:
return None
# Try ISO date first (YYYY-MM-DD), then full ISO datetime
for fmt in ("%Y-%m-%d", "%Y-%m-%dT%H:%M:%SZ", "%Y-%m-%dT%H:%M:%S"):
try:
return datetime.strptime(marker, fmt).replace(tzinfo=timezone.utc)
except ValueError:
continue
except Exception:
pass
verbose_proxy_logger.warning(
"Mavvrik FOCUS: could not parse metricsMarker %r — skipping catch-up", marker
)
return None
class MavvrikFocusLogger(FocusLogger):
"""FOCUS-based export logger that routes to the Mavvrik destination."""
def __init__(self, **kwargs: Any) -> None:
frequency = os.getenv("MAVVRIK_FOCUS_FREQUENCY", "daily").lower()
if frequency != "daily":
raise ValueError(
f"MAVVRIK_FOCUS_FREQUENCY='{frequency}' is not supported. "
"Only 'daily' is allowed -- the Mavvrik ingestion protocol stores one "
"file per calendar date (metrics/YYYY-MM-DD). Hourly or interval "
"exports would overwrite each other within the same day."
)
super().__init__(
provider="mavvrik",
export_format="csv",
frequency="daily",
prefix="mavvrik_focus_exports",
destination_config={
"api_key": os.getenv("MAVVRIK_API_KEY"),
"api_endpoint": os.getenv("MAVVRIK_API_ENDPOINT"),
"connection_id": os.getenv("MAVVRIK_CONNECTION_ID"),
},
**kwargs,
)
raw = os.getenv("MAVVRIK_FOCUS_MAX_ROWS")
self._max_rows: Optional[int] = int(raw) if raw else 500_000
async def _export_window(
self,
*,
window: FocusTimeWindow,
limit: Optional[int],
) -> None:
"""Export with Mavvrik row cap applied when no explicit limit is passed."""
effective_limit = limit if limit is not None else self._max_rows
engine = self._ensure_engine()
data = await engine._database.get_usage_data(
limit=effective_limit,
start_time_utc=window.start_time,
end_time_utc=window.end_time,
)
if effective_limit is not None and len(data) >= effective_limit:
verbose_proxy_logger.warning(
"Mavvrik FOCUS export: row cap reached (%d rows). "
"Some data for window %s→%s may be excluded. "
"Increase MAVVRIK_FOCUS_MAX_ROWS to export all rows.",
effective_limit,
window.start_time.date(),
window.end_time.date(),
)
if data.is_empty():
verbose_proxy_logger.debug(
"Mavvrik FOCUS export: no usage data for window %s", window
)
return
normalized = engine._transformer.transform(data)
if normalized.is_empty():
return
payload = engine._serializer.serialize(normalized)
if not payload:
return
await engine._destination.deliver(
content=payload,
time_window=window,
filename=engine._build_filename(window),
)
# Maximum number of days to catch up in a single run. Prevents runaway
# loops if the connector was disabled for a long time, and avoids querying
# data that has likely been cleaned up from LiteLLM_DailyUserSpend.
_MAX_CATCHUP_DAYS = 7
async def _run_scheduled_export(self) -> None:
"""Export today's window, catching up any dates Mavvrik has not yet received.
On each run:
1. Register with Mavvrik → get metricsMarker (last successfully ingested date)
2. If metricsMarker is behind yesterday, catch up missed dates (capped at
_MAX_CATCHUP_DAYS to avoid runaway loops on long outages)
3. Export yesterday (today's daily window)
This ensures a failed export on day N is automatically retried on day N+1
without any manual intervention.
"""
engine = self._ensure_engine()
from litellm.integrations.focus.destinations.mavvrik_destination import ( # noqa: PLC0415
FocusMavvrikDestination,
)
destination = engine._destination
if not isinstance(destination, FocusMavvrikDestination):
await super()._run_scheduled_export()
return
# Register and get the last date Mavvrik has processed.
# metricsMarker may be a Unix timestamp (int/float) or an ISO date string.
marker = await destination.get_metrics_marker()
now = datetime.now(timezone.utc)
yesterday = now.replace(hour=0, minute=0, second=0, microsecond=0) - timedelta(
days=1
)
last_ingested = _parse_metrics_marker(marker)
# Catch up missed dates, capped at _MAX_CATCHUP_DAYS
if last_ingested and last_ingested < yesterday:
# Never go further back than _MAX_CATCHUP_DAYS from yesterday
earliest_catchup = yesterday - timedelta(days=self._MAX_CATCHUP_DAYS - 1)
catch_up_date = max(last_ingested + timedelta(days=1), earliest_catchup)
if last_ingested + timedelta(days=1) < earliest_catchup:
verbose_proxy_logger.warning(
"Mavvrik FOCUS export: metricsMarker is more than %d days behind "
"(%s). Catching up from %s only; earlier data will not be re-exported.",
self._MAX_CATCHUP_DAYS,
last_ingested.date(),
catch_up_date.date(),
)
while catch_up_date < yesterday:
verbose_proxy_logger.info(
"Mavvrik FOCUS export: catching up missed date %s",
catch_up_date.date(),
)
window = FocusTimeWindow(
start_time=catch_up_date,
end_time=catch_up_date + timedelta(days=1),
frequency="daily",
)
await self._export_window(window=window, limit=None)
catch_up_date += timedelta(days=1)
# Export yesterday's window (the normal daily run)
window = FocusTimeWindow(
start_time=yesterday,
end_time=yesterday + timedelta(days=1),
frequency="daily",
)
await self._export_window(window=window, limit=None)
async def initialize_mavvrik_focus_export_job(self) -> None:
"""Scheduler entry point — uses Mavvrik-specific pod-lock key."""
from litellm.proxy.proxy_server import proxy_logging_obj # noqa: PLC0415
pod_lock_manager = None
if proxy_logging_obj is not None:
writer = getattr(proxy_logging_obj, "db_spend_update_writer", None)
if writer is not None:
pod_lock_manager = getattr(writer, "pod_lock_manager", None)
if pod_lock_manager and pod_lock_manager.redis_cache:
acquired = await pod_lock_manager.acquire_lock(
cronjob_id=MAVVRIK_FOCUS_EXPORT_JOB_NAME
)
if not acquired:
verbose_proxy_logger.debug(
"Mavvrik FOCUS export: unable to acquire pod lock"
)
return
try:
await self._run_scheduled_export()
finally:
await pod_lock_manager.release_lock(
cronjob_id=MAVVRIK_FOCUS_EXPORT_JOB_NAME
)
else:
await self._run_scheduled_export()
@staticmethod
async def init_mavvrik_focus_background_job(
scheduler: AsyncIOScheduler,
) -> None:
"""Register the Mavvrik FOCUS export job on the provided scheduler."""
loggers: List[MavvrikFocusLogger] = [
cb
for cb in litellm.logging_callback_manager.get_custom_loggers_for_type(
callback_type=MavvrikFocusLogger
)
if type(cb) is MavvrikFocusLogger
]
if not loggers:
verbose_proxy_logger.debug(
"No MavvrikFocusLogger registered; skipping scheduler"
)
return
logger = loggers[0]
trigger_kwargs = logger._build_scheduler_trigger()
scheduler.add_job( # type: ignore[attr-defined]
logger.initialize_mavvrik_focus_export_job,
id=MAVVRIK_FOCUS_EXPORT_JOB_NAME,
replace_existing=True,
**trigger_kwargs,
)
verbose_proxy_logger.info(
"mavvrik_focus: background export job scheduled (%s)", trigger_kwargs
)

View file

@ -0,0 +1,10 @@
"""
New Relic AI Monitoring Integration for LiteLLM
This module provides integration with New Relic's AI Monitoring feature to track
LLM requests, responses, and usage metrics.
"""
from litellm.integrations.newrelic.newrelic import NewRelicLogger
__all__ = ["NewRelicLogger"]

View file

@ -0,0 +1,926 @@
"""
New Relic AI Monitoring Integration for LiteLLM
This module provides integration with New Relic's AI Monitoring feature to track
LLM requests, responses, and usage metrics.
Environment Variables (consumed by the New Relic agent at process bootstrap -
set via container env, or before invoking `newrelic-admin run-program`):
NEW_RELIC_LICENSE_KEY: Your New Relic license key (required)
NEW_RELIC_APP_NAME: Your application name (required)
UI- and runtime-toggleable:
NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED: Whether to record message
content (optional, default: true)
Configuration:
Message logging can be controlled via (both must agree to record):
1. turn_off_message_logging parameter - pass via callback initialization or config YAML
2. NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED env var
Default behavior: Messages ARE recorded unless explicitly disabled by either method
Either method can disable recording - both must enable for recording to occur
Usage - Python SDK:
import litellm
litellm.callbacks = ["newrelic"]
# Or with explicit configuration:
from litellm.integrations.newrelic import NewRelicLogger
litellm.callbacks = [NewRelicLogger(turn_off_message_logging=True)]
Usage - Proxy Server (config.yaml):
litellm_settings:
callbacks: ["newrelic"]
newrelic_params:
turn_off_message_logging: true # Disable message content recording
# Or disable via environment variable:
# export NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED=false
# Ensure New Relic agent is initialized (use newrelic-admin or initialize manually)
# newrelic-admin run-program python your_app.py
"""
import json
import os
import threading
import time
import uuid
from typing import Any, Dict, List, Optional, Tuple, Union
import litellm
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.redact_messages import should_redact_message_logging
from litellm.types.integrations.newrelic import NewRelicInitParams
from litellm.types.integrations.base_health_check import IntegrationHealthCheckStatus
from litellm.types.utils import ModelResponse, Message, StandardLoggingPayload
try:
import newrelic.agent as _newrelic_agent
except ImportError:
_newrelic_agent = None # type: ignore
class NewRelicLogger(CustomLogger):
"""
New Relic logger for LiteLLM to send AI monitoring events.
This logger creates two types of New Relic custom events:
1. LlmChatCompletionSummary - One per completion request
2. LlmChatCompletionMessage - One per message (request and response)
"""
# Class-level state for supportability metric emission, shared across all instances.
# Protected by _metric_lock to ensure thread-safe access.
_last_metric_emission_time: float = 0.0
_metric_lock = threading.Lock()
def __init__(self, **kwargs):
#########################################################
# Handle newrelic_params set as litellm.newrelic_params
#########################################################
dict_newrelic_params = self._get_newrelic_params()
# Use setdefault so constructor kwargs take priority over global params.
# model_dump() always returns all fields (including defaults), so update()
# would silently overwrite explicit constructor args like turn_off_message_logging=True.
for k, v in dict_newrelic_params.items():
kwargs.setdefault(k, v)
# CustomLogger.__init__ will set self.turn_off_message_logging from kwargs
super().__init__(**kwargs)
# Check for required environment variables
self.license_key = os.getenv("NEW_RELIC_LICENSE_KEY")
self.app_name = os.getenv("NEW_RELIC_APP_NAME")
# Validate configuration
if not self.license_key or not self.app_name:
verbose_logger.warning(
"New Relic integration requires NEW_RELIC_LICENSE_KEY and "
"NEW_RELIC_APP_NAME environment variables. Integration will be disabled."
)
self.enabled = False
elif _newrelic_agent is None:
verbose_logger.error(
"New Relic Python agent not installed. Review the New Relic integration documentation at https://docs.litellm.ai/docs/observability/newrelic."
)
self.enabled = False
else:
try:
# timeout=0 forces non-blocking startup: the agent connects in a
# background thread regardless of newrelic.ini / NEW_RELIC_STARTUP_TIMEOUT.
_newrelic_agent.register_application(timeout=0)
self.enabled = True
verbose_logger.info(
f"New Relic AI Monitoring initialized for app: {self.app_name}, "
f"content recording: {self.record_content}"
)
except Exception as e:
verbose_logger.error(
f"Failed to initialize New Relic agent: {e}. "
"Integration will be disabled."
)
self.enabled = False
def _get_newrelic_params(self) -> Dict:
"""
Get the newrelic_params from litellm.newrelic_params
These are params specific to initializing the NewRelicLogger e.g. turn_off_message_logging
"""
dict_newrelic_params: Dict = {}
if litellm.newrelic_params is not None:
if isinstance(litellm.newrelic_params, NewRelicInitParams):
dict_newrelic_params = litellm.newrelic_params.model_dump()
elif isinstance(litellm.newrelic_params, Dict):
# only allow params that are of NewRelicInitParams
dict_newrelic_params = NewRelicInitParams(
**litellm.newrelic_params
).model_dump()
return dict_newrelic_params
@property
def record_content(self) -> bool:
"""Whether to record message content in New Relic.
Both turn_off_message_logging param AND NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED
env var must agree to record content. If either disables recording, content will not
be recorded. Read at call time so UI config changes take effect without a restart.
Default: True (record content) unless explicitly disabled by either method.
"""
return (not self.turn_off_message_logging) and self._parse_bool_env(
"NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED", True
)
def _parse_bool_env(self, var_name: str, default: bool = False) -> bool:
"""Parse a boolean environment variable.
Accepts true/false, 1/0, yes/no, on/off (case-insensitive,
whitespace-tolerant) — matching the convention used in
``litellm/__init__.py`` and the standard library's
``configparser.BOOLEAN_STATES``. Unrecognised values log a
warning and fall back to ``default`` rather than silently
flipping user intent.
"""
raw = os.getenv(var_name)
if not raw:
return default
value = raw.strip().lower()
if value in ("1", "true", "yes", "on"):
return True
if value in ("0", "false", "no", "off"):
return False
verbose_logger.warning(
f"{var_name}={raw!r} is not a recognised boolean "
f"(accepts true/false, 1/0, yes/no, on/off). "
f"Falling back to default ({default})."
)
return default
def _get_litellm_version(self) -> str:
"""
Get litellm version for supportability metrics.
Returns:
Version string (e.g., "1.80.0") or "unknown" if unable to determine
"""
try:
from importlib.metadata import version
return version("litellm")
except Exception as e:
verbose_logger.warning(f"Unable to determine litellm version: {e}")
return "unknown"
def _emit_supportability_metric(self):
"""
Emit New Relic supportability metric for LiteLLM usage.
Per spec, this metric should be emitted at least once every 27 hours
to indicate the library is in use. Format:
Supportability/Python/ML/LiteLLM/{version}
This method updates _last_metric_emission_time and should
be called within a lock when checking periodic emission.
"""
try:
litellm_version = self._get_litellm_version()
metric_name = f"Supportability/Python/ML/LiteLLM/{litellm_version}"
# Record metric with value of 1 (will be aggregated by New Relic)
app = _newrelic_agent.application()
# Always update the timestamp so the 27-hour back-off applies
# regardless of whether the app is ready, preventing lock contention
# on every request when the agent is slow to register or never starts.
NewRelicLogger._last_metric_emission_time = time.time()
if app and app.enabled:
app.record_custom_metric(metric_name, 1)
verbose_logger.info(
f"Emitted New Relic supportability metric: {metric_name}"
)
else:
verbose_logger.info(
"New Relic application is not enabled; skipping metric recording."
)
except Exception as e:
verbose_logger.warning(f"Failed to emit supportability metric: {e}")
def _check_and_emit_periodic_metric(self):
"""
Check if 27 hours have passed since last metric emission and re-emit if needed.
Uses a mutex to ensure only one thread emits the metric even if multiple
requests are being processed concurrently.
"""
# Quick check without lock to avoid unnecessary locking
current_time = time.time()
time_since_last_emission = (
current_time - NewRelicLogger._last_metric_emission_time
)
if time_since_last_emission >= 97200: # 27 hours = 97200 seconds
# Acquire lock to ensure only one thread emits
with NewRelicLogger._metric_lock:
# Double-check inside lock in case another thread just emitted
current_time = time.time()
time_since_last_emission = (
current_time - NewRelicLogger._last_metric_emission_time
)
if time_since_last_emission >= 97200:
self._emit_supportability_metric()
def _get_trace_context(
self,
kwargs: Dict,
standard_logging_object: Optional[StandardLoggingPayload] = None,
) -> str:
"""
Get the New Relic trace ID for AI monitoring events.
This integration runs in LiteLLM's async logging worker, outside the
New Relic agent's current transaction. Because we can't call
`newrelic.agent.current_trace_id()` to let the agent populate the
trace_id on AIM custom events, we manually simulate what the agent
would do. An AIM event without a trace_id is malformed per the NR
schema, so this method always returns a valid string.
Resolution order:
1. W3C traceparent header (litellm_params.metadata.headers.traceparent) -
what the agent would link to if we were in-transaction.
2. StandardLoggingPayload.trace_id - LiteLLM's internal trace for
retry/fallback grouping.
3. Generated UUID - synthetic grouping key when upstream context is
absent or parsing it fails.
Span IDs are intentionally not emitted: any span ID recoverable from
the inbound traceparent is the caller's parent span, not ours.
Returns:
trace_id: always a non-empty string.
"""
trace_id: Optional[str] = None
try:
litellm_params = kwargs.get("litellm_params") or {}
metadata = litellm_params.get("metadata") or {}
headers = metadata.get("headers") or {}
# Normalize header key lookup to be case-insensitive per W3C spec
traceparent = next(
(v for k, v in headers.items() if k.lower() == "traceparent"), None
)
if traceparent:
# Extract trace_id from traceparent header if available
# traceparent format: "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-00"
parts = traceparent.split("-")
if len(parts) == 4:
trace_id = parts[1]
if not trace_id and standard_logging_object:
slo_trace_id = standard_logging_object.get("trace_id")
if slo_trace_id:
trace_id = slo_trace_id
except Exception as e:
verbose_logger.warning(
f"Unable to parse New Relic trace context from upstream sources: {e}"
)
if not trace_id:
trace_id = uuid.uuid4().hex
verbose_logger.debug(
f"New Relic trace_id not available from distributed tracing headers or "
f"StandardLoggingPayload. Generated trace_id={trace_id} for AI monitoring "
f"event grouping."
)
return trace_id
def _extract_completion_id(self, kwargs: Dict, response_obj: ModelResponse) -> str:
"""
Extract completion ID from kwargs or response_obj, or generate one.
"""
completion_id = None
if response_obj:
completion_id = response_obj.get("id")
if not completion_id:
completion_id = kwargs.get("litellm_call_id")
# If still not found, generate UUID and log warning per spec
if not completion_id:
completion_id = str(uuid.uuid4())
return completion_id
def _get_vendor(
self,
kwargs: Dict,
standard_logging_object: Optional[StandardLoggingPayload] = None,
) -> str:
"""Extract vendor/provider, preferring StandardLoggingPayload."""
if standard_logging_object:
vendor = standard_logging_object.get("custom_llm_provider")
if vendor:
return vendor
litellm_params = kwargs.get("litellm_params", {}) or {}
return litellm_params.get("custom_llm_provider") or "litellm"
def _get_model_names(
self,
kwargs: Dict,
response_obj: ModelResponse,
standard_logging_object: Optional[StandardLoggingPayload] = None,
) -> Tuple[str, str]:
"""
Extract request and response model names, preferring StandardLoggingPayload
for the request model.
Returns:
Tuple of (request_model, response_model)
"""
request_model = None
if standard_logging_object:
slo_model = standard_logging_object.get("model")
if slo_model:
request_model = str(slo_model)
if not request_model:
request_model = str(kwargs.get("model") or "unknown")
response_model: str = str(response_obj.get("model") or request_model)
return request_model, response_model
def _extract_usage(
self,
response_obj: ModelResponse,
standard_logging_object: Optional[StandardLoggingPayload] = None,
) -> Dict[str, int]:
"""Extract usage statistics, preferring StandardLoggingPayload."""
if standard_logging_object:
prompt = standard_logging_object.get("prompt_tokens")
completion = standard_logging_object.get("completion_tokens")
total = standard_logging_object.get("total_tokens")
if any(x is not None for x in [prompt, completion, total]):
return {
"prompt_tokens": prompt or 0,
"completion_tokens": completion or 0,
"total_tokens": total or 0,
}
usage = response_obj.get("usage", None)
if not usage:
return {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
return {
"prompt_tokens": usage.get("prompt_tokens") or 0,
"completion_tokens": usage.get("completion_tokens") or 0,
"total_tokens": usage.get("total_tokens") or 0,
}
def _get_finish_reason(self, response_obj: ModelResponse) -> str:
"""
Extract finish reason from first choice in the response.
Returns "unknown" if choices are not present or finish_reason is not found.
"""
choices = response_obj.get("choices") or []
if choices and len(choices) > 0:
return choices[0].get("finish_reason") or "unknown"
return "unknown"
def _to_epoch_ms(self, t: Any) -> float:
"""Convert a datetime or float timestamp to epoch milliseconds."""
if hasattr(t, "timestamp"):
return t.timestamp() * 1000.0
return float(t) * 1000.0
def _get_duration(
self,
kwargs: Dict,
start_time: Any,
end_time: Any,
standard_logging_object: Optional[StandardLoggingPayload] = None,
) -> Optional[float]:
"""
Extract duration in milliseconds.
Resolution order:
1. StandardLoggingPayload.response_time (already computed by LiteLLM)
2. llm_api_duration_ms from kwargs
3. Calculated from start_time and end_time
"""
if standard_logging_object:
response_time = standard_logging_object.get("response_time")
if response_time is not None:
return (
float(response_time) * 1000.0
) # SLO stores seconds; convert to ms
duration_ms = kwargs.get("llm_api_duration_ms")
if duration_ms is not None:
return float(duration_ms)
if start_time is not None and end_time is not None:
return self._to_epoch_ms(end_time) - self._to_epoch_ms(start_time)
return None
def _get_request_params(
self,
kwargs: Dict,
standard_logging_object: Optional[StandardLoggingPayload] = None,
) -> Dict[str, Any]:
"""
Extract request parameters like temperature and max_tokens, preferring
StandardLoggingPayload.model_parameters.
Returns dict with available parameters, omitting those not present.
"""
if standard_logging_object:
source_params = standard_logging_object.get("model_parameters") or {}
else:
source_params = kwargs.get("optional_params") or {}
params = {}
temperature = source_params.get("temperature")
if temperature is not None:
params["temperature"] = temperature
max_tokens = source_params.get("max_tokens")
if max_tokens is not None:
params["max_tokens"] = max_tokens
return params
def _extract_message_content(self, message: Union[Message, Dict]) -> str:
"""
Extract content from a message, handling various formats.
Handles tool calls, multimodal content (as JSON), and standard text content.
Returns empty string if content is None or missing.
"""
content = message.get("content")
# Handle tool calls
if message.get("tool_calls"):
try:
return json.dumps(message["tool_calls"])
except Exception:
return str(message["tool_calls"])
# Handle None or missing content
if content is None:
return ""
# Handle list content (multimodal)
if isinstance(content, list):
try:
return json.dumps(content)
except Exception:
return str(content)
# Handle non-string content
if not isinstance(content, str):
return str(content)
return content
def _extract_all_messages(
self,
kwargs: Dict,
response_obj: ModelResponse,
response_model: str,
vendor: str,
standard_logging_object: Optional[StandardLoggingPayload] = None,
) -> List[Dict[str, Any]]:
"""
Extract all messages (request + response) with sequence numbers and timestamps.
Processes request messages from StandardLoggingPayload.messages (preferred) or
kwargs["messages"] (fallback), and response messages from response_obj["choices"].
Assigns sequential numbers starting at 0.
Adds timestamps from StandardLoggingPayload (preferred) or kwargs if available
(converted to epoch milliseconds).
"""
messages = []
sequence = 0
# Extract timestamps, preferring StandardLoggingPayload
start_time = None
if standard_logging_object:
start_time = standard_logging_object.get("startTime")
if not start_time:
start_time = kwargs.get("start_time")
end_time = None
if standard_logging_object:
end_time = standard_logging_object.get("endTime")
if not end_time:
end_time = kwargs.get("end_time")
# Content is recorded only when the NR-specific switches allow it AND
# LiteLLM's wider redaction decision (turn_off_message_logging, dynamic
# params, headers) does not require redaction. Async streaming hands the
# callback an unredacted async_complete_streaming_response, so without
# this gate generated content would still reach NR even when the user
# has globally disabled message logging.
record_content = self.record_content and not should_redact_message_logging(
kwargs
)
# Extract request messages, preferring StandardLoggingPayload.
# SLO messages can be a string (serialized/redacted), so only use it when it's a list.
slo_messages = (
standard_logging_object.get("messages") if standard_logging_object else None
)
if isinstance(slo_messages, list):
request_messages = slo_messages
else:
request_messages = kwargs.get("messages") or []
for msg in request_messages:
message_data = {
"role": msg.get("role") or "user",
"sequence": sequence,
"response.model": response_model,
"vendor": vendor,
}
# Add timestamp for request message if available (convert to milliseconds)
if start_time is not None:
message_data["timestamp"] = int(self._to_epoch_ms(start_time))
if record_content:
message_data["content"] = self._extract_message_content(msg)
messages.append(message_data)
sequence += 1
# Extract response messages from choices
choices = response_obj.get("choices") or []
if choices and len(choices) > 0:
for choice in choices:
# Prefer "message" (non-streaming); fall back to "delta" (streaming-assembled)
message = choice.get("message", None) or choice.get("delta", None)
if message:
message_data = {
"role": message.get("role") or "assistant",
"sequence": sequence,
"response.model": response_model,
"vendor": vendor,
"is_response": True,
}
# Add timestamp for response message if available (convert to milliseconds)
if end_time is not None:
message_data["timestamp"] = int(self._to_epoch_ms(end_time))
if record_content:
message_data["content"] = self._extract_message_content(message)
messages.append(message_data)
sequence += 1
return messages
def _record_summary_event(
self,
request_id: str,
trace_id: Optional[str],
request_model: str,
response_model: str,
vendor: str,
finish_reason: str,
num_messages: int,
usage: Dict[str, int],
duration: Optional[float] = None,
request_params: Optional[Dict[str, Any]] = None,
):
"""Record LlmChatCompletionSummary event to New Relic."""
try:
event_data = {
"id": request_id,
"request_id": request_id,
"request.model": request_model,
"response.model": response_model,
"response.choices.finish_reason": finish_reason,
"response.number_of_messages": num_messages,
"vendor": vendor,
"ingest_source": "litellm",
"response.usage.prompt_tokens": usage["prompt_tokens"],
"response.usage.completion_tokens": usage["completion_tokens"],
"response.usage.total_tokens": usage["total_tokens"],
}
# Add optional attributes if present
if trace_id:
event_data["trace_id"] = trace_id
if duration is not None:
event_data["duration"] = duration
# Add request parameters if present
if request_params:
if "temperature" in request_params:
event_data["request.temperature"] = request_params["temperature"]
if "max_tokens" in request_params:
event_data["request.max_tokens"] = request_params["max_tokens"]
app = _newrelic_agent.application()
if app and app.enabled:
app.record_custom_event("LlmChatCompletionSummary", event_data)
else:
verbose_logger.warning(
"New Relic application is not enabled; skipping summary event recording."
)
except Exception as e:
verbose_logger.warning(f"Failed to record New Relic summary event: {e}")
self.handle_callback_failure("newrelic")
def _record_message_events(
self,
request_id: str,
llm_response_id: str,
trace_id: Optional[str],
messages: List[Dict[str, Any]],
):
"""Record LlmChatCompletionMessage events to New Relic.
Args:
request_id: Agent-generated UUID that links to Summary event's id
llm_response_id: LLM's response ID (e.g., "chatcmpl-...") for message id format
trace_id: Trace ID for distributed tracing (None if not available)
messages: List of message dicts to record
"""
try:
app = _newrelic_agent.application()
if not (app and app.enabled):
verbose_logger.warning(
"New Relic application is not enabled; skipping message event recording."
)
return
for message in messages:
sequence = message["sequence"]
event_data = {
"id": f"{llm_response_id}-{sequence}",
"request_id": request_id,
"completion_id": request_id,
"role": message["role"],
"sequence": sequence,
"response.model": message["response.model"],
"vendor": message["vendor"],
"ingest_source": "litellm",
"token_count": 0, # Per-message token counts are not available from LiteLLM
}
# Add trace context if available
if trace_id:
event_data["trace_id"] = trace_id
# Add content only if it was included in the message data
if "content" in message:
event_data["content"] = message["content"]
# Add is_response only if True (per spec, omit for request messages)
if message.get("is_response"):
event_data["is_response"] = True
# Forward actual request/response timestamp (ms) so NR uses the
# real LLM call window rather than the async-logger fire time.
# Requires newrelic>=11.2.0 which reads params["timestamp"] as
# the intrinsic event timestamp.
if "timestamp" in message:
event_data["timestamp"] = message["timestamp"]
app.record_custom_event("LlmChatCompletionMessage", event_data)
except Exception as e:
verbose_logger.warning(f"Failed to record New Relic message events: {e}")
self.handle_callback_failure("newrelic")
def _record_error_metric(self):
"""Record error metric to New Relic."""
try:
if not self.enabled:
return
self._check_and_emit_periodic_metric()
app = _newrelic_agent.application()
if app and app.enabled:
app.record_custom_metric("LLM/LiteLLM/Error", 1)
except Exception as e:
verbose_logger.warning(f"Failed to record New Relic error metric: {e}")
self.handle_callback_failure("newrelic")
def _process_success(
self,
kwargs: Dict,
response_obj: ModelResponse,
start_time: Optional[float] = None,
end_time: Optional[float] = None,
):
"""
Core logic for processing successful LLM calls.
Used by both sync and async success event handlers.
"""
# Early exit if not enabled
if not self.enabled:
return
# Check and emit periodic supportability metric if 27 hours have passed
self._check_and_emit_periodic_metric()
# Use StandardLoggingPayload where available for normalized, pre-computed values
standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get(
"standard_logging_object"
)
# Get trace context
trace_id = self._get_trace_context(kwargs, standard_logging_object)
# Generate unique request ID for this request (used as Summary event id)
request_id = str(uuid.uuid4())
# Extract data from response
llm_response_id = self._extract_completion_id(kwargs, response_obj)
vendor = self._get_vendor(kwargs, standard_logging_object)
request_model, response_model = self._get_model_names(
kwargs, response_obj, standard_logging_object
)
usage = self._extract_usage(response_obj, standard_logging_object)
finish_reason = self._get_finish_reason(response_obj)
# Extract additional summary event fields
duration = self._get_duration(
kwargs, start_time, end_time, standard_logging_object
)
request_params = self._get_request_params(kwargs, standard_logging_object)
# Extract all messages
messages = self._extract_all_messages(
kwargs, response_obj, response_model, vendor, standard_logging_object
)
# Record summary event
self._record_summary_event(
request_id=request_id,
trace_id=trace_id,
request_model=request_model,
response_model=response_model,
vendor=vendor,
finish_reason=finish_reason,
num_messages=len(messages),
usage=usage,
duration=duration,
request_params=request_params,
)
# Record message events
self._record_message_events(
request_id=request_id,
llm_response_id=llm_response_id,
trace_id=trace_id,
messages=messages,
)
async def async_health_check(self) -> IntegrationHealthCheckStatus:
"""
Check if the New Relic integration is healthy.
Verifies that the integration is enabled and the New Relic agent
has an active, connected application, then records a small
`LiteLLMConnectionTest` custom event so the user can confirm the
end-to-end pipeline in the New Relic UI via NRQL:
`SELECT * FROM LiteLLMConnectionTest SINCE 1 hour ago`.
The `LiteLLMConnectionTest` event type is intentionally outside the
`Llm*` family that AI Monitoring queries, so test events do not
appear in AI Monitoring dashboards.
"""
if not self.enabled:
return IntegrationHealthCheckStatus(
status="unhealthy",
error_message="New Relic integration is disabled. Check that "
"NEW_RELIC_LICENSE_KEY and NEW_RELIC_APP_NAME are set and the "
"newrelic package is installed.",
)
try:
app = _newrelic_agent.application()
if not (app and app.enabled):
return IntegrationHealthCheckStatus(
status="unhealthy",
error_message=(
"New Relic Python agent not installed. Review the New Relic integration documentation at https://docs.litellm.ai/docs/observability/newrelic."
),
)
app.record_custom_event(
"LiteLLMConnectionTest",
{
"is_test_event": True,
"app_name": self.app_name,
"source": "litellm-proxy",
"timestamp": time.time(),
},
)
return IntegrationHealthCheckStatus(status="healthy", error_message=None)
except Exception as e:
return IntegrationHealthCheckStatus(
status="unhealthy",
error_message=str(e),
)
# CustomLogger interface implementation
def log_pre_api_call(self, model, messages, kwargs):
"""Unused per spec."""
pass
def log_post_api_call(self, kwargs, response_obj, start_time, end_time):
"""Unused per spec."""
pass
def log_success_event(self, kwargs, response_obj, start_time, end_time):
"""
Main success path for non-streaming requests.
Note: New Relic's record_custom_event is synchronous but non-blocking
(in-memory operation), so it's safe to call from sync context.
"""
try:
self._process_success(kwargs, response_obj, start_time, end_time)
except Exception as e:
verbose_logger.warning(f"Error in New Relic log_success_event: {e}")
self.handle_callback_failure("newrelic")
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
"""
Main success path for async/streaming requests.
Note: New Relic's SDK is thread-safe and record_custom_event is fast,
so we can call it directly without asyncio.to_thread().
"""
try:
self._process_success(kwargs, response_obj, start_time, end_time)
except Exception as e:
verbose_logger.warning(f"Error in New Relic async_log_success_event: {e}")
self.handle_callback_failure("newrelic")
def log_failure_event(self, kwargs, response_obj, start_time, end_time):
"""
Log error metric for failed LLM calls (sync).
Per spec: Do not send AI events on failure, only record error metric.
"""
try:
self._record_error_metric()
except Exception as e:
verbose_logger.warning(f"Error in New Relic log_failure_event: {e}")
self.handle_callback_failure("newrelic")
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
"""
Log error metric for failed LLM calls (async).
Per spec: Do not send AI events on failure, only record error metric.
"""
try:
self._record_error_metric()
except Exception as e:
verbose_logger.warning(f"Error in New Relic async_log_failure_event: {e}")
self.handle_callback_failure("newrelic")

View file

@ -25,6 +25,7 @@ from litellm.integrations.datadog.datadog_metrics import DatadogMetricsLogger
from litellm.integrations.deepeval import DeepEvalLogger
from litellm.integrations.dotprompt import DotpromptManager
from litellm.integrations.focus.focus_logger import FocusLogger
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import MavvrikFocusLogger
from litellm.integrations.vantage.vantage_logger import VantageLogger
from litellm.integrations.galileo import GalileoObserve
from litellm.integrations.gcs_bucket.gcs_bucket import GCSBucketLogger
@ -39,6 +40,7 @@ from litellm.integrations.langsmith import LangsmithLogger
from litellm.integrations.litellm_agent import LiteLLMAgentModelResolver
from litellm.integrations.literal_ai import LiteralAILogger
from litellm.integrations.mlflow import MlflowLogger
from litellm.integrations.newrelic import NewRelicLogger
from litellm.integrations.openmeter import OpenMeterLogger
from litellm.integrations.opentelemetry import OpenTelemetry
from litellm.integrations.opik.opik import OpikLogger
@ -102,8 +104,10 @@ class CustomLoggerRegistry:
"gitlab": GitLabPromptManager,
"cloudzero": CloudZeroLogger,
"focus": FocusLogger,
"mavvrik": MavvrikFocusLogger,
"vantage": VantageLogger,
"posthog": PostHogLogger,
"newrelic": NewRelicLogger,
}
try:

View file

@ -131,6 +131,8 @@ def get_next_standardized_reset_time(
# Handle different time units
if unit == "d":
return _handle_day_reset(current_time, base_midnight, value, tz)
elif unit == "w":
return _handle_day_reset(current_time, base_midnight, value * 7, tz)
elif unit == "h":
return _handle_hour_reset(current_time, base_midnight, value)
elif unit == "m":

View file

@ -158,6 +158,7 @@ from ..integrations.litellm_agent import LiteLLMAgentModelResolver
from ..integrations.literal_ai import LiteralAILogger
from ..integrations.logfire_logger import LogfireLevel, LogfireLogger
from ..integrations.lunary import LunaryLogger
from ..integrations.newrelic import NewRelicLogger
from ..integrations.openmeter import OpenMeterLogger
from ..integrations.opik.opik import OpikLogger
from ..integrations.posthog import PostHogLogger
@ -3507,9 +3508,7 @@ class Logging(LiteLLMLoggingBaseClass):
else:
return None
def _handle_anthropic_messages_response_logging(
self, result: Any
) -> Union[ModelResponse, ResponsesAPIResponse]:
def _handle_anthropic_messages_response_logging(self, result: Any) -> ModelResponse:
"""
Handles logging for Anthropic messages responses.
@ -3528,15 +3527,14 @@ class Logging(LiteLLMLoggingBaseClass):
return result
elif isinstance(result, ModelResponse):
return result
elif isinstance(
if isinstance(
result,
(ResponseCompletedEvent, ResponseIncompleteEvent, ResponseFailedEvent),
):
# anthropic_messages() can route to OpenAI Responses API; in that path
# the assembled streaming result is one of these terminal events rather than
# a ModelResponse. Return the inner response so downstream handlers
# (_transform_usage_objects, normalize_logging_result) can process it.
return result.response
result = result.response
if isinstance(result, ResponsesAPIResponse):
return self._translate_responses_api_response_to_model_response(result)
httpx_response = self.model_call_details.get("httpx_response", None)
if httpx_response and isinstance(httpx_response, httpx.Response):
@ -3570,6 +3568,55 @@ class Logging(LiteLLMLoggingBaseClass):
)
return result
def _translate_responses_api_response_to_model_response(
self, result: ResponsesAPIResponse
) -> ModelResponse:
"""
Convert a Responses API response into a ModelResponse for spend_logs.
The proxy UI parses spend_log rows expecting chat-completion shape
(response.choices[0].message); a raw ResponsesAPIResponse dump (output[...])
would render as empty in the Logs tab. Translation also yields full
choices/message detail downstream consumers can rely on.
"""
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
)
try:
return LiteLLMResponsesTransformationHandler().transform_response(
model=self.model,
raw_response=result,
model_response=litellm.ModelResponse(),
logging_obj=self,
request_data={},
messages=[],
optional_params={},
litellm_params={},
encoding=litellm.encoding,
)
except Exception as e:
verbose_logger.debug(
"Responses API -> ModelResponse translation failed for "
"anthropic_messages logging (%s); falling back to minimal "
"usage-only ModelResponse to keep the spend_logs row.",
str(e),
)
model_response = litellm.ModelResponse()
model_response.model = self.model
usage = getattr(result, "usage", None)
if usage is not None and ResponseAPILoggingUtils._is_response_api_usage(
usage
):
setattr(
model_response,
"usage",
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
),
)
return model_response
def _handle_non_streaming_google_genai_generate_content_response_logging(
self, result: Any
) -> ModelResponse:
@ -4124,6 +4171,17 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
focus_logger = FocusLogger()
_in_memory_loggers.append(focus_logger)
return focus_logger # type: ignore
elif logging_integration == "mavvrik":
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import (
MavvrikFocusLogger,
)
for callback in _in_memory_loggers:
if type(callback) is MavvrikFocusLogger:
return callback # type: ignore
mavvrik_focus_logger = MavvrikFocusLogger()
_in_memory_loggers.append(mavvrik_focus_logger)
return mavvrik_focus_logger # type: ignore
elif logging_integration == "vantage":
from litellm.integrations.vantage.vantage_logger import VantageLogger
@ -4419,6 +4477,13 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
gitlab_logger = GitLabPromptManager(gitlab_config=gitlab_config)
_in_memory_loggers.append(gitlab_logger)
return gitlab_logger # type: ignore
elif logging_integration == "newrelic":
for callback in _in_memory_loggers:
if isinstance(callback, NewRelicLogger):
return callback # type: ignore
newrelic_logger = NewRelicLogger()
_in_memory_loggers.append(newrelic_logger)
return newrelic_logger # type: ignore
return None
except Exception as e:
verbose_logger.exception(
@ -4720,6 +4785,10 @@ def get_custom_logger_compatible_class( # noqa: PLR0915
for callback in _in_memory_loggers:
if isinstance(callback, SMTPEmailLogger):
return callback
elif logging_integration == "newrelic":
for callback in _in_memory_loggers:
if isinstance(callback, NewRelicLogger):
return callback
return None
except Exception as e:

View file

@ -4741,6 +4741,12 @@ class BedrockConverseMessagesProcessor:
guardContent={"text": {"text": element["text"]}}
)
_parts.append(_part)
elif element["type"] in ("grounding_source", "query"):
# Contextual grounding tags are guardrail metadata; the
# model only needs the underlying text, so render them
# as plain text on the generate path.
_part = BedrockContentBlock(text=element["text"])
_parts.append(_part)
elif element["type"] == "image_url":
format: Optional[str] = None
if isinstance(element["image_url"], dict):
@ -5173,6 +5179,12 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915
guardContent={"text": {"text": element["text"]}}
)
_parts.append(_part)
elif element["type"] in ("grounding_source", "query"):
# Contextual grounding tags are guardrail metadata; the
# model only needs the underlying text, so render them as
# plain text on the generate path.
_part = BedrockContentBlock(text=element["text"])
_parts.append(_part)
elif element["type"] == "image_url":
format: Optional[str] = None
if isinstance(element["image_url"], dict):

View file

@ -47,6 +47,7 @@ class RealTimeStreaming:
user_api_key_dict: Optional[Any] = None,
request_data: Optional[Dict] = None,
backend_uses_beta_protocol: Optional[bool] = None,
force_transcription_model: Optional[str] = None,
):
self.websocket = websocket
self.backend_ws = backend_ws
@ -100,6 +101,11 @@ class RealTimeStreaming:
self._flushing_pending_messages_until_setup: bool = False
self._pending_messages_until_setup: List[str] = []
self._pending_messages_byte_total: int = 0
# Whether this is a transcription-only session (session.type == "transcription",
# e.g. gpt-realtime-whisper). Such sessions must not be sent response.create and
# their input_audio_transcription.completed usage drives duration-based cost.
self._force_transcription_model = force_transcription_model
self._is_transcription_session: bool = force_transcription_model is not None
# Per-connection caps for pre-setup audio frames (message count + total bytes).
_MAX_BUFFERED_MESSAGES: int = 200
@ -211,6 +217,8 @@ class RealTimeStreaming:
self.session_tools = tools
# GA: session.type is required; log it for traceability but no action needed
verbose_logger.debug(f"Realtime session.type: {session.get('type')}")
if session.get("type") == "transcription":
self._is_transcription_session = True
except (json.JSONDecodeError, AttributeError, TypeError):
pass
@ -227,6 +235,55 @@ class RealTimeStreaming:
except (AttributeError, TypeError):
pass
def _detect_transcription_session_from_backend(
self, event_obj: Union[dict, OpenAIRealtimeEvents]
) -> None:
"""Flag transcription-only sessions from backend session events."""
try:
event_type = event_obj.get("type", "")
if event_type in (
"transcription_session.created",
"transcription_session.updated",
):
self._is_transcription_session = True
elif event_type in ("session.created", "session.updated"):
session = cast(dict, event_obj).get("session", {}) or {}
if session.get("type") == "transcription":
self._is_transcription_session = True
except (AttributeError, TypeError):
pass
def _capture_transcription_usage(
self, event_obj: Union[dict, OpenAIRealtimeEvents]
) -> None:
"""
Append a usage-only transcription completed event to the logged results so
the cost calculator can bill it by audio duration. The default logged event
types exclude this event, so it is captured here directly for transcription
sessions rather than widening logging for every realtime session. Only the
type and usage are kept — the transcript is already captured separately in
input_messages, so it is not duplicated into the response log here.
"""
try:
usage = event_obj.get("usage")
if usage is None:
return
# If this event type is already captured by store_message (e.g. the user
# logs all realtime events), don't append a second copy.
if self._should_store_message(event_obj):
return
self.messages.append(
cast(
OpenAIRealtimeEvents,
{
"type": "conversation.item.input_audio_transcription.completed",
"usage": usage,
},
)
)
except (AttributeError, TypeError):
pass
def _collect_tool_calls_from_response_done(
self, event_obj: Union[dict, OpenAIRealtimeEvents]
) -> None:
@ -287,6 +344,7 @@ class RealTimeStreaming:
backend, False if the provider transformation produced no output and
the message was effectively dropped.
"""
message = self._enforce_transcription_session_model(message)
if self.provider_config:
transformed = self.provider_config.transform_realtime_request(
message, self.model, self.session_configuration_request
@ -306,6 +364,80 @@ class RealTimeStreaming:
await self.backend_ws.send(message) # type: ignore[union-attr, attr-defined]
return True
def _enforce_transcription_session_model(self, message: str) -> str:
"""Force client transcription session updates to the authorized model.
`/v1/realtime?intent=transcription` may intentionally omit `model` from
the upstream URL for Azure compatibility, but the proxy still authorizes
a resolved LiteLLM model before opening the backend websocket. If a
client later sends a transcription `session.update`, any model embedded
in that update must be rewritten to the same authorized model instead of
allowing a post-auth model/deployment switch.
Normal realtime sessions keep their independent nested transcription
model behavior because `_force_transcription_model` is only set for
transcription-intent websocket routes.
"""
if self._force_transcription_model is None:
return message
try:
message_obj = json.loads(message)
except (json.JSONDecodeError, TypeError):
return message
if message_obj.get("type") not in (
"session.update",
"transcription_session.update",
):
return message
session = message_obj.get("session")
if not isinstance(session, dict):
return message
if session.get("type") == "transcription":
self._is_transcription_session = True
authorized_model = self._force_transcription_model
changed = False
transcription = session.get("input_audio_transcription")
if (
isinstance(transcription, dict)
and transcription.get("model") != authorized_model
):
session["input_audio_transcription"] = {
**transcription,
"model": authorized_model,
}
changed = True
audio = session.get("audio")
if isinstance(audio, dict):
audio_input = audio.get("input")
if isinstance(audio_input, dict):
nested_transcription = audio_input.get("transcription")
if (
isinstance(nested_transcription, dict)
and nested_transcription.get("model") != authorized_model
):
session["audio"] = {
**audio,
"input": {
**audio_input,
"transcription": {
**nested_transcription,
"model": authorized_model,
},
},
}
changed = True
if not changed:
return message
return json.dumps(message_obj)
def _uses_deferred_backend_setup(self) -> bool:
"""True when setup is deferred until the client's first session.update."""
if self.provider_config is None:
@ -792,6 +924,8 @@ class RealTimeStreaming:
"""
event_type = event_obj.get("type")
self._detect_transcription_session_from_backend(event_obj)
# Send session.created to the client FIRST so it stays in sync, then inject
# the disable-auto-response session.update; otherwise a backend error could
# reach the client before it sees session.created.
@ -809,6 +943,14 @@ class RealTimeStreaming:
self._collect_user_input_from_backend_event(event_obj)
self.store_message(event_obj)
await self.websocket.send_text(raw_response)
# Transcription-only sessions (e.g. gpt-realtime-whisper) have no
# assistant turn: capture audio-duration usage for cost and never
# trigger response.create.
if self._is_transcription_session:
self._capture_transcription_usage(event_obj)
return True
blocked = await self.run_realtime_guardrails(
transcript,
item_id=event_obj.get("item_id"),

View file

@ -469,12 +469,16 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
if should_start_new_block and not self.sent_content_block_finish:
# Queue the sequence: content_block_stop -> content_block_start
# For text blocks the trigger chunk is not emitted as a separate
# delta because content_block_start carries the information.
# For tool_use blocks we must also emit the trigger chunk's delta
# when it carries input_json_delta data, because some providers
# (e.g. xAI, Gemini) include tool arguments in the same streaming
# chunk as the function name/id.
# -> (optionally) the trigger chunk's delta.
#
# The synthesized content_block_start always carries an
# empty body, so the chunk that *triggered* the transition
# also carries the new block's first delta. It must be
# re-emitted or the first token of the new block is lost.
# This applies to text_delta and thinking_delta (the first
# non-empty text/thinking token) as well as input_json_delta
# (providers like xAI/Gemini bundle tool arguments with the
# function name/id in a single chunk).
# 1. Stop current content block
self.chunk_queue.append(
@ -493,14 +497,9 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
}
)
# 3. If the trigger chunk carries tool argument data, queue it
# so the input_json_delta is not silently dropped.
if (
processed_chunk.get("type") == "content_block_delta"
and isinstance(processed_chunk.get("delta"), dict)
and processed_chunk["delta"].get("type") == "input_json_delta"
and processed_chunk["delta"].get("partial_json")
):
# 3. If the trigger chunk carries delta content, queue it
# so the first delta of the new block is not silently dropped.
if self._trigger_delta_has_content(processed_chunk):
self.chunk_queue.append(processed_chunk)
self.sent_content_block_finish = False
@ -711,12 +710,16 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
if not self.queued_usage_chunk:
if should_start_new_block and not self.sent_content_block_finish:
# Queue the sequence: content_block_stop -> content_block_start
# For text blocks the trigger chunk is not emitted as a separate
# delta because content_block_start carries the information.
# For tool_use blocks we must also emit the trigger chunk's delta
# when it carries input_json_delta data, because some providers
# (e.g. xAI, Gemini) include tool arguments in the same streaming
# chunk as the function name/id.
# -> (optionally) the trigger chunk's delta.
#
# The synthesized content_block_start always carries an
# empty body, so the chunk that *triggered* the transition
# also carries the new block's first delta. It must be
# re-emitted or the first token of the new block is lost.
# This applies to text_delta and thinking_delta (the
# first non-empty text/thinking token) as well as
# input_json_delta (providers like xAI/Gemini bundle tool
# arguments with the function name/id in a single chunk).
# 1. Stop current content block
self.chunk_queue.append(
@ -733,15 +736,9 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
}
)
# 3. If the trigger chunk carries tool argument data, queue it
# so the input_json_delta is not silently dropped.
if (
processed_chunk.get("type") == "content_block_delta"
and isinstance(processed_chunk.get("delta"), dict)
and processed_chunk["delta"].get("type")
== "input_json_delta"
and processed_chunk["delta"].get("partial_json")
):
# 3. If the trigger chunk carries delta content, queue it
# so the first delta of the new block is not silently dropped.
if self._trigger_delta_has_content(processed_chunk):
self.chunk_queue.append(processed_chunk)
# Reset state for new block
@ -898,6 +895,38 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
def _increment_content_block_index(self):
self.current_content_block_index += 1
@staticmethod
def _trigger_delta_has_content(processed_chunk: Dict[str, Any]) -> bool:
"""Return True if a translated trigger chunk carries a non-empty
``content_block_delta`` payload that must be re-emitted after a
block transition.
When an upstream chunk both *triggers* a new content block (its type
differs from the active block) and *carries* delta content, that
content belongs to the new block. The synthesized
``content_block_start`` only ever carries an empty body — see
``_translate_streaming_openai_chunk_to_anthropic_content_block``,
which returns an empty ``TextBlock``/``ToolUseBlock``/thinking block —
so the trigger chunk's delta must be re-queued or the first token of
the new block (the first non-empty text/thinking delta, or bundled
tool arguments) is silently dropped.
"""
if processed_chunk.get("type") != "content_block_delta":
return False
delta = processed_chunk.get("delta")
if not isinstance(delta, dict):
return False
delta_type = delta.get("type")
if delta_type == "text_delta":
return bool(delta.get("text"))
if delta_type == "input_json_delta":
return bool(delta.get("partial_json"))
if delta_type == "thinking_delta":
return bool(delta.get("thinking"))
if delta_type == "signature_delta":
return bool(delta.get("signature"))
return False
def _should_start_new_content_block(self, chunk: "ModelResponseStream") -> bool:
"""
Determine if we should start a new content block based on the processed chunk.

View file

@ -102,9 +102,9 @@ def _build_responses_kwargs(
from litellm.types.utils import CallTypes
if isinstance(value, LiteLLMLoggingObject):
# Reclassify as acompletion so the success handler doesn't try to
# validate the Responses API event as an AnthropicResponse.
# (Mirrors the pattern used in LiteLLMMessagesToCompletionTransformationHandler.)
# Keep call_type as anthropic_messages so spend_logs are billed
# against /v1/messages; the success handler translates the
# Responses API result back to a ModelResponse for the row.
setattr(value, "call_type", CallTypes.anthropic_messages.value)
responses_kwargs[key] = value
elif key not in excluded and key not in responses_kwargs and value is not None:

View file

@ -8,6 +8,7 @@ from typing import Any, Optional, cast
from litellm._logging import _redact_string, verbose_proxy_logger
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
from litellm.types.realtime import RealtimeQueryParams
from ....litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from ....litellm_core_utils.realtime_streaming import RealTimeStreaming
@ -35,6 +36,7 @@ class AzureOpenAIRealtime(AzureChatCompletion):
model: str,
api_version: Optional[str],
realtime_protocol: Optional[str] = None,
query_params: Optional[RealtimeQueryParams] = None,
) -> str:
"""
Construct Azure realtime WebSocket URL.
@ -46,6 +48,7 @@ class AzureOpenAIRealtime(AzureChatCompletion):
realtime_protocol: Protocol version to use:
- "GA" or "v1": Uses /openai/v1/realtime (GA path)
- "beta" or None: Uses /openai/realtime (beta path, default)
query_params: Extra query params to forward (e.g. intent=transcription).
Returns:
WebSocket URL string
@ -54,6 +57,8 @@ class AzureOpenAIRealtime(AzureChatCompletion):
beta/default: "wss://.../openai/realtime?api-version=2024-10-01-preview&deployment=gpt-4o-realtime-preview"
GA/v1: "wss://.../openai/v1/realtime?model=gpt-realtime-deployment"
"""
from urllib.parse import urlencode
api_base = api_base.replace("https://", "wss://")
# Determine path based on realtime_protocol (case-insensitive)
@ -61,13 +66,25 @@ class AzureOpenAIRealtime(AzureChatCompletion):
"GA",
"V1",
)
intent = (query_params or {}).get("intent")
if _is_ga:
path = "/openai/v1/realtime"
return f"{api_base}{path}?model={model}"
query_parts = []
if intent != "transcription" and (
query_params is None or "model" in query_params
):
query_parts.append(urlencode({"model": model}))
else:
# Default to beta path for backwards compatibility
path = "/openai/realtime"
return f"{api_base}{path}?api-version={api_version}&deployment={model}"
query_parts = [urlencode({"api-version": api_version, "deployment": model})]
if intent:
query_parts.append(urlencode({"intent": intent}))
qs = "&".join(query_parts)
return f"{api_base}{path}?{qs}" if qs else f"{api_base}{path}"
async def async_realtime(
self,
@ -81,6 +98,7 @@ class AzureOpenAIRealtime(AzureChatCompletion):
client: Optional[Any] = None,
timeout: Optional[float] = None,
realtime_protocol: Optional[str] = None,
query_params: Optional[RealtimeQueryParams] = None,
user_api_key_dict: Optional[Any] = None,
litellm_metadata: Optional[dict] = None,
):
@ -96,7 +114,11 @@ class AzureOpenAIRealtime(AzureChatCompletion):
raise ValueError("api_version is required for Azure OpenAI calls")
url = self._construct_url(
api_base, model, api_version, realtime_protocol=realtime_protocol
api_base,
model,
api_version,
realtime_protocol=realtime_protocol,
query_params=query_params,
)
try:
@ -113,9 +135,15 @@ class AzureOpenAIRealtime(AzureChatCompletion):
websocket,
cast(ClientConnection, backend_ws),
logging_obj,
model=model,
user_api_key_dict=user_api_key_dict,
request_data={"litellm_metadata": litellm_metadata or {}},
backend_uses_beta_protocol=backend_uses_beta_protocol,
force_transcription_model=(
model
if (query_params or {}).get("intent") == "transcription"
else None
),
)
await realtime_streaming.bidirectional_forward()

View file

@ -40,6 +40,13 @@ class AzureRealtimeHTTPConfig(BaseRealtimeHTTPConfig):
version = api_version or get_secret_str("AZURE_API_VERSION") or "2024-12-17"
return f"{base}/openai/realtime/calls?api-version={version}"
def get_transcription_session_url(
self, api_base: Optional[str], model: str, api_version: Optional[str] = None
) -> str:
base = self.get_api_base(api_base).rstrip("/")
version = api_version or get_secret_str("AZURE_API_VERSION") or "2024-12-17"
return f"{base}/openai/realtime/transcription_sessions?api-version={version}"
def get_realtime_calls_headers(self, ephemeral_key: str) -> dict:
return {
"api-key": ephemeral_key,

View file

@ -59,6 +59,15 @@ class BaseRealtimeHTTPConfig(ABC):
) -> str:
"""Return the full URL for POST /realtime/client_secrets."""
def get_transcription_session_url(
self, api_base: Optional[str], model: str, api_version: Optional[str] = None
) -> str:
"""Return the full URL for POST /realtime/transcription_sessions."""
base = (api_base or "").rstrip("/")
if base.endswith("/v1"):
base = base[:-3]
return f"{base}/v1/realtime/transcription_sessions"
@abstractmethod
def validate_environment(
self,

View file

@ -125,6 +125,7 @@ from litellm.types.vector_stores import (
VectorStoreSearchOptionalRequestParams,
VectorStoreSearchResponse,
)
from litellm.types.realtime import RealtimeQueryParams
from litellm.types.videos.main import VideoObject
from litellm.utils import (
CustomStreamWrapper,
@ -2305,6 +2306,7 @@ class BaseLLMHTTPHandler:
if extra_body:
data.update(extra_body)
stream = bool(stream or data.get("stream"))
# Preserve the OpenAI-style request context (not sent to the provider) for streaming
# hooks/metadata; the streaming iterator now consumes this to run deployment hooks
@ -2467,6 +2469,7 @@ class BaseLLMHTTPHandler:
if extra_body:
data.update(extra_body)
stream = bool(stream or data.get("stream"))
# Preserve the OpenAI-style request context (not sent to the provider) for streaming
# hooks/metadata; the streaming iterator now consumes this to run deployment hooks
@ -5315,6 +5318,23 @@ class BaseLLMHTTPHandler:
headers=error_headers,
)
@staticmethod
def _append_query_params(
url: str, query_params: Optional[RealtimeQueryParams]
) -> str:
"""Append query_params to url, skipping keys already present in the URL."""
if not query_params:
return url
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
parsed = urlparse(url)
existing = dict(parse_qsl(parsed.query))
extras = {k: v for k, v in query_params.items() if k not in existing}
if not extras:
return url
new_query = parsed.query + ("&" if parsed.query else "") + urlencode(extras)
return urlunparse(parsed._replace(query=new_query))
async def async_realtime(
self,
model: str,
@ -5328,11 +5348,14 @@ class BaseLLMHTTPHandler:
timeout: Optional[float] = None,
user_api_key_dict: Optional[Any] = None,
litellm_metadata: Optional[Dict[str, Any]] = None,
query_params: Optional[RealtimeQueryParams] = None,
):
import websockets
from websockets.asyncio.client import ClientConnection
url = provider_config.get_complete_url(api_base, model, api_key)
url = self._append_query_params(
provider_config.get_complete_url(api_base, model, api_key), query_params
)
headers = provider_config.validate_environment(
headers=headers,
model=model,
@ -5373,6 +5396,11 @@ class BaseLLMHTTPHandler:
model,
user_api_key_dict=user_api_key_dict,
request_data=_request_data,
force_transcription_model=(
model
if (query_params or {}).get("intent") == "transcription"
else None
),
)
if _session_config:
realtime_streaming.session_configuration_request = _session_config
@ -5437,6 +5465,69 @@ class BaseLLMHTTPHandler:
"""
Forward POST /v1/realtime/client_secrets to upstream provider.
Uses provider_config (BaseRealtimeHTTPConfig) for URL construction and
header auth when available; falls back to the legacy OpenAI-style defaults.
"""
return await self._async_realtime_session_post(
endpoint="client_secrets",
api_base=api_base,
api_key=api_key,
request_data=request_data,
logging_obj=logging_obj,
timeout=timeout,
provider_config=provider_config,
model=model,
extra_headers=extra_headers,
client=client,
api_version=api_version,
)
async def async_realtime_transcription_session_handler(
self,
api_base: str,
api_key: str,
request_data: Dict[str, Any],
logging_obj: LiteLLMLoggingObj,
timeout: Union[float, httpx.Timeout],
provider_config: Optional[Any] = None,
model: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
api_version: Optional[str] = None,
) -> httpx.Response:
"""Forward POST /v1/realtime/transcription_sessions to upstream provider."""
return await self._async_realtime_session_post(
endpoint="transcription_sessions",
api_base=api_base,
api_key=api_key,
request_data=request_data,
logging_obj=logging_obj,
timeout=timeout,
provider_config=provider_config,
model=model,
extra_headers=extra_headers,
client=client,
api_version=api_version,
)
async def _async_realtime_session_post(
self,
endpoint: Literal["client_secrets", "transcription_sessions"],
api_base: str,
api_key: str,
request_data: Dict[str, Any],
logging_obj: LiteLLMLoggingObj,
timeout: Union[float, httpx.Timeout],
provider_config: Optional[Any] = None,
model: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
api_version: Optional[str] = None,
) -> httpx.Response:
"""
Shared POST flow for the realtime HTTP session endpoints
(client_secrets and transcription_sessions).
Uses provider_config (BaseRealtimeHTTPConfig) for URL construction and
header auth when available; falls back to the legacy OpenAI-style defaults.
"""
@ -5448,14 +5539,19 @@ class BaseLLMHTTPHandler:
async_httpx_client = client
if provider_config is not None:
url = provider_config.get_complete_url(
api_base=api_base, model=model or "", api_version=api_version
)
if endpoint == "transcription_sessions":
url = provider_config.get_transcription_session_url(
api_base=api_base, model=model or "", api_version=api_version
)
else:
url = provider_config.get_complete_url(
api_base=api_base, model=model or "", api_version=api_version
)
headers: Dict[str, Any] = provider_config.validate_environment(
headers={}, model=model or "", api_key=api_key
)
else:
url = f"{api_base.rstrip('/')}/v1/realtime/client_secrets"
url = f"{api_base.rstrip('/')}/v1/realtime/{endpoint}"
headers = {
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",

View file

@ -101,6 +101,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig):
def __init__(self) -> None:
super().__init__()
self.authenticator = Authenticator()
self._stream_item_ids_by_output_index: Dict[int, str] = {}
@property
def custom_llm_provider(self) -> LlmProviders:
@ -129,6 +130,61 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig):
"""
return dict(response_api_optional_params)
def transform_streaming_response(
self,
model: str,
parsed_chunk: dict,
logging_obj: LiteLLMLoggingObj,
) -> Any:
parsed_chunk = self._normalize_stream_item_id(parsed_chunk)
return super().transform_streaming_response(
model=model,
parsed_chunk=parsed_chunk,
logging_obj=logging_obj,
)
def _normalize_stream_item_id(self, parsed_chunk: dict) -> dict:
"""Rewrite streamed item ids to one stable id per output_index.
GitHub Copilot tags each event of a single output item with a different
item id, so clients that key streaming state by item id (e.g. the Vercel
AI SDK) crash with "reasoning part <id> not found" / "text part <id> not
found". Every sub-event carries a top-level ``item_id`` (whatever the
item type), so its presence is the rewrite signal; output_item.added /
.done instead nest the id under ``item``. The anchor is keyed by
output_index and taken from output_item.added, which the protocol always
emits first, so it is written before any sub-event reads it. Copilot
accepts that id paired with the final encrypted_content next turn, so
multi-turn replay is unaffected.
State is keyed by output_index on this config, which
ProviderConfigManager builds fresh per request, so it is stream-scoped.
"""
output_index = parsed_chunk.get("output_index")
if not isinstance(output_index, int):
return parsed_chunk
if parsed_chunk.get("type") == "response.output_item.added":
item = parsed_chunk.get("item")
if isinstance(item, dict) and isinstance(item.get("id"), str):
self._stream_item_ids_by_output_index[output_index] = item["id"]
return parsed_chunk
stable_id = self._stream_item_ids_by_output_index.get(output_index)
if stable_id is None:
return parsed_chunk
if isinstance(parsed_chunk.get("item_id"), str):
parsed_chunk = dict(parsed_chunk)
parsed_chunk["item_id"] = stable_id
elif parsed_chunk.get("type") == "response.output_item.done":
item = parsed_chunk.get("item")
if isinstance(item, dict):
parsed_chunk = dict(parsed_chunk)
parsed_chunk["item"] = {**item, "id": stable_id}
return parsed_chunk
def validate_environment(
self,
headers: dict,

View file

@ -25,6 +25,7 @@ from typing import (
import httpx
import litellm
from litellm.constants import DEFAULT_OCI_CHAT_MAX_TOKENS
from litellm.litellm_core_utils.logging_utils import track_llm_api_timing
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.llms.custom_httpx.http_handler import (
@ -87,15 +88,20 @@ STREAMING_TIMEOUT = 60 * 5
def _model_uses_max_completion_tokens(model: str) -> bool:
"""Return True for OCI-hosted models that require ``maxCompletionTokens``.
Reasoning models on OCI (e.g. the OpenAI GPT-5 family) reject ``maxTokens``
with HTTP 400 and require ``maxCompletionTokens`` per OpenAI's reasoning-API
convention. Driven by ``supports_reasoning`` in
``model_prices_and_context_window.json`` so new model families are picked
up via a catalog update rather than a code change.
OpenAI commercial models proxied through OCI (``openai.*``) reject
``maxTokens`` with HTTP 400 on the reasoning families (gpt-5.x, o-series)
and accept ``maxCompletionTokens`` everywhere, so route the whole vendor
prefix to it rather than chasing each new release in
``model_prices_and_context_window.json``. The ``openai.gpt-oss-*`` open
weights are served by OCI's own stack and keep ``maxTokens``. Any other
vendor falls back to the catalog's ``supports_reasoning`` flag.
"""
if not model:
return False
name = model[4:] if model.lower().startswith("oci/") else model
lowered = name.lower()
if lowered.startswith("openai."):
return not lowered.startswith("openai.gpt-oss")
return supports_reasoning(model=name, custom_llm_provider="oci")
@ -193,19 +199,49 @@ def _normalize_response_format(selected_params: Dict, vendor: OCIVendors) -> Non
rf = selected_params.get("responseFormat")
if not isinstance(rf, dict) or "type" not in rf:
return
rf_payload = dict(rf)
selected_params["responseFormat"] = rf_payload
response_type = rf_payload["type"]
if "json_schema" in rf_payload:
raw_schema = rf_payload.pop("json_schema")
rf_payload["jsonSchema"] = (
dict(raw_schema) if isinstance(raw_schema, dict) else raw_schema
)
rf_type = str(rf["type"]).lower()
raw_schema = rf.get("json_schema")
json_schema = raw_schema if isinstance(raw_schema, dict) else None
if rf_type == "text":
selected_params["responseFormat"] = {"type": "TEXT"}
return
if vendor == OCIVendors.COHERE:
rf_payload["type"] = response_type
else:
fmt = response_type.upper()
rf_payload["type"] = "JSON_OBJECT" if fmt == "JSON" else fmt
# OCI Cohere has no JSON_SCHEMA type; a schema rides on JSON_OBJECT.
payload: Dict[str, Any] = {"type": "JSON_OBJECT"}
if json_schema is not None and json_schema.get("schema") is not None:
payload["schema"] = json_schema["schema"]
selected_params["responseFormat"] = payload
return
if rf_type == "json_schema":
if json_schema is None:
raise OCIError(
status_code=400,
message="response_format type 'json_schema' requires a 'json_schema' object",
)
# OCI's ResponseJsonSchema accepts only name/description/schema/isStrict.
# OpenAI sends `strict` instead of `isStrict`; forwarding it (or any
# other extra key) makes OCI reject the whole request with HTTP 400.
oci_schema: Dict[str, Any] = {"name": json_schema.get("name") or "response"}
if json_schema.get("description") is not None:
oci_schema["description"] = json_schema["description"]
if json_schema.get("schema") is not None:
oci_schema["schema"] = json_schema["schema"]
if json_schema.get("strict") is not None:
oci_schema["isStrict"] = json_schema["strict"]
selected_params["responseFormat"] = {
"type": "JSON_SCHEMA",
"jsonSchema": oci_schema,
}
return
fmt = rf_type.upper()
selected_params["responseFormat"] = {
"type": "JSON_OBJECT" if fmt == "JSON" else fmt
}
def get_vendor_from_model(model: str) -> OCIVendors:
@ -297,6 +333,11 @@ class OCIChatConfig(BaseConfig):
if get_vendor_from_model(model) == OCIVendors.COHERE
else self.openai_to_oci_generic_param_map
)
# `n` is intentionally not advertised for Cohere even though n=1 is
# tolerated: Cohere has no numGenerations field, so n>1 cannot be
# honoured and advertising it would be misleading. Callers that gate on
# this list strip n=1 (a no-op, matching what map_openai_params does);
# callers that bypass it have n=1 dropped there. Both paths converge.
return [key for key, value in param_map.items() if value]
def map_openai_params(
@ -317,6 +358,19 @@ class OCIChatConfig(BaseConfig):
for key, value in {**non_default_params, **optional_params}.items():
alias = param_map.get(key)
if alias is False:
# max_retries is a litellm-level control param (litellm applies
# retries itself); it is never a generation param OCI accepts, so
# drop it silently. The litellm proxy injects it on every request,
# which otherwise 500s OCI calls unless drop_params is set.
if key == "max_retries":
continue
# n=1 (or None) is the OpenAI default: a single generation, which
# every OCI model produces anyway. Drop it silently so standard
# clients that always send n=1 (e.g. the MLflow gateway) are not
# rejected; only n>1 is genuinely unsupported on Cohere, which
# has no numGenerations field.
if key == "n" and (value is None or value == 1):
continue
if drop_params or litellm.drop_params:
continue
raise OCIError(
@ -451,6 +505,13 @@ class OCIChatConfig(BaseConfig):
elif oci_alias in optional_params:
selected_params[target] = optional_params[oci_alias] # type: ignore[index]
# OCI's server-side default token cap is tiny (~20 tokens), so an
# omitted max_tokens silently truncates the response mid-string. Most
# callers never send a limit (MLflow judges among them), so inject a
# sane default when one is absent, mirroring litellm's Anthropic config.
if max_tokens_key not in selected_params:
selected_params[max_tokens_key] = DEFAULT_OCI_CHAT_MAX_TOKENS
# OCI expects uppercase reasoning levels (LOW/MEDIUM/HIGH/NONE); OpenAI
# clients send lowercase. OpenAI's "disable" maps to OCI's "NONE".
if "reasoningEffort" in selected_params:

View file

@ -157,8 +157,14 @@ class OpenAIRealtime(OpenAIChatCompletion):
websocket,
cast(ClientConnection, backend_ws),
logging_obj,
model=model,
user_api_key_dict=user_api_key_dict,
request_data={"litellm_metadata": litellm_metadata or {}},
force_transcription_model=(
model
if (query_params or {}).get("intent") == "transcription"
else None
),
)
await realtime_streaming.bidirectional_forward()

View file

@ -41,6 +41,14 @@ class OpenAIRealtimeHTTPConfig(BaseRealtimeHTTPConfig):
base = base[:-3]
return f"{base}/v1/realtime/calls"
def get_transcription_session_url(
self, api_base: Optional[str], model: str, api_version: Optional[str] = None
) -> str:
base = self.get_api_base(api_base).rstrip("/")
if base.endswith("/v1"):
base = base[:-3]
return f"{base}/v1/realtime/transcription_sessions"
def validate_environment(
self,
headers: dict,

View file

@ -1,7 +1,7 @@
"""
Calls Parallel AI's /search endpoint to search the web.
Calls Parallel AI's /v1/search endpoint to search the web.
Parallel AI API Reference: https://docs.parallel.ai/api-reference/search-and-extract-api-beta/search
Parallel AI API Reference: https://docs.parallel.ai/api-reference/search/search
"""
from typing import Dict, List, Optional, TypedDict, Union
@ -18,36 +18,43 @@ from litellm.secret_managers.main import get_secret_str
class _ParallelAISourcePolicy(TypedDict, total=False):
"""Source policy for Parallel AI search results."""
allowed_domains: List[str] # Optional - list of allowed domains
disallowed_domains: List[str] # Optional - list of disallowed domains
include_domains: List[str]
exclude_domains: List[str]
after_date: str
class _ParallelAISearchRequestRequired(TypedDict):
"""Required fields for Parallel AI Search API request."""
# Note: At least one of objective or search_queries must be provided
pass
class _ParallelAIExcerptSettings(TypedDict, total=False):
max_chars_per_result: int
class ParallelAISearchRequest(_ParallelAISearchRequestRequired, total=False):
class _ParallelAIAdvancedSettings(TypedDict, total=False):
source_policy: _ParallelAISourcePolicy
excerpt_settings: _ParallelAIExcerptSettings
fetch_policy: Dict
location: str
max_results: int
class ParallelAISearchRequest(TypedDict, total=False):
"""
Parallel AI Search API request format.
Based on: https://docs.parallel.ai/api-reference/search-and-extract-api-beta/search
Parallel AI v1 Search API request format.
Based on: https://docs.parallel.ai/api-reference/search/search
"""
search_queries: List[str] # Required - at least one keyword search query
objective: str # Optional - natural-language description of search goal
search_queries: List[str] # Optional - list of keyword search queries
processor: str # Optional - search processor ('base', 'pro'), default 'base'
max_results: int # Optional - maximum number of results, default 10
max_chars_per_result: int # Optional - max characters per result excerpt
source_policy: _ParallelAISourcePolicy # Optional - source policy for allowed/disallowed domains
mode: str # Optional - 'turbo', 'basic', or 'advanced' (default 'advanced')
max_chars_total: int # Optional - upper bound on total excerpt characters
session_id: str # Optional - tracks calls across search/extract requests
client_model: str # Optional - model consuming the results
advanced_settings: _ParallelAIAdvancedSettings
LEGACY_PROCESSOR_TO_MODE = {"base": "basic", "pro": "advanced"}
class ParallelAISearchConfig(BaseSearchConfig):
PARALLEL_AI_API_BASE = "https://api.parallel.ai"
PARALLEL_HEADER_SEARCH_EXTRACT_VALUE = "search-extract-2025-10-10"
@staticmethod
def ui_friendly_name() -> str:
@ -60,9 +67,6 @@ class ParallelAISearchConfig(BaseSearchConfig):
api_base: Optional[str] = None,
**kwargs,
) -> Dict:
"""
Validate environment and return headers.
"""
api_key = (
api_key
or get_secret_str("PARALLEL_AI_API_KEY")
@ -74,7 +78,6 @@ class ParallelAISearchConfig(BaseSearchConfig):
)
headers["x-api-key"] = api_key
headers["Content-Type"] = "application/json"
headers["parallel-beta"] = self.PARALLEL_HEADER_SEARCH_EXTRACT_VALUE
return headers
def get_complete_url(
@ -84,32 +87,18 @@ class ParallelAISearchConfig(BaseSearchConfig):
data: Optional[Union[Dict, List[Dict]]] = None,
**kwargs,
) -> str:
"""
Get complete URL for Search endpoint.
"""
api_base = (
api_base
or get_secret_str("PARALLEL_AI_API_BASE")
or self.PARALLEL_AI_API_BASE
)
# Parallel AI search endpoint is at /v1beta/search
if not api_base.endswith("/v1beta/search"):
if api_base.endswith("/"):
api_base = f"{api_base}v1beta/search"
else:
api_base = f"{api_base}/v1beta/search"
api_base = api_base.rstrip("/")
if not api_base.endswith("/v1/search"):
api_base = f"{api_base.removesuffix('/v1')}/v1/search"
return api_base
def _transform_query_to_objective(self, query: Union[str, List[str]]) -> str:
"""
Transform query to objective.
"""
if isinstance(query, list):
return " ".join(query)
return query
def transform_search_request(
self,
query: Union[str, List[str]],
@ -117,57 +106,78 @@ class ParallelAISearchConfig(BaseSearchConfig):
**kwargs,
) -> Dict:
"""
Transform Search request to Parallel AI API format.
Transform Search request to Parallel AI v1 API format.
Args:
query: Search query (string or list of strings)
- If string: maps to `objective` (natural language)
- If string: maps to `search_queries` (single item) and `objective`
- If list: maps to `search_queries` (keyword queries)
optional_params: Optional parameters for the request
- max_results: Maximum number of search results (default 10)
- search_domain_filter: List of domains to include -> maps to `source_policy.allowed_domains`
- exclude_domains: List of domains to exclude -> maps to `source_policy.disallowed_domains`
- processor: Search processor ('base', 'pro')
- max_chars_per_result: Max characters per result excerpt
- mode: Search mode ('turbo', 'basic', 'advanced'); defaults to 'basic'
- processor: Legacy v1beta param; 'base' maps to mode 'basic', 'pro' to 'advanced'
- max_results: Maximum number of search results -> `advanced_settings.max_results`
- search_domain_filter: Domains to include -> `advanced_settings.source_policy.include_domains`
- exclude_domains: Domains to exclude -> `advanced_settings.source_policy.exclude_domains`
- country: ISO 3166-1 alpha-2 code -> `advanced_settings.location`
- max_chars_per_result: -> `advanced_settings.excerpt_settings.max_chars_per_result`
- Any other params are passed through to the request body as-is
Returns:
Dict with typed request data following ParallelAISearchRequest spec
Dict with request data following the v1 search request spec
"""
params = dict(optional_params)
request_data: ParallelAISearchRequest = {}
# Map query to objective (string or list both become objective)
if isinstance(query, list):
request_data["objective"] = self._transform_query_to_objective(query)
request_data["search_queries"] = query
else:
request_data["search_queries"] = [query]
request_data["objective"] = query
# Transform Perplexity unified spec parameters to Parallel AI format
if "max_results" in optional_params:
request_data["max_results"] = optional_params["max_results"]
mode = params.pop("mode", None)
processor = params.pop("processor", None)
if mode is None and processor is not None:
mode = LEGACY_PROCESSOR_TO_MODE.get(processor, processor)
# the v1 API defaults to 'advanced' when mode is omitted; default to 'basic'
# instead to keep v1beta's default tier (processor 'base') and litellm's
# $0.004/query cost map entry for `parallel_ai/search` accurate
request_data["mode"] = mode or "basic"
advanced_settings: _ParallelAIAdvancedSettings = {}
if "max_results" in params:
advanced_settings["max_results"] = params.pop("max_results")
if "country" in params:
advanced_settings["location"] = params.pop("country")
if "max_chars_per_result" in params:
advanced_settings["excerpt_settings"] = {
"max_chars_per_result": params.pop("max_chars_per_result")
}
# Map domain filters to source_policy
source_policy: _ParallelAISourcePolicy = {}
if "search_domain_filter" in optional_params:
source_policy["allowed_domains"] = optional_params["search_domain_filter"]
if "search_domain_filter" in params:
source_policy["include_domains"] = params.pop("search_domain_filter")
if "exclude_domains" in optional_params:
source_policy["disallowed_domains"] = optional_params["exclude_domains"]
if "exclude_domains" in params:
source_policy["exclude_domains"] = params.pop("exclude_domains")
if source_policy:
request_data["source_policy"] = source_policy
advanced_settings["source_policy"] = source_policy
# Convert to dict before dynamic key assignments
result_data = dict(request_data)
advanced_settings.update(params.pop("advanced_settings", {}))
# pass through all other parameters as-is
for param, value in optional_params.items():
if (
param not in self.get_supported_perplexity_optional_params()
and param not in result_data
):
result_data[param] = value
if advanced_settings:
request_data["advanced_settings"] = advanced_settings
# unified-spec param with no v1 equivalent
params.pop("max_tokens_per_page", None)
result_data: Dict = dict(request_data)
result_data.update(params)
return result_data
def transform_search_response(
@ -177,36 +187,27 @@ class ParallelAISearchConfig(BaseSearchConfig):
**kwargs,
) -> SearchResponse:
"""
Transform Parallel AI API response to LiteLLM unified SearchResponse format.
Transform Parallel AI v1 API response to LiteLLM unified SearchResponse format.
Parallel AI → LiteLLM mappings:
- results[].title → SearchResult.title
- results[].url → SearchResult.url
- results[].excerpts (array) → SearchResult.snippet (joined string)
- No date/last_updated fields in Parallel AI response (set to None)
Args:
raw_response: Raw httpx response from Parallel AI API
logging_obj: Logging object for tracking
Returns:
SearchResponse with standardized format
Parallel AI -> LiteLLM mappings:
- results[].title -> SearchResult.title
- results[].url -> SearchResult.url
- results[].excerpts (array) -> SearchResult.snippet (joined string)
- results[].publish_date -> SearchResult.date
"""
response_json = raw_response.json()
# Transform results to SearchResult objects
results = []
for result in response_json.get("results", []):
# Join excerpts array into a single snippet string
excerpts = result.get("excerpts", [])
excerpts = result.get("excerpts") or []
snippet = " ... ".join(excerpts) if excerpts else ""
search_result = SearchResult(
title=result.get("title", ""),
url=result.get("url", ""),
title=result.get("title") or "",
url=result.get("url") or "",
snippet=snippet,
date=None, # Parallel AI doesn't provide date in response
last_updated=None, # Parallel AI doesn't provide last_updated in response
date=result.get("publish_date"),
last_updated=None,
)
results.append(search_result)

View file

@ -3422,14 +3422,18 @@ class ModelResponseIterator:
self.has_seen_tool_calls = True
break
# Handle final chunk with finishReason but no content.
# _process_candidates skips candidates without "content",
# so the finish_reason from the final chunk is lost.
# _process_candidates skips candidates without a "content" part, so a
# content-less chunk leaves choices empty and the downstream streaming
# handler hits IndexError on choices[0]. This covers the final chunk
# (finishReason, no content) and mid-stream metadata-only chunks
# (grounding/web-search/thought, no content and no finishReason — seen
# with web_search + reasoning) by emitting an empty-delta choice.
if not model_response.choices and _candidates:
from litellm.types.utils import Delta, StreamingChoices
for candidate in _candidates:
finish_reason_str = candidate.get("finishReason")
mapped_finish_reason = None
if finish_reason_str is not None:
if self.has_seen_tool_calls:
mapped_finish_reason = "tool_calls"
@ -3437,14 +3441,14 @@ class ModelResponseIterator:
mapped_finish_reason = VertexGeminiConfig._check_finish_reason(
None, finish_reason_str
)
choice = StreamingChoices(
finish_reason=mapped_finish_reason,
index=candidate.get("index", 0),
delta=Delta(content=None, role=None),
logprobs=None,
enhancements=None,
)
model_response.choices.append(choice)
choice = StreamingChoices(
finish_reason=mapped_finish_reason,
index=candidate.get("index", 0),
delta=Delta(content=None, role=None),
logprobs=None,
enhancements=None,
)
model_response.choices.append(choice)
# Also handle the case where the final chunk has empty
# content (e.g. text:"") WITH finishReason. In this case

View file

@ -4409,6 +4409,23 @@
"/v1/audio/transcriptions"
]
},
"azure/gpt-realtime-whisper": {
"input_cost_per_second": 0.0002833333333333333,
"litellm_provider": "azure",
"mode": "audio_transcription",
"source": "https://learn.microsoft.com/en-us/azure/ai-foundry/openai/concepts/gpt-realtime-whisper",
"supported_endpoints": [
"/v1/realtime",
"/v1/realtime/transcription_sessions"
],
"supported_modalities": [
"audio"
],
"supported_output_modalities": [
"text"
],
"supports_audio_input": true
},
"azure/gpt-5.1-2025-11-13": {
"cache_read_input_token_cost": 1.25e-07,
"cache_read_input_token_cost_priority": 2.5e-07,
@ -7557,6 +7574,45 @@
"supports_function_calling": true,
"supports_tool_choice": true
},
"azure_ai/deepseek-v3.1": {
"input_cost_per_token": 1.23e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.94e-06,
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/deepseek/",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"azure_ai/deepseek-v4-pro": {
"input_cost_per_token": 1.74e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 1000000,
"max_output_tokens": 384000,
"max_tokens": 384000,
"mode": "chat",
"output_cost_per_token": 3.48e-06,
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/deepseek/",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"azure_ai/deepseek-v4-flash": {
"input_cost_per_token": 1.9e-07,
"litellm_provider": "azure_ai",
"max_input_tokens": 1000000,
"max_output_tokens": 384000,
"max_tokens": 384000,
"mode": "chat",
"output_cost_per_token": 5.1e-07,
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/deepseek/",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"azure_ai/embed-v-4-0": {
"input_cost_per_token": 1.2e-07,
"litellm_provider": "azure_ai",
@ -40916,6 +40972,23 @@
"supports_system_messages": true,
"supports_tool_choice": true
},
"gpt-realtime-whisper": {
"input_cost_per_second": 0.0002833333333333333,
"litellm_provider": "openai",
"mode": "audio_transcription",
"source": "https://platform.openai.com/docs/models/gpt-realtime-whisper",
"supported_endpoints": [
"/v1/realtime",
"/v1/realtime/transcription_sessions"
],
"supported_modalities": [
"audio"
],
"supported_output_modalities": [
"text"
],
"supports_audio_input": true
},
"sora-2": {
"litellm_provider": "openai",
"mode": "video_generation",

View file

@ -10,8 +10,9 @@ class MCPUpstreamAuthError(Exception):
(typically HTTP 401) and the gateway should surface it transparently to
the client instead of swallowing it.
Only relevant for pass-through MCP servers (see
``MCPServer.is_oauth_passthrough``). The gateway converts this exception
Relevant for MCP servers that delegate OAuth to the upstream server,
including pass-through servers and OAuth2 servers with
``delegate_auth_to_upstream`` enabled. The gateway converts this exception
into an HTTP 401 response on single-server routes, preserving any
``WWW-Authenticate`` challenge emitted by the upstream so standards-
compliant MCP clients can trigger the upstream OAuth flow.

View file

@ -2777,28 +2777,40 @@ class MCPServerManager:
Uses anyio.fail_after() instead of asyncio.wait_for() to avoid conflicts
with the MCP SDK's anyio TaskGroup. See GitHub issue #20715 for details.
For pass-through MCP servers (``MCPServer.is_oauth_passthrough``) an
For OAuth pass-through and upstream-delegated OAuth2 MCP servers, an
upstream HTTP 401 is converted into :class:`MCPUpstreamAuthError`
instead of being swallowed to an empty tool list. That lets the
single-server HTTP routes surface a proper 401 + ``WWW-Authenticate``
challenge so standards-compliant MCP clients trigger the upstream
OAuth flow. Non-pass-through servers keep today's swallow-and-log
behaviour so the multi-server ``/mcp`` aggregator doesn't get
tainted by a single bad server.
OAuth flow. Other servers keep today's swallow-and-log behaviour so
the multi-server ``/mcp`` aggregator doesn't get tainted by a single
bad server.
Args:
client: MCP client instance
server_name: Name of the server for logging
server: Optional MCPServer; when pass-through, auth errors are
re-raised as :class:`MCPUpstreamAuthError`.
server: Optional MCPServer; when upstream auth is delegated, auth
errors are re-raised as :class:`MCPUpstreamAuthError`.
Returns:
List of tools from the server
"""
is_passthrough = bool(server is not None and server.is_oauth_passthrough)
should_surface_upstream_auth = bool(
server is not None
and (
server.is_oauth_passthrough
or (
server.auth_type == MCPAuth.oauth2
and getattr(server, "delegate_auth_to_upstream", False) is True
and not server.has_client_credentials
)
)
)
try:
with anyio.fail_after(MCP_TOOL_LISTING_TIMEOUT):
tools = await client.list_tools(raise_on_error=is_passthrough)
tools = await client.list_tools(
raise_on_error=should_surface_upstream_auth
)
verbose_logger.debug(f"Tools from {server_name}: {tools}")
return tools
except TimeoutError:
@ -2815,12 +2827,12 @@ class MCPServerManager:
)
return []
except Exception as e:
if is_passthrough:
if should_surface_upstream_auth:
auth_info = _extract_upstream_auth_failure(e)
if auth_info is not None:
status_code, www_authenticate = auth_info
verbose_logger.info(
f"Upstream auth failure from pass-through MCP server "
f"Upstream auth failure from MCP server "
f"{server_name}: HTTP {status_code}"
)
raise MCPUpstreamAuthError(

View file

@ -2552,6 +2552,7 @@ if MCP_AVAILABLE:
standard_logging_mcp_tool_call
)
litellm_logging_obj.model = f"MCP: {name}"
litellm_logging_obj.model_call_details["model"] = f"MCP: {name}"
# Resolve the MCP server early so BYOK checks and credential injection
# apply to ALL dispatch paths (local tool registry AND managed MCP server).
if mcp_server is None:
@ -3426,6 +3427,8 @@ if MCP_AVAILABLE:
)
if stored_oauth_headers:
continue
if getattr(server, "delegate_auth_to_upstream", False) is True:
continue
request = StarletteRequest(scope)
base_url = get_request_base_url(request)
@ -3960,7 +3963,7 @@ if MCP_AVAILABLE:
):
_stateful_session_locks.pop(active_request_session_id, None)
except MCPUpstreamAuthError as e:
# Pass-through server returned 401 — surface it to the client so
# Upstream delegated auth returned 401; surface it to the client so
# standards-compliant MCP clients trigger the upstream OAuth flow.
raise e.to_http_exception(
base_url=get_request_base_url(StarletteRequest(scope)),
@ -4076,7 +4079,7 @@ if MCP_AVAILABLE:
):
await sse_session_manager.handle_request(scope, receive, send)
except MCPUpstreamAuthError as e:
# Pass-through server returned 401 — surface it to the client so
# Upstream delegated auth returned 401; surface it to the client so
# standards-compliant MCP clients trigger the upstream OAuth flow.
raise e.to_http_exception(
base_url=get_request_base_url(StarletteRequest(scope)),

Binary file not shown.

After

Width:  |  Height:  |  Size: 862 B

View file

@ -1856,6 +1856,10 @@ class BlockKeyRequest(LiteLLMPydanticObjectBase):
key: str # required
class BlockModelRequest(LiteLLMPydanticObjectBase):
model_id: str # required
class AddTeamCallback(LiteLLMPydanticObjectBase):
callback_name: str
callback_type: Optional[Literal["success", "failure", "success_and_failure"]] = (
@ -2231,8 +2235,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
health_check_concurrency: Optional[int] = Field(
None,
description=(
"limit concurrent health checks per cycle; when unset, "
"health checks run without a concurrency cap"
"limit concurrent health checks per cycle; when unset, health checks run without a concurrency cap"
),
)
health_check_skip_disabled_background_models: bool = Field(
@ -3094,6 +3097,14 @@ class AllCallbacks(LiteLLMPydanticObjectBase):
ui_callback_name="Galileo",
)
newrelic: CallbackOnUI = CallbackOnUI(
litellm_callback_name="newrelic",
ui_callback_name="New Relic",
litellm_callback_params=[
"NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED",
],
)
class SpendLogsMetadata(TypedDict):
"""

View file

@ -3155,6 +3155,98 @@ async def can_key_call_model(
raise
async def can_key_call_resolved_model(
model: str,
llm_model_list: Optional[list],
valid_token: UserAPIKeyAuth,
llm_router: Optional[litellm.Router],
) -> None:
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
skip_key_model_check = valid_token.config or (
isinstance(valid_token.models, list)
and SpecialModelNames.all_team_models.value in valid_token.models
)
if not skip_key_model_check:
await can_key_call_model(
model=model,
llm_model_list=llm_model_list,
valid_token=valid_token,
llm_router=llm_router,
)
team_object: Optional[LiteLLM_TeamTableCachedObj] = None
team_object_from_lookup = False
if valid_token.team_id is not None:
try:
team_object = await get_team_object(
team_id=valid_token.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=valid_token.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
team_object_from_lookup = True
except Exception:
team_object = LiteLLM_TeamTableCachedObj(
team_id=valid_token.team_id,
models=valid_token.team_models,
blocked=valid_token.team_blocked,
team_alias=valid_token.team_alias,
metadata=valid_token.team_metadata,
object_permission_id=valid_token.team_object_permission_id,
object_permission=valid_token.team_object_permission,
)
if team_object is not None:
try:
await can_team_access_model(
model=model,
team_object=team_object,
llm_router=llm_router,
team_model_aliases=valid_token.team_model_aliases,
)
except ProxyException as team_denial:
if team_denial.type != ProxyErrorTypes.team_model_access_denied:
raise
if not await _key_access_group_grants_model(
model=model,
valid_token=valid_token,
team_object=team_object,
llm_router=llm_router,
):
raise
if valid_token.user_id is not None and team_object_from_lookup:
await _check_team_member_model_access(
model=model,
team_object=team_object,
valid_token=valid_token,
llm_router=llm_router,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
if valid_token.project_id is not None:
project_object = await get_project_object(
project_id=valid_token.project_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
if project_object is not None and len(project_object.models) > 0:
can_project_access_model(
model=model,
project_object=project_object,
llm_router=llm_router,
)
def can_org_access_model(
model: str,
org_object: Optional[LiteLLM_OrganizationTable],

View file

@ -2075,6 +2075,17 @@ class ProxyBaseLLMRequestProcessing:
)
)
yield serialize_chunk(chunk)
except (asyncio.CancelledError, GeneratorExit):
# Client disconnected mid-stream. CancelledError / GeneratorExit
# are BaseException and bypass the success/failure logging
# callbacks that release the pre-call max_parallel_requests +1;
# release it here. This is the outermost generator Starlette closes
# on disconnect, so the nested iterator hook (which only sees
# GeneratorExit on GC) cannot own the refund.
proxy_logging_obj._release_max_parallel_requests_on_disconnect(
user_api_key_dict
)
raise
except Exception as e:
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.async_data_generator(): Exception occured - {}".format(

View file

@ -49,6 +49,7 @@ from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessag
from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
BedrockContentItem,
BedrockGuardrailOutput,
BedrockGuardrailQualifier,
BedrockGuardrailResponse,
BedrockRequest,
BedrockTextContent,
@ -74,6 +75,29 @@ from litellm.types.utils import (
GUARDRAIL_NAME = "bedrock"
_BEDROCK_DYNAMIC_BODY_DENYLIST = frozenset({"content", "source"})
# Maps an OpenAI message content-block ``type`` to the Bedrock guardrail qualifier
# it represents, so callers can drive contextual grounding by tagging their content.
# The model response is qualified as ``guard_content`` directly by the OUTPUT builder;
# the existing ``guarded_text`` marker is intentionally left unmapped here so its
# guardrail-hook payload is unchanged by this feature.
_CONTENT_TYPE_TO_QUALIFIER: Dict[str, BedrockGuardrailQualifier] = {
"grounding_source": "grounding_source",
"query": "query",
}
# Roles whose ``grounding_source`` blocks are trusted as reference material for the
# contextual-grounding check. Only app-authored roles qualify: ``tool``/``function``
# results and ``user`` content can carry caller- or externally-influenced text, which
# must not be graded against as if it were the application's own source material.
_GROUNDING_SOURCE_TRUSTED_ROLES = frozenset({"system", "developer"})
class QualifiedTextBlock(NamedTuple):
"""A piece of message text paired with its Bedrock grounding qualifier (if any)."""
text: str
qualifier: Optional[BedrockGuardrailQualifier]
class GuardrailMessageFilterResult(NamedTuple):
payload_messages: Optional[List[AllMessageValues]]
@ -164,41 +188,71 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
if messages is None:
return bedrock_request
for message in messages:
message_text_content: Optional[List[str]] = self.get_content_for_message(
message=message
)
if message_text_content is None:
blocks = self.get_content_items_for_message(message=message)
if blocks is None:
continue
for text_content in message_text_content:
bedrock_content_item = BedrockContentItem(
text=BedrockTextContent(text=text_content)
for block in blocks:
# INPUT scans send plain text only. Grounding qualifiers are attached
# exclusively when assembling the OUTPUT request, so a caller cannot use
# a grounding_source/query tag to change how input-safety policies treat
# their content (which would be an input-guardrail bypass).
bedrock_request_content.append(
BedrockContentItem(text=BedrockTextContent(text=block.text))
)
bedrock_request_content.append(bedrock_content_item)
bedrock_request["content"] = bedrock_request_content
return bedrock_request
def _create_bedrock_output_content_request(
self, response: Union[Any, ModelResponse]
self,
response: Union[Any, ModelResponse],
messages: Optional[List[AllMessageValues]] = None,
) -> BedrockRequest:
"""
Create a bedrock request for the output content - the LLM response.
Contextual grounding grades the response against the reference source and
the user query from the request. When the request tagged any
``grounding_source``/``query`` blocks, they are emitted first and the
response is qualified as ``guard_content`` so Bedrock can score grounding.
Without such tags the payload is the legacy single response block.
"""
bedrock_request: BedrockRequest = BedrockRequest(source="OUTPUT")
bedrock_request_content: List[BedrockContentItem] = []
if isinstance(response, litellm.ModelResponse):
for choice in response.choices:
if isinstance(choice, litellm.Choices):
if choice.message.content and isinstance(
choice.message.content, str
):
bedrock_content_item = BedrockContentItem(
text=BedrockTextContent(text=choice.message.content)
)
bedrock_request_content.append(bedrock_content_item)
bedrock_request["content"] = bedrock_request_content
grounding_blocks = self._collect_grounding_blocks(messages)
bedrock_request_content: List[BedrockContentItem] = [
self._build_content_item(block) for block in grounding_blocks
]
has_grounding = len(bedrock_request_content) > 0
# Append the response (the content to guard) after any grounding blocks; assign
# unconditionally so harvested grounding blocks survive a non-ModelResponse input.
bedrock_request_content.extend(
self._build_response_content_items(response, has_grounding=has_grounding)
)
bedrock_request["content"] = bedrock_request_content
return bedrock_request
def _build_response_content_items(
self, response: Union[Any, ModelResponse], has_grounding: bool
) -> List[BedrockContentItem]:
"""Build content item(s) from the model response. When the request supplied
grounding, the response is qualified ``guard_content`` so Bedrock can score it.
"""
items: List[BedrockContentItem] = []
if not isinstance(response, litellm.ModelResponse):
return items
for choice in response.choices:
if (
isinstance(choice, litellm.Choices)
and isinstance(choice.message.content, str)
and choice.message.content
):
block = QualifiedTextBlock(
text=choice.message.content,
qualifier="guard_content" if has_grounding else None,
)
items.append(self._build_content_item(block))
return items
def convert_to_bedrock_format(
self,
source: Literal["INPUT", "OUTPUT"],
@ -221,10 +275,68 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
)
elif source == "OUTPUT":
bedrock_request = self._create_bedrock_output_content_request(
response=response
response=response, messages=messages
)
return bedrock_request
def get_content_items_for_message(
self, message: AllMessageValues
) -> Optional[List[QualifiedTextBlock]]:
"""
Flatten a message into text blocks, preserving any contextual-grounding
qualifier carried by the content-block ``type`` (grounding_source / query).
Untagged text keeps ``qualifier=None`` so the payload is unchanged for
callers that do not use grounding.
"""
content = message.get("content")
if content is None:
return None
blocks: List[QualifiedTextBlock] = []
if isinstance(content, str):
blocks.append(QualifiedTextBlock(text=content, qualifier=None))
elif isinstance(content, list):
for item in content:
if isinstance(item, dict) and "text" in item:
qualifier = _CONTENT_TYPE_TO_QUALIFIER.get(item.get("type", ""))
blocks.append(
QualifiedTextBlock(text=item["text"], qualifier=qualifier)
)
elif isinstance(item, str):
blocks.append(QualifiedTextBlock(text=item, qualifier=None))
return blocks
def _build_content_item(self, block: QualifiedTextBlock) -> BedrockContentItem:
"""Build a Bedrock content item, attaching qualifiers only when present."""
text_content = BedrockTextContent(text=block.text)
if block.qualifier is not None:
text_content["qualifiers"] = [block.qualifier]
return BedrockContentItem(text=text_content)
def _collect_grounding_blocks(
self, messages: Optional[List[AllMessageValues]]
) -> List[QualifiedTextBlock]:
"""Harvest grounding_source/query blocks from the request for an OUTPUT scan.
``grounding_source`` is honored only from app-authored roles (system /
developer). A grounding_source tag on a ``user``, ``tool`` or ``function``
message is ignored, so neither a forwarded end-user message nor a tool/function
result carrying externally-influenced content can supply fake evidence for the
contextual-grounding check to grade the response against. ``query`` is accepted
from any role (it is the user's question).
"""
grounding: List[QualifiedTextBlock] = []
for message in messages or []:
role = message.get("role")
for block in self.get_content_items_for_message(message=message) or []:
if block.qualifier == "query":
grounding.append(block)
elif (
block.qualifier == "grounding_source"
and role in _GROUNDING_SOURCE_TRUSTED_ROLES
):
grounding.append(block)
return grounding
def _prepare_guardrail_messages_for_role(
self,
messages: Optional[List[AllMessageValues]],
@ -1169,6 +1281,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
output_content_bedrock = await self.make_bedrock_api_request(
source="OUTPUT",
response=response,
messages=new_messages,
request_data=data,
logging_event_type=GuardrailEventHooks.post_call,
)
@ -1281,6 +1394,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
output_guardrail_response = await self.make_bedrock_api_request(
source="OUTPUT",
response=assembled_model_response,
messages=request_data.get("messages"),
request_data=request_data,
logging_event_type=GuardrailEventHooks.post_call,
)
@ -1414,28 +1528,6 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
return new_content, masking_index
def get_content_for_message(self, message: AllMessageValues) -> Optional[List[str]]:
"""
Get the content for a message.
For bedrock guardrails we create a list of all the text content in the message.
If a message has a list of content items, we flatten the list and return a list of text content.
"""
message_text_content = []
content = message.get("content")
if content is None:
return None
if isinstance(content, str):
message_text_content.append(content)
elif isinstance(content, list):
for item in content:
if isinstance(item, dict) and "text" in item:
message_text_content.append(item["text"])
elif isinstance(item, str):
message_text_content.append(item)
return message_text_content
def _apply_masking_to_response(
self,
response: Union[ModelResponse, Any],

View file

@ -0,0 +1,46 @@
"""Ovalix guardrail hook: registration and initialization for the proxy."""
from typing import TYPE_CHECKING
from litellm.types.guardrails import SupportedGuardrailIntegrations
from .ovalix import OvalixGuardrail
if TYPE_CHECKING:
from litellm.types.guardrails import Guardrail, LitellmParams
def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"):
"""Create and register an Ovalix guardrail callback from proxy config."""
import litellm
tracker_api_base = getattr(litellm_params, "tracker_api_base", None)
tracker_api_key = getattr(litellm_params, "tracker_api_key", None)
application_id = getattr(litellm_params, "application_id", None)
pre_checkpoint_id = getattr(litellm_params, "pre_checkpoint_id", None)
post_checkpoint_id = getattr(litellm_params, "post_checkpoint_id", None)
_ovalix_callback = OvalixGuardrail(
guardrail_name=guardrail.get("guardrail_name", ""),
tracker_api_base=tracker_api_base,
tracker_api_key=tracker_api_key,
application_id=application_id,
pre_checkpoint_id=pre_checkpoint_id,
post_checkpoint_id=post_checkpoint_id,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
)
litellm.logging_callback_manager.add_litellm_callback(_ovalix_callback)
return _ovalix_callback
# Registry of guardrail name -> initializer for proxy config loading.
guardrail_initializer_registry = {
SupportedGuardrailIntegrations.OVALIX.value: initialize_guardrail,
}
# Registry of guardrail name -> guardrail class (e.g. for apply_guardrail API).
guardrail_class_registry = {
SupportedGuardrailIntegrations.OVALIX.value: OvalixGuardrail,
}

View file

@ -0,0 +1,330 @@
"""Ovalix guardrail integration: pre- and post-call checks via the Tracker service.
Use Ovalix Guardrails for your LLM calls. Supports pre_call (user input) and
post_call (model output) checkpoints with optional correction/blocking.
"""
import datetime
import hashlib
import os
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Type
import httpx
from litellm._logging import verbose_proxy_logger
from litellm.exceptions import GuardrailRaisedException
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
BLOCKED_BY_OVALIX_FALLBACK_MESSAGE = "This message was blocked by Ovalix"
BLOCKED_ACTION_TYPE = "block"
class OvalixGuardrailMissingSecrets(Exception):
"""Raised when required Ovalix config (API base, key, application/checkpoint IDs) is missing."""
pass
class OvalixGuardrailBlockedException(GuardrailRaisedException):
"""
Raised when Ovalix blocks a message. Sets status_code=400 so the proxy
returns 400 and HTTP clients do not retry (they retry on 5xx).
"""
status_code = 400
def __init__(
self,
guardrail_name: Optional[str] = None,
message: str = "",
should_wrap_with_default_message: bool = True,
):
super().__init__(
guardrail_name=guardrail_name,
message=message,
should_wrap_with_default_message=should_wrap_with_default_message,
)
class OvalixGuardrail(CustomGuardrail):
"""
Ovalix guardrail: pre-prompt (pre_call) and post-prompt (post_call) checks
via the Tracker service, with application and checkpoint resolution from the
Monolith backend.
"""
def __init__(
self,
tracker_api_base: Optional[str] = None,
tracker_api_key: Optional[str] = None,
application_id: Optional[str] = None,
pre_checkpoint_id: Optional[str] = None,
post_checkpoint_id: Optional[str] = None,
**kwargs: Any,
):
self._tracker_api_base = tracker_api_base or os.environ.get(
"OVALIX_TRACKER_API_BASE"
)
self._tracker_api_key = tracker_api_key or os.environ.get(
"OVALIX_TRACKER_API_KEY"
)
self._application_id = application_id or os.environ.get("OVALIX_APPLICATION_ID")
self._pre_checkpoint_id = pre_checkpoint_id or os.environ.get(
"OVALIX_PRE_CHECKPOINT_ID"
)
self._post_checkpoint_id = post_checkpoint_id or os.environ.get(
"OVALIX_POST_CHECKPOINT_ID"
)
if "supported_event_hooks" not in kwargs:
kwargs["supported_event_hooks"] = []
self._validate_config(kwargs["supported_event_hooks"])
self._tracker_headers = httpx.Headers(
{
"Authorization": f"Bearer {self._tracker_api_key}",
"Content-Type": "application/json",
},
encoding="utf-8",
)
self._async_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback
)
super().__init__(**kwargs)
verbose_proxy_logger.debug(
"Ovalix Guardrail initialized: tracker=%s, application_id=%s, pre_checkpoint_id=%s, post_checkpoint_id=%s",
self._tracker_api_base,
self._application_id,
self._pre_checkpoint_id,
self._post_checkpoint_id,
)
def _validate_config(
self, supported_event_hooks: List[GuardrailEventHooks]
) -> None:
"""Ensure required secrets and checkpoint IDs are set; auto-add hooks when IDs are present."""
errors: List[str] = []
if not self._tracker_api_base:
errors.append(
"Tracker API base, set OVALIX_TRACKER_API_BASE or pass tracker_api_base"
)
if not self._tracker_api_key:
errors.append(
"Tracker API key, set OVALIX_TRACKER_API_KEY or pass tracker_api_key"
)
if not self._application_id:
errors.append(
"Application ID, set OVALIX_APPLICATION_ID or pass application_id"
)
if (
not self._pre_checkpoint_id
and GuardrailEventHooks.pre_call in supported_event_hooks
):
errors.append(
"Pre-checkpoint ID, set OVALIX_PRE_CHECKPOINT_ID or pass pre_checkpoint_id"
)
if (
not self._post_checkpoint_id
and GuardrailEventHooks.post_call in supported_event_hooks
):
errors.append(
"Post-checkpoint ID, set OVALIX_POST_CHECKPOINT_ID or pass post_checkpoint_id"
)
if not self._pre_checkpoint_id and not self._post_checkpoint_id:
errors.append(
"Pre-checkpoint ID or Post-checkpoint ID, set OVALIX_PRE_CHECKPOINT_ID or OVALIX_POST_CHECKPOINT_ID or pass pre_checkpoint_id or post_checkpoint_id"
)
if errors:
raise OvalixGuardrailMissingSecrets(
"Missing Ovalix guardrail configuration errors: " + ". ".join(errors)
)
# auto-add hooks when checkpoint IDs are present
if (
self._pre_checkpoint_id
and GuardrailEventHooks.pre_call not in supported_event_hooks
):
supported_event_hooks.append(GuardrailEventHooks.pre_call)
if (
self._post_checkpoint_id
and GuardrailEventHooks.post_call not in supported_event_hooks
):
supported_event_hooks.append(GuardrailEventHooks.post_call)
def _get_actor(self, data: dict) -> str:
"""Return a stable actor identifier from request metadata (e.g. user email or id)."""
metadata = data.get("metadata") or data.get("litellm_metadata") or {}
if metadata.get("user_api_key_user_email"):
return metadata["user_api_key_user_email"]
if metadata.get("user_api_key_user_id"):
return metadata["user_api_key_user_id"]
return "unknown"
def _get_tracker_actor_id(self, data: dict) -> str:
"""Normalize the actor string into a short, stable id for Tracker API payloads."""
# NOTE: this hash is purely for normalization — it collapses an arbitrary actor
# string (email, user id, or "unknown") into a compact, fixed-length, consistent
# key. It is not a privacy/security measure and the actor value is not sensitive,
# so a plain SHA-256 (truncated) is sufficient; no salting/KDF is needed here.
actor_id = self._get_actor(data).encode()
normalized_actor_id = hashlib.sha256(actor_id).hexdigest()[:8]
return normalized_actor_id
def _get_session_id(self, data: dict) -> str:
"""Return a unique identifier for the chat/session (actor + date + application_id)."""
actor_hash = self._get_tracker_actor_id(data)
today = datetime.datetime.now(datetime.timezone.utc).strftime("%Y-%m-%d")
return f"{actor_hash}_{today}_{self._application_id}"
async def _call_checkpoint(
self,
content: str,
checkpoint_id: str,
actor: str,
session_id: str,
) -> Dict[str, Any]:
"""Call the Ovalix Tracker checkpoint API and return the JSON response."""
application_id = self._application_id
if not application_id or not checkpoint_id:
raise ValueError("Ovalix: application_id or checkpoint_id not resolved")
url = f"{self._tracker_api_base}/tracking/custom_application/checkpoint"
headers = dict(self._tracker_headers)
payload = {
"application_id": application_id,
"checkpoint_id": checkpoint_id,
"actor": actor,
"session_id": session_id,
"data_type": "TEXT",
"data": {"content": content},
}
response = await self._async_handler.post(url, headers=headers, json=payload)
response.raise_for_status()
return response.json()
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
"""
Apply Ovalix guardrail to the given inputs (request or response text).
Used by the unified guardrail flow and the /apply_guardrail API.
For "request", uses the pre-checkpoint; for "response", uses the post-checkpoint.
Args:
inputs: Guardrail API inputs (e.g. texts to check).
request_data: Full request payload (messages, metadata, response).
input_type: "request" (pre_call) or "response" (post_call).
logging_obj: Optional logging context.
Returns:
Updated inputs (e.g. with replaced/corrected texts, or unchanged).
"""
if not self._pre_checkpoint_id and not self._post_checkpoint_id:
return inputs
tracker_actor_id = self._get_tracker_actor_id(request_data)
session_id = self._get_session_id(request_data)
texts = inputs.get("texts") or []
if not texts or not isinstance(texts, list):
return inputs
if input_type == "response":
if not self._post_checkpoint_id:
return inputs
corrected_llm_responses = await self._generate_post_guardrail_llm_texts(
texts, tracker_actor_id, session_id, self._post_checkpoint_id
)
return {**inputs, "texts": corrected_llm_responses}
if self._pre_checkpoint_id:
post_guardrail_texts = await self._generate_post_guardrail_llm_texts(
texts, tracker_actor_id, session_id, self._pre_checkpoint_id
)
return {**inputs, "texts": post_guardrail_texts}
return inputs
async def _generate_post_guardrail_llm_texts(
self, texts: List[str], actor: str, session_id: str, checkpoint_id: str
) -> List[str]:
"""Generate post-guardrail LLM responses for the given LLM responses."""
post_guardrail_texts: List[str] = []
is_first_response = True
for llm_response in reversed(texts):
try:
resp = await self._call_checkpoint(
llm_response, checkpoint_id, actor, session_id
)
except Exception as e:
verbose_proxy_logger.exception(
"Ovalix apply_guardrail checkpoint call failed: %s", e
)
raise GuardrailRaisedException(
guardrail_name=self.guardrail_name,
message=f"Ovalix guardrail error: {e!s}",
should_wrap_with_default_message=False,
) from e
action_type = (resp.get("action_type") or "").lower()
blocking_message = (
self._get_trackers_corrected_message(resp)
or BLOCKED_BY_OVALIX_FALLBACK_MESSAGE
)
if action_type == BLOCKED_ACTION_TYPE and is_first_response:
self._block_current_message(blocking_message)
elif action_type == BLOCKED_ACTION_TYPE:
post_guardrail_texts.insert(0, blocking_message)
else:
corrected_text = (
self._get_trackers_corrected_message(resp) or llm_response
)
post_guardrail_texts.insert(0, corrected_text)
is_first_response = False
return post_guardrail_texts
def _block_current_message(self, blocking_message: str) -> None:
"""Raise OvalixGuardrailBlockedException with the given message (no default wrapper)."""
raise OvalixGuardrailBlockedException(
guardrail_name=self.guardrail_name,
message=blocking_message,
should_wrap_with_default_message=False,
)
def _get_trackers_corrected_message(self, resp: dict) -> Optional[str]:
"""Extract corrected/blocking message content from Tracker checkpoint response."""
modified = resp.get("modified_data")
if isinstance(modified, dict) and "content" in modified:
return modified["content"]
return None
@staticmethod
def get_config_model() -> Optional[Type["GuardrailConfigModel"]]:
from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import (
OvalixGuardrailConfigModel,
)
return OvalixGuardrailConfigModel

View file

@ -130,6 +130,7 @@ services = Union[
"generic_api",
"arize",
"galileo",
"newrelic",
"sqs",
],
str,
@ -208,6 +209,7 @@ async def health_services_endpoint( # noqa: PLR0915
"generic_api",
"arize",
"galileo",
"newrelic",
"sqs",
]:
raise HTTPException(
@ -325,6 +327,26 @@ async def health_services_endpoint( # noqa: PLR0915
"status": "success",
"message": "Mock LLM request made - check langfuse.",
}
elif service == "newrelic":
if not _is_proxy_admin(user_api_key_dict):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={
"error": "Only proxy admins can trigger the New Relic test event."
},
)
from litellm.integrations.newrelic.newrelic import NewRelicLogger
newrelic_logger = NewRelicLogger()
response = await newrelic_logger.async_health_check()
return {
"status": response["status"],
"message": (
response["error_message"]
if response["status"] == "unhealthy"
else "New Relic is healthy — test event sent"
),
}
if service == "webhook":
user_info = CallInfo(

View file

@ -2954,6 +2954,46 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
f"Error in rate limit failure event: {str(e)}"
)
async def async_release_max_parallel_requests_on_disconnect(
self, user_api_key_dict: UserAPIKeyAuth
) -> None:
"""
Release the api-key ``max_parallel_requests`` slot that
``async_pre_call_hook`` reserved, for a request that ended without
either logging callback firing.
The +1 is normally undone by ``async_log_success_event`` (natural
stream completion) or ``async_log_failure_event`` (LLM error). When a
client cancels a stream mid-flight, the cancellation surfaces as
``asyncio.CancelledError`` / ``GeneratorExit`` and neither callback
runs, so without this the counter leaks one slot per cancelled stream
until the key wedges at its limit.
"""
if (
not user_api_key_dict.api_key
or user_api_key_dict.max_parallel_requests is None
):
return
await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline(
increment_list=[
RedisPipelineIncrementOperation(
key=self.create_rate_limit_keys(
key="api_key",
value=user_api_key_dict.api_key,
rate_limit_type="max_parallel_requests",
),
increment_value=-1,
# Refresh the window TTL on the decrement, matching the
# failure path. max_parallel_requests is a concurrency
# gauge, not a rolling-window count, so the key must
# outlive in-flight requests rather than expire mid-stream.
ttl=self.window_size,
)
],
litellm_parent_otel_span=None,
)
async def async_post_call_success_hook(
self, data: dict, user_api_key_dict: UserAPIKeyAuth, response
):

View file

@ -924,11 +924,21 @@ async def get_daily_activity(
where=where_conditions
)
# Fetch paginated results
# Fetch paginated results.
# ``date`` alone is not a unique sort key -- a busy tenant has many
# rows per date (one per api_key, model, model_group, provider,
# endpoint, ...), so offset pagination over ``date desc`` lands on
# arbitrary boundaries and the same row can be skipped on one page
# and returned on another. A client that pages through and sums the
# per-page metrics (the Usage dashboard) then gets a non-deterministic
# total. Adding ``id`` (the row's UUID primary key, present on both
# LiteLLM_DailyUserSpend and LiteLLM_DailyTeamSpend) as a tiebreaker
# gives every page a stable cursor (#30164).
daily_spend_data = await getattr(prisma_client.db, table_name).find_many(
where=where_conditions,
order=[
{"date": "desc"},
{"id": "asc"},
],
skip=(page - 1) * page_size,
take=page_size,

View file

@ -1860,18 +1860,23 @@ async def prepare_key_update_data(
non_default_values["budget_reset_at"] = key_reset_at
non_default_values["budget_duration"] = budget_duration
if "budget_limits" in non_default_values and non_default_values["budget_limits"]:
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
if "budget_limits" in non_default_values:
raw_windows = non_default_values["budget_limits"]
initialized_windows = []
for window in raw_windows:
w = window if isinstance(window, dict) else window.model_dump()
w["reset_at"] = get_budget_reset_time(
budget_duration=w["budget_duration"]
).isoformat()
initialized_windows.append(w)
non_default_values["budget_limits"] = json.dumps(initialized_windows)
if raw_windows:
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
initialized_windows = []
for window in raw_windows:
w = window if isinstance(window, dict) else window.model_dump()
w["reset_at"] = get_budget_reset_time(
budget_duration=w["budget_duration"]
).isoformat()
initialized_windows.append(w)
non_default_values["budget_limits"] = json.dumps(initialized_windows)
else:
# [] / None clears the field; prisma-client-py has no DbNull
# sentinel for Json? columns, so store the JSON literal null
non_default_values["budget_limits"] = json.dumps(None)
if "object_permission" in non_default_values:
non_default_values = await _handle_update_object_permission(
@ -2248,14 +2253,18 @@ async def _validate_update_key_data(
# - Anyone else (non-PROXY_ADMIN, not the owner, not a team member
# on a team key): must pass _check_key_admin_access (PROXY_ADMIN
# / key-owner / team-admin / org-admin of the key).
# - max_budget / spend: always require the admin check, even for the
# key owner or a team member (matches the existing admin-only
# budget semantics).
# - max_budget / spend / budget_limits: always require the admin
# check, even for the key owner or a team member (matches the
# existing admin-only budget semantics). budget_limits uses
# model_fields_set because an explicit null/[] clears the field
# and must gate the same as setting or changing it.
_is_budget_change = (
data.max_budget is not None and data.max_budget != existing_key_row.max_budget
) or (
data.spend is not None
and data.spend != getattr(existing_key_row, "spend", None)
(data.max_budget is not None and data.max_budget != existing_key_row.max_budget)
or (
data.spend is not None
and data.spend != getattr(existing_key_row, "spend", None)
)
or "budget_limits" in data.model_fields_set
)
# Personal-key bypass: the caller both created the key AND still owns it

View file

@ -15,13 +15,14 @@ import datetime
import json
from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast
from fastapi import APIRouter, Depends, HTTPException, Request, status
from fastapi import APIRouter, Depends, HTTPException, Header, Request, status
from pydantic import BaseModel, ConfigDict, Field
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import LITELLM_PROXY_ADMIN_NAME
from litellm.proxy._types import (
BlockModelRequest,
CommonProxyErrors,
LiteLLM_ProxyModelTable,
LiteLLM_TeamTable,
@ -331,6 +332,168 @@ async def patch_model(
)
async def _set_model_blocked_status(
data: BlockModelRequest,
user_api_key_dict: UserAPIKeyAuth,
blocked: bool,
action: Literal["blocked", "unblocked"],
litellm_changed_by: Optional[str],
) -> Optional[LiteLLM_ProxyModelTable]:
from litellm.proxy.proxy_server import (
litellm_proxy_admin_name,
llm_router,
prisma_client,
store_model_in_db,
)
try:
if prisma_client is None:
raise HTTPException(
status_code=500,
detail={"error": CommonProxyErrors.db_not_connected_error.value},
)
if store_model_in_db is not True:
raise ProxyException(
message="Model updates only supported for DB-stored models",
type=ProxyErrorTypes.validation_error.value,
code=status.HTTP_400_BAD_REQUEST,
param=None,
)
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
raise ProxyException(
message="Only proxy admins can change a model's blocked flag.",
type=ProxyErrorTypes.auth_error.value,
code=status.HTTP_403_FORBIDDEN,
param="blocked",
)
db_model = await get_db_model(
model_id=data.model_id,
prisma_client=prisma_client,
)
if db_model is None:
if (
llm_router
and llm_router.get_deployment(model_id=data.model_id) is not None
):
raise ProxyException(
message="Cannot edit config-based model. Store model in DB via /model/new first.",
type=ProxyErrorTypes.validation_error.value,
code=status.HTTP_400_BAD_REQUEST,
param=None,
)
raise ProxyException(
message=f"Model {data.model_id} not found on proxy.",
type=ProxyErrorTypes.not_found_error,
code=status.HTTP_404_NOT_FOUND,
param=None,
)
updated_model = await ModelRepository(prisma_client).table.update(
where={"model_id": data.model_id},
data={
"blocked": blocked,
"updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
"updated_at": cast(str, get_utc_datetime()),
},
)
await clear_cache()
asyncio.create_task(
create_object_audit_log(
object_id=data.model_id,
action=action,
user_api_key_dict=user_api_key_dict,
table_name=LitellmTableNames.PROXY_MODEL_TABLE_NAME,
before_value=db_model.model_dump_json(exclude_none=True),
after_value=(
updated_model.model_dump_json(exclude_none=True)
if isinstance(updated_model, BaseModel)
else None
),
litellm_changed_by=litellm_changed_by,
litellm_proxy_admin_name=litellm_proxy_admin_name,
)
)
return updated_model
except Exception as e:
verbose_proxy_logger.exception(f"Error in model {action}: {str(e)}")
if isinstance(e, (HTTPException, ProxyException)):
raise e
raise ProxyException(
message=f"Error updating model blocked status: {str(e)}",
type=ProxyErrorTypes.internal_server_error,
code=status.HTTP_500_INTERNAL_SERVER_ERROR,
param=None,
)
@router.post(
"/model/block",
tags=["model management"],
dependencies=[Depends(user_api_key_auth)],
)
async def block_model(
data: BlockModelRequest,
http_request: Request,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
litellm_changed_by: Optional[str] = Header(
None,
description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability",
),
) -> Optional[LiteLLM_ProxyModelTable]:
"""
Block a DB-stored model deployment from serving requests.
Parameters:
- model_id: str - The model deployment id to block.
"""
return await _set_model_blocked_status(
data=data,
user_api_key_dict=user_api_key_dict,
blocked=True,
action="blocked",
litellm_changed_by=litellm_changed_by,
)
@router.post(
"/model/unblock",
tags=["model management"],
dependencies=[Depends(user_api_key_auth)],
)
async def unblock_model(
data: BlockModelRequest,
http_request: Request,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
litellm_changed_by: Optional[str] = Header(
None,
description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability",
),
) -> Optional[LiteLLM_ProxyModelTable]:
"""
Unblock a DB-stored model deployment so it can serve requests again.
Parameters:
- model_id: str - The model deployment id to unblock.
"""
return await _set_model_blocked_status(
data=data,
user_api_key_dict=user_api_key_dict,
blocked=False,
action="unblocked",
litellm_changed_by=litellm_changed_by,
)
################################# Helper Functions #################################
####################################################################################
####################################################################################

View file

@ -517,6 +517,13 @@ class AnthropicPassthroughLoggingHandler:
# Process each individual event
for event_str in individual_events:
try:
# Skip OpenAI-style [DONE] sentinels some Anthropic-compatible
# providers emit. Match the whole SSE line so a valid chunk whose
# text payload happens to contain "[DONE]" is not dropped.
if any(
line.strip() == "data: [DONE]" for line in event_str.split("\n")
):
continue
transformed_openai_chunk = anthropic_model_response_iterator.convert_str_chunk_to_generic_chunk(
chunk=event_str
)
@ -525,6 +532,14 @@ class AnthropicPassthroughLoggingHandler:
except (StopIteration, StopAsyncIteration):
break
except json.JSONDecodeError:
# Some upstreams emit non-JSON SSE lines; skip them so the
# logging pipeline is not broken by a single bad frame.
verbose_proxy_logger.debug(
"Skipping non-JSON SSE event: %s",
event_str[:200],
)
continue
complete_streaming_response = litellm.stream_chunk_builder(
chunks=all_openai_chunks,

View file

@ -174,9 +174,7 @@ async def chat_completion_pass_through_endpoint( # noqa: PLR0915
data["adapter_id"] = adapter_id
verbose_proxy_logger.debug(
"Request received by LiteLLM:\n{}".format(json.dumps(data, indent=4)),
)
verbose_proxy_logger.debug("Request received by LiteLLM:\n%s", data)
data["model"] = (
general_settings.get("completion_model", None) # server default
or user_model # model name passed via cli args
@ -298,7 +296,7 @@ async def chat_completion_pass_through_endpoint( # noqa: PLR0915
)
)
verbose_proxy_logger.debug("\nResponse from Litellm:\n{}".format(response))
verbose_proxy_logger.debug("\nResponse from Litellm:\n%s", response)
return response
except Exception as e:
await proxy_logging_obj.post_call_failure_hook(
@ -811,9 +809,10 @@ async def pass_through_request( # noqa: PLR0915
else:
_parsed_body = await _read_request_body(request)
verbose_proxy_logger.debug(
"Pass through endpoint sending request to \nURL {}\nheaders: {}\nbody: {}\n".format(
url, headers, _parsed_body
)
"Pass through endpoint sending request to \nURL %s\nheaders: %s\nbody: %s\n",
url,
headers,
_parsed_body,
)
### COLLECT GUARDRAILS FOR PASSTHROUGH ENDPOINT ###

View file

@ -252,6 +252,7 @@ from litellm.proxy.analytics_endpoints.analytics_endpoints import (
)
from litellm.proxy.auth.auth_checks import (
ExperimentalUIJWTToken,
can_key_call_resolved_model,
get_team_object,
log_db_metrics,
)
@ -7112,6 +7113,17 @@ async def async_data_generator( # noqa: PLR0915
if not request_data.get("_litellm_skip_openai_stream_done"):
done_message = "[DONE]"
yield f"data: {done_message}\n\n"
except (asyncio.CancelledError, GeneratorExit):
# Client disconnected mid-stream. CancelledError / GeneratorExit are
# BaseException, so they bypass the success/failure logging callbacks
# that normally release the pre-call max_parallel_requests +1; release
# it here. This is the outermost generator Starlette closes on
# disconnect, so it fires reliably regardless of needs_iterator_wrap
# (a nested iterator hook would only see GeneratorExit on GC).
proxy_logging_obj._release_max_parallel_requests_on_disconnect(
user_api_key_dict
)
raise
except Exception as e:
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.async_data_generator(): Exception occured - {}".format(
@ -7839,6 +7851,15 @@ class ProxyStartupEvent:
)
await VantageLogger.init_vantage_background_job(scheduler=scheduler)
########################################################
# Mavvrik FOCUS Background Job
########################################################
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import ( # noqa: PLC0415
MavvrikFocusLogger,
)
await MavvrikFocusLogger.init_mavvrik_focus_background_job(scheduler=scheduler)
########################################################
# Prometheus Background Job
########################################################
@ -8224,6 +8245,7 @@ async def model_list(
include_metadata: Optional[bool] = False,
fallback_type: Optional[str] = None,
scope: Optional[str] = None,
healthy_only: Optional[bool] = False,
):
"""
Use `/model/info` - to get detailed model information, example - pricing, mode, etc.
@ -8237,6 +8259,15 @@ async def model_list(
- scope: Optional scope parameter. Currently only accepts "expand".
When scope=expand is passed, proxy admins, team admins, and org admins
will receive all proxy models as if they are a proxy admin.
- healthy_only: When true, hide models whose backing deployments are all marked
unhealthy by background health checks. Requires
`background_health_checks: true` in general_settings; without
health state the listing is returned unfiltered (fail open).
Models expanded from wildcard routes (e.g. `openai/*`) are not
filtered, and nothing is hidden when `allowed_fails_policy` is
configured (cooldown remains the sole exclusion mechanism).
Hiding is presentation-only: a hidden model can still be
called directly.
"""
global llm_model_list, general_settings, llm_router, prisma_client, user_api_key_cache, proxy_logging_obj
@ -8270,6 +8301,19 @@ async def model_list(
llm_router.get_fully_blocked_model_names() if llm_router is not None else set()
)
# Opt-in: also hide models whose deployments are all unhealthy per background
# health checks. Empty when health state is unavailable or stale (fail open).
unhealthy_names: Set[str] = set()
if healthy_only and llm_router is not None:
unhealthy_names = await llm_router.async_get_fully_unhealthy_model_names()
if not unhealthy_names:
verbose_proxy_logger.debug(
"healthy_only=true but no unhealthy deployment state is available "
"(requires background_health_checks); returning unfiltered model list"
)
hidden_names = blocked_names | unhealthy_names
# If scope=expand and user has admin privileges, return all proxy models
if should_expand_scope:
# Get all proxy models as if user is a proxy admin
@ -8302,9 +8346,9 @@ async def model_list(
only_model_access_groups=only_model_access_groups or False,
)
# Hide paused models from the public listing (admins manage them via /model/info)
if blocked_names:
all_models = [m for m in all_models if m not in blocked_names]
# 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]
# Build response data with all proxy models
model_data = []
@ -8339,9 +8383,9 @@ async def model_list(
user_api_key_cache=user_api_key_cache,
)
# Hide paused models from the public listing (admins manage them via /model/info)
if blocked_names:
all_models = [m for m in all_models if m not in blocked_names]
# 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]
# Build response data
model_data = []
@ -9458,13 +9502,15 @@ async def vertex_ai_live_passthrough_endpoint(
@lru_cache(maxsize=_REALTIME_BODY_CACHE_SIZE)
def _realtime_query_params_template(
model: str, intent: Optional[str]
model: Optional[str], intent: Optional[str]
) -> Tuple[Tuple[str, str], ...]:
"""
Build a hashable representation of the realtime query params so we can cache
the repetitive model/intent combinations.
"""
params: List[Tuple[str, str]] = [("model", model)]
params: List[Tuple[str, str]] = []
if model is not None:
params.append(("model", model))
if intent is not None:
params.append(("intent", intent))
return tuple(params)
@ -9475,8 +9521,10 @@ def _realtime_query_params_template(
@app.websocket("/realtime")
async def realtime_websocket_endpoint(
websocket: WebSocket,
model: str,
intent: str = fastapi.Query(
model: Optional[str] = fastapi.Query(
None, description="The model to use for the websocket connection."
),
intent: Optional[str] = fastapi.Query(
None, description="The intent of the websocket connection."
),
guardrails: Optional[str] = fastapi.Query(
@ -9493,6 +9541,25 @@ async def realtime_websocket_endpoint(
accept_kwargs: dict = {}
if requested_protocols:
accept_kwargs["subprotocol"] = requested_protocols[0]
route_model = model
if route_model is None:
if intent == "transcription":
route_model = "gpt-realtime-whisper"
else:
await websocket.close(code=1008, reason="model query parameter is required")
return
assert route_model is not None
try:
await can_key_call_resolved_model(
model=route_model,
llm_model_list=llm_model_list,
valid_token=user_api_key_dict,
llm_router=llm_router,
)
except ProxyException as e:
await websocket.close(code=1008, reason=e.message[:120])
return
await websocket.accept(**accept_kwargs)
# Only use explicit parameters, not all query params
@ -9501,7 +9568,7 @@ async def realtime_websocket_endpoint(
)
data: Dict[str, Any] = {
"model": model,
"model": route_model,
"websocket": websocket,
"query_params": query_params, # Only explicit params
}
@ -9521,7 +9588,7 @@ async def realtime_websocket_endpoint(
request._url = websocket.url
async def return_body():
return _realtime_request_body(model)
return _realtime_request_body(route_model)
request.body = return_body # type: ignore
@ -9547,7 +9614,7 @@ async def realtime_websocket_endpoint(
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
model=model,
model=route_model,
route_type="_arealtime",
)
except Exception as e:

View file

@ -10,6 +10,7 @@ from fastapi import status as http_status
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import can_key_call_resolved_model
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
@ -19,11 +20,143 @@ from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.types.realtime import (
RealtimeClientSecretRequest,
RealtimeClientSecretResponse,
RealtimeTranscriptionSessionRequest,
RealtimeTranscriptionSessionResponse,
)
router = APIRouter()
_REALTIME_TOKEN_VERSION = "realtime_v1"
_DEFAULT_REALTIME_MODEL = "gpt-4o-realtime-preview"
_DEFAULT_TRANSCRIPTION_MODEL = "gpt-realtime-whisper"
_ALLOWED_SESSION_TYPES = ("realtime", "transcription")
def _coerce_realtime_session_type(session_type: Optional[str]) -> str:
if session_type in _ALLOWED_SESSION_TYPES:
return session_type
return "realtime"
def _append_model_candidate(candidates: list[str], model: Any) -> None:
if isinstance(model, str) and model and model not in candidates:
candidates.append(model)
def _transcription_model_candidates_from_session(session: dict) -> list[str]:
candidates: list[str] = []
audio = session.get("audio")
if isinstance(audio, dict):
audio_input = audio.get("input")
if isinstance(audio_input, dict):
nested_transcription = audio_input.get("transcription")
if isinstance(nested_transcription, dict):
_append_model_candidate(
candidates,
nested_transcription.get("model"),
)
flat_transcription = session.get("input_audio_transcription")
if isinstance(flat_transcription, dict):
_append_model_candidate(candidates, flat_transcription.get("model"))
return candidates
def _set_transcription_model_on_session(
session: dict,
model: str,
create_if_missing: bool = False,
) -> None:
updated_existing_config = False
flat_transcription = session.get("input_audio_transcription")
if isinstance(flat_transcription, dict):
session["input_audio_transcription"] = {
**flat_transcription,
"model": model,
}
updated_existing_config = True
audio = session.get("audio")
if isinstance(audio, dict):
audio_input = audio.get("input")
if isinstance(audio_input, dict):
nested_transcription = audio_input.get("transcription")
if isinstance(nested_transcription, dict):
session["audio"] = {
**audio,
"input": {
**audio_input,
"transcription": {
**nested_transcription,
"model": model,
},
},
}
updated_existing_config = True
if updated_existing_config or not create_if_missing:
return
audio = audio if isinstance(audio, dict) else {}
audio_input = audio.get("input")
audio_input = audio_input if isinstance(audio_input, dict) else {}
session["audio"] = {
**audio,
"input": {
**audio_input,
"transcription": {"model": model},
},
}
async def _prepare_client_secret_session(
req: RealtimeClientSecretRequest,
user_api_key_dict: UserAPIKeyAuth,
llm_model_list: Optional[list],
llm_router: Any,
) -> tuple[str, Optional[dict], str]:
session_type = _coerce_realtime_session_type(
req.session.type if req.session else None
)
session_data: Optional[dict] = (
req.session.model_dump(exclude_none=True) if req.session else None
)
if session_data is not None:
session_data["type"] = session_type
session_model = req.session.model if req.session else None
model: str = session_model or req.model or _DEFAULT_REALTIME_MODEL
if session_type != "transcription":
return model, session_data, session_type
transcription_model_candidates = _transcription_model_candidates_from_session(
session_data or {}
)
if not transcription_model_candidates:
_append_model_candidate(transcription_model_candidates, session_model)
_append_model_candidate(transcription_model_candidates, req.model)
if not transcription_model_candidates:
transcription_model_candidates.append(_DEFAULT_TRANSCRIPTION_MODEL)
model = transcription_model_candidates[0]
for transcription_model in transcription_model_candidates:
await can_key_call_resolved_model(
model=transcription_model,
valid_token=user_api_key_dict,
llm_model_list=llm_model_list,
llm_router=llm_router,
)
if session_data is not None:
_set_transcription_model_on_session(
session=session_data,
model=model,
create_if_missing=True,
)
session_data.pop("model", None)
return model, session_data, session_type
def _encode_realtime_token_payload(
@ -32,6 +165,7 @@ def _encode_realtime_token_payload(
user_id: Optional[str],
team_id: Optional[str],
expires_at: Optional[int],
session_type: str = "realtime",
) -> str:
"""
Encode metadata with the upstream ephemeral key so /realtime/calls can
@ -44,6 +178,7 @@ def _encode_realtime_token_payload(
"user_id": user_id or "",
"team_id": team_id or "",
"expires_at": expires_at,
"session_type": session_type,
}
return json.dumps(payload, separators=(",", ":"))
@ -94,6 +229,7 @@ async def create_realtime_client_secret(
add_litellm_data_to_request,
general_settings,
llm_router,
llm_model_list,
proxy_config,
proxy_logging_obj,
route_request,
@ -106,17 +242,18 @@ async def create_realtime_client_secret(
body = await _read_request_body(request=request)
req = RealtimeClientSecretRequest(**body)
model: str = (
(req.session.model if req.session else None)
or req.model
or "gpt-4o-realtime-preview"
model, session_data, session_type = await _prepare_client_secret_session(
req=req,
user_api_key_dict=user_api_key_dict,
llm_model_list=llm_model_list,
llm_router=llm_router,
)
data = {"model": model}
# If session is provided, use it; otherwise create one from model
if req.session:
data["session"] = req.session.model_dump(exclude_none=True)
if session_data is not None:
data["session"] = session_data
elif req.model:
# User provided model at root level, convert to session format
data["session"] = {"type": "realtime", "model": model}
@ -161,6 +298,8 @@ async def create_realtime_client_secret(
"litellm.proxy.realtime_endpoints.webrtc.create_realtime_client_secret(): Exception - %s",
str(e),
)
if isinstance(e, ProxyException):
raise e
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "message", str(e)),
@ -199,6 +338,7 @@ async def create_realtime_client_secret(
user_id=getattr(user_api_key_dict, "user_id", None),
team_id=getattr(user_api_key_dict, "team_id", None),
expires_at=expires_at if isinstance(expires_at, int) else None,
session_type=session_type,
)
encrypted_token: str = encrypt_value_helper(token_payload)
upstream_json["value"] = encrypted_token
@ -279,16 +419,20 @@ async def proxy_realtime_calls(
model = (
decoded_payload.get("model_id")
or request.query_params.get("model")
or "gpt-4o-realtime-preview"
or _DEFAULT_REALTIME_MODEL
)
user_id = decoded_payload.get("user_id") or None
team_id = decoded_payload.get("team_id") or None
session_type = _coerce_realtime_session_type(
decoded_payload.get("session_type")
)
else:
# Backward compatibility: older tokens contained only encrypted upstream key.
openai_ephemeral_key = decrypted_token_value
model = request.query_params.get("model", "gpt-4o-realtime-preview")
model = request.query_params.get("model", _DEFAULT_REALTIME_MODEL)
user_id = None
team_id = None
session_type = "realtime"
# Build a minimal UserAPIKeyAuth with user/team IDs from the token
# so spend tracking and budget enforcement work correctly.
@ -299,11 +443,17 @@ async def proxy_realtime_calls(
data: dict = {}
try:
# Build session config for the multipart form data
session_config = {
"type": "realtime",
"model": model,
"type": session_type,
}
if session_type == "transcription":
_set_transcription_model_on_session(
session=session_config,
model=model,
create_if_missing=True,
)
else:
session_config["model"] = model
data = {
"model": model,
@ -366,3 +516,145 @@ async def proxy_realtime_calls(
status_code=upstream_resp.status_code,
media_type=upstream_resp.headers.get("content-type", "application/sdp"),
)
@router.post(
"/v1/realtime/transcription_sessions",
dependencies=[Depends(user_api_key_auth)],
tags=["realtime"],
)
@router.post(
"/realtime/transcription_sessions",
dependencies=[Depends(user_api_key_auth)],
tags=["realtime"],
)
@router.post(
"/openai/v1/realtime/transcription_sessions",
dependencies=[Depends(user_api_key_auth)],
tags=["realtime"],
)
async def create_realtime_transcription_session(
request: Request,
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> RealtimeTranscriptionSessionResponse:
"""
Create an ephemeral Realtime transcription session
(POST /v1/realtime/transcription_sessions) for the WebRTC/WebSocket flow.
Mirrors the client_secrets route but targets the transcription_sessions
endpoint and encrypts the ephemeral key returned under `client_secret.value`.
"""
from litellm.proxy.proxy_server import (
add_litellm_data_to_request,
general_settings,
llm_router,
llm_model_list,
proxy_config,
proxy_logging_obj,
route_request,
user_model,
version,
)
data: dict = {}
try:
body = await _read_request_body(request=request)
req = RealtimeTranscriptionSessionRequest(**body)
model: str = req.resolved_model() or "gpt-realtime-whisper"
await can_key_call_resolved_model(
model=model,
valid_token=user_api_key_dict,
llm_model_list=llm_model_list,
llm_router=llm_router,
)
transcription_session = {k: v for k, v in body.items() if k != "model"}
data = {"model": model, "transcription_session": transcription_session}
data = await add_litellm_data_to_request(
data=data,
request=request,
general_settings=general_settings,
user_api_key_dict=user_api_key_dict,
version=version,
proxy_config=proxy_config,
)
data = await proxy_logging_obj.pre_call_hook(
user_api_key_dict=user_api_key_dict,
data=data,
call_type="acreate_realtime_transcription_session",
)
verbose_proxy_logger.debug(
"Realtime: /v1/realtime/transcription_sessions (model=%s)", model
)
llm_call = await route_request(
data=data,
route_type="acreate_realtime_transcription_session",
llm_router=llm_router,
user_model=user_model,
)
upstream_resp: httpx.Response = await llm_call # type: ignore
except Exception as e:
await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict,
original_exception=e,
request_data=data,
)
verbose_proxy_logger.error(
"litellm.proxy.realtime_endpoints.create_realtime_transcription_session(): Exception - %s",
str(e),
)
if isinstance(e, ProxyException):
raise e
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "detail", getattr(e, "message", str(e))),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", http_status.HTTP_400_BAD_REQUEST),
)
raise ProxyException(
message=getattr(e, "message", str(e)),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", 500),
)
if upstream_resp.status_code != 200:
verbose_proxy_logger.error(
"Realtime transcription_sessions upstream error %s: %s",
upstream_resp.status_code,
upstream_resp.text,
)
return Response( # type: ignore[return-value]
content=upstream_resp.content,
status_code=upstream_resp.status_code,
media_type="application/json",
)
upstream_json: dict = upstream_resp.json()
# Encrypt the ephemeral key (returned under client_secret.value) with routing
# metadata so the follow-up /realtime/calls request can recover the model.
client_secret = upstream_json.get("client_secret")
if isinstance(client_secret, dict) and "value" in client_secret:
raw_value: str = client_secret.get("value", "")
expires_at = client_secret.get("expires_at")
token_payload = _encode_realtime_token_payload(
ephemeral_key=raw_value,
model_id=model,
user_id=getattr(user_api_key_dict, "user_id", None),
team_id=getattr(user_api_key_dict, "team_id", None),
expires_at=expires_at if isinstance(expires_at, int) else None,
session_type="transcription",
)
client_secret["value"] = encrypt_value_helper(token_payload)
upstream_json["client_secret"] = client_secret
return RealtimeTranscriptionSessionResponse(**upstream_json)

View file

@ -1,6 +1,7 @@
import asyncio
from typing import TYPE_CHECKING, Any, Literal, Optional
import httpx
from fastapi import HTTPException, status
import litellm
@ -46,6 +47,30 @@ def _is_a2a_agent_model(model_name: Any) -> bool:
return isinstance(model_name, str) and model_name.startswith("a2a/")
def _raise_if_model_fully_blocked(
llm_router: LitellmRouter, model_name: Any, team_id: Optional[str]
) -> None:
if not isinstance(model_name, str) or not model_name:
return
if not isinstance(llm_router, litellm.Router):
return
deployments = (
llm_router.get_model_list(model_name=model_name, team_id=team_id) or []
)
if llm_router._are_all_deployments_blocked(deployments):
raise litellm.PermissionDeniedError(
message="Model is blocked",
model=model_name,
llm_provider="",
response=httpx.Response(
status_code=403,
request=httpx.Request(
method="POST", url="https://github.com/BerriAI/litellm"
),
),
)
ROUTE_ENDPOINT_MAPPING = {
"acompletion": "/chat/completions",
"atext_completion": "/completions",
@ -74,6 +99,7 @@ ROUTE_ENDPOINT_MAPPING = {
"avideo_extension": "/videos/extensions",
"acreate_realtime_client_secret": "/realtime/client_secrets",
"arealtime_calls": "/realtime/calls",
"acreate_realtime_transcription_session": "/realtime/transcription_sessions",
"acreate_container": "/containers",
"alist_containers": "/containers",
"aretrieve_container": "/containers/{container_id}",
@ -261,6 +287,7 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin
"_arealtime", # private function for realtime API
"acreate_realtime_client_secret",
"arealtime_calls",
"acreate_realtime_transcription_session",
"_aresponses_websocket", # private function for responses WebSocket mode
"aimage_edit",
"agenerate_content",
@ -411,6 +438,9 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin
else:
return getattr(litellm, f"{route_type}")(**data)
elif llm_router is not None:
_raise_if_model_fully_blocked(
llm_router=llm_router, model_name=data.get("model"), team_id=team_id
)
# Evals API: always route to litellm directly (not through router)
# But extract model credentials if a model is provided
if route_type in [
@ -427,6 +457,7 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin
"adelete_run",
"acreate_realtime_client_secret",
"arealtime_calls",
"acreate_realtime_transcription_session",
]:
# If a model is provided, get its credentials from the router
model = data.get("model")

View file

@ -2142,6 +2142,20 @@ async def ui_view_spend_logs( # noqa: PLR0915
data = await prisma_client.db.query_raw(sql_query, *sql_params)
# query_raw returns the JSONB `metadata` column as a string (the Prisma
# serialiser bypasses the model-layer JSON hydration we get on the ORM
# path). The UI reads `metadata.status` / `metadata.error_information`
# as object fields, so failure rows looked like successes (#29674).
# Re-hydrate to dict here.
for row in data:
if isinstance(row, dict):
md = row.get("metadata")
if isinstance(md, str):
try:
row["metadata"] = json.loads(md)
except (ValueError, TypeError):
row["metadata"] = {}
# Calculate total pages
total_pages = (total_records + page_size - 1) // page_size

View file

@ -137,6 +137,9 @@ from litellm.proxy.hooks.max_budget_limiter import _PROXY_MaxBudgetLimiter
from litellm.proxy.hooks.parallel_request_limiter import (
_PROXY_MaxParallelRequestsHandler,
)
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor
from litellm.repositories.budget_repository import BudgetRepository
@ -2688,6 +2691,42 @@ class ProxyLogging:
logging_obj._deferred_stream_complete_args = None
asyncio.create_task(_deferred_cb(*_args))
def _release_max_parallel_requests_on_disconnect(
self, user_api_key_dict: UserAPIKeyAuth
) -> None:
"""
Release the api-key max_parallel_requests slot when a streaming
response is cancelled mid-flight (client disconnect). Neither the
success nor failure logging callback fires on the resulting
CancelledError / GeneratorExit, so the pre-call +1 would otherwise
leak.
Must be called from the outermost streaming generator (the one
Starlette drives and closes on disconnect). A nested iterator-hook
generator only receives GeneratorExit when it is garbage collected,
which is non-deterministic, so the refund cannot live there.
Scheduled fire-and-forget (no await) because awaiting is not
permitted while unwinding a GeneratorExit.
"""
limiter = self.get_proxy_hook("parallel_request_limiter")
if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3):
return
try:
asyncio.create_task(
limiter.async_release_max_parallel_requests_on_disconnect(
user_api_key_dict
)
)
except RuntimeError:
# No running event loop (e.g. interpreter/loop shutdown); the
# counter's window TTL will reclaim the slot.
verbose_proxy_logger.warning(
"parallel_request_limiter_v3: could not schedule "
"max_parallel_requests release on disconnect; no running "
"event loop. Slot will be reclaimed when its window TTL expires"
)
def _init_response_taking_too_long_task(self, data: Optional[dict] = None):
"""
Initialize the response taking too long task if user is using slack alerting

View file

@ -1 +1,9 @@
Abstraction / Routing logic for OpenAI's `/v1/realtime` endpoint.
Abstraction / Routing logic for OpenAI's `/v1/realtime` endpoints.
Supported endpoints:
- WebSocket: `/v1/realtime` (with `intent=transcription` for transcription-only sessions)
- HTTP: `/v1/realtime/client_secrets`, `/v1/realtime/transcription_sessions`
Supported providers: OpenAI, Azure OpenAI, Bedrock, Vertex AI, xAI.
For user-facing documentation and usage examples, see the litellm-docs repo.

View file

@ -15,6 +15,7 @@ from litellm.types.realtime import (
RealtimeExpiresAfter,
RealtimeQueryParams,
RealtimeSessionConfig,
RealtimeTranscriptionSessionRequest,
)
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders
@ -159,6 +160,78 @@ async def acreate_realtime_client_secret(
)
@wrapper_client
async def acreate_realtime_transcription_session(
model: Optional[str] = None,
transcription_session: Optional[Dict[str, Any]] = None,
timeout: Optional[float] = None,
**kwargs,
):
"""
Create an ephemeral transcription session via POST
/v1/realtime/transcription_sessions.
``transcription_session`` is the upstream request body (input_audio_format,
input_audio_transcription, turn_detection, …). ``model`` is a LiteLLM-only
routing hint; the provider model lives in
``transcription_session.input_audio_transcription.model``.
"""
req = RealtimeTranscriptionSessionRequest(
model=model,
**(transcription_session or {}),
)
model_name = req.resolved_model() or "gpt-realtime-whisper"
litellm_logging_obj: LiteLLMLogging = kwargs.get("litellm_logging_obj") # type: ignore
litellm_params = GenericLiteLLMParams(**kwargs)
(
model_name,
custom_llm_provider,
dynamic_api_key,
dynamic_api_base,
) = get_llm_provider(
model=model_name,
api_base=litellm_params.api_base,
api_key=litellm_params.api_key,
)
(
provider_config,
resolved_api_base,
resolved_api_key,
) = _get_realtime_http_provider_config(
custom_llm_provider=custom_llm_provider,
dynamic_api_base=dynamic_api_base,
dynamic_api_key=dynamic_api_key,
litellm_params=litellm_params,
)
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model_name,
optional_params={"transcription_session": transcription_session},
litellm_params={"api_base": resolved_api_base},
custom_llm_provider=custom_llm_provider,
)
request_data = req.model_dump(exclude_none=True, exclude={"model"})
# Ensure the upstream body's input_audio_transcription.model matches the
# authorized routing model. This prevents a caller from supplying an allowed
# top-level model for auth while sneaking a different model into the nested
# transcription config that gets forwarded to the provider.
if isinstance(request_data.get("input_audio_transcription"), dict):
request_data["input_audio_transcription"]["model"] = model_name
return await base_llm_http_handler.async_realtime_transcription_session_handler(
api_base=resolved_api_base,
api_key=resolved_api_key,
request_data=request_data,
logging_obj=litellm_logging_obj,
timeout=timeout or request_timeout,
provider_config=provider_config,
model=model_name,
extra_headers=kwargs.get("extra_headers"),
client=kwargs.get("client"),
api_version=litellm_params.api_version,
)
@wrapper_client
async def arealtime_calls(
openai_ephemeral_key: str,
@ -246,9 +319,13 @@ async def _arealtime( # noqa: PLR0915
api_key=api_key,
)
# Ensure query params use the normalized provider model (no proxy aliases).
# If the client supplied `model` in the URL, ensure it uses the normalized
# provider model (no proxy aliases). If they omitted it, preserve that shape
# for transcription-only sessions like OpenAI's `?intent=transcription`.
if query_params is not None:
query_params = {**query_params, "model": model}
query_params = {**query_params}
if "model" in query_params:
query_params["model"] = model
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
@ -278,6 +355,7 @@ async def _arealtime( # noqa: PLR0915
headers=headers,
user_api_key_dict=kwargs.get("user_api_key_dict"),
litellm_metadata=_build_litellm_metadata(kwargs),
query_params=query_params,
)
elif _custom_llm_provider == "azure":
api_base = (
@ -300,8 +378,13 @@ async def _arealtime( # noqa: PLR0915
kwargs.get("realtime_protocol")
or litellm_params.get("realtime_protocol")
or os.environ.get("LITELLM_AZURE_REALTIME_PROTOCOL")
or "beta"
)
if (
realtime_protocol is None
and (query_params or {}).get("intent") == "transcription"
):
realtime_protocol = "GA"
realtime_protocol = realtime_protocol or "beta"
await azure_realtime.async_realtime(
model=model,
websocket=websocket,
@ -313,6 +396,7 @@ async def _arealtime( # noqa: PLR0915
timeout=timeout,
logging_obj=litellm_logging_obj,
realtime_protocol=realtime_protocol,
query_params=query_params,
user_api_key_dict=kwargs.get("user_api_key_dict"),
litellm_metadata=_build_litellm_metadata(kwargs),
)
@ -450,6 +534,7 @@ async def _arealtime( # noqa: PLR0915
headers=headers,
user_api_key_dict=kwargs.get("user_api_key_dict"),
litellm_metadata=_build_litellm_metadata(kwargs),
query_params=query_params,
)
else:
raise ValueError(f"Unsupported model: {model}")

View file

@ -3381,7 +3381,10 @@ class Router:
# Request Number X, Model Number Y
_tasks.append(
_async_completion_no_exceptions_return_idx(
model=model, idx=idx, messages=message, **kwargs # type: ignore
model=model,
idx=idx,
messages=message, # type: ignore[arg-type]
**kwargs,
)
)
responses = await asyncio.gather(*_tasks)
@ -3544,7 +3547,7 @@ class Router:
self, model: str, messages: List[AllMessageValues], priority: int, stream: Literal[False] = False, **kwargs
) -> ModelResponse:
...
@overload
async def schedule_acompletion(
self, model: str, messages: List[AllMessageValues], priority: int, stream: Literal[True], **kwargs
@ -7937,8 +7940,7 @@ class Router:
re.compile(pattern)
except re.error as exc:
raise ValueError(
f"Invalid regex in tag_regex for model '{deployment.model_name}': "
f"{pattern!r} — {exc}"
f"Invalid regex in tag_regex for model '{deployment.model_name}': {pattern!r} — {exc}"
) from exc
deployment = self._add_deployment(deployment=deployment)
@ -8196,8 +8198,7 @@ class Router:
if deployment.model_name in self.adaptive_routers:
raise ValueError(
f"Adaptive-router deployment {deployment.model_name} already exists. "
"Please use a different model name."
f"Adaptive-router deployment {deployment.model_name} already exists. Please use a different model name."
)
adaptive_router = AdaptiveRouter(
@ -9407,8 +9408,7 @@ class Router:
):
model_group_info.supports_parallel_function_calling = True
if (
model_info.get("supports_vision", None) is not None
and model_info["supports_vision"] is True # type: ignore
model_info.get("supports_vision", None) is not None and model_info["supports_vision"] is True # type: ignore
):
model_group_info.supports_vision = True
if (
@ -9428,8 +9428,7 @@ class Router:
model_group_info.supports_url_context = True
if (
model_info.get("supports_reasoning", None) is not None
and model_info["supports_reasoning"] is True # type: ignore
model_info.get("supports_reasoning", None) is not None and model_info["supports_reasoning"] is True # type: ignore
):
model_group_info.supports_reasoning = True
if (
@ -10052,6 +10051,76 @@ class Router:
name for name, fully_blocked in blocked_by_name.items() if fully_blocked
}
@staticmethod
def _are_all_deployments_blocked(
deployments: List[DeploymentTypedDict],
) -> bool:
return len(deployments) > 0 and all(
(deployment.get("model_info") or {}).get("blocked") is True
for deployment in deployments
)
def _is_model_fully_blocked(self, model: str) -> bool:
deployments = self.get_model_list(model_name=model) or []
return self._are_all_deployments_blocked(deployments=deployments)
async def async_get_fully_unhealthy_model_names(self) -> Set[str]:
"""
Returns the set of model names where every backing deployment is currently
marked unhealthy by background health checks (and the health state is not stale).
Used by `/v1/models?healthy_only=true` to hide models that cannot serve any
request. A model with at least one healthy (or unknown-health) deployment
remains visible. Returns an empty set when no health state is available, so
callers fail open to the unfiltered listing.
Notes:
- Mirrors `_async_filter_health_check_unhealthy_deployments`: when
`allowed_fails_policy` is set, cooldown is the sole routing exclusion
mechanism, so nothing is hidden here either.
- Team-specific public model names (`team_public_model_name`) are
aggregated alongside `model_name`, so team aliases of fully-unhealthy
deployments are hidden too (unlike `get_fully_blocked_model_names`,
which matches `model_name` only).
- Wildcard routes (e.g. `openai/*`) are matched by their literal
deployment name only; models expanded from a wildcard route are not
hidden (fail open).
- Intentionally diverges from the routing-time safety net (which
bypasses the health filter when every candidate is unhealthy and
still attempts the request): hiding here is presentation-only —
it answers "should this model be advertised?", not "should a
request for it still be attempted?". A hidden model can still be
called directly.
"""
if self.allowed_fails_policy is not None:
return set()
unhealthy_ids = (
await self.health_state_cache.async_get_unhealthy_deployment_ids()
)
if not unhealthy_ids:
return set()
deployments = self.get_model_list() or []
unhealthy_by_name: Dict[str, bool] = {}
for deployment in deployments:
model_info = deployment.get("model_info") or {}
names = [deployment.get("model_name") or ""]
team_public_model_name = model_info.get("team_public_model_name")
if team_public_model_name:
names.append(team_public_model_name)
is_unhealthy = model_info.get("id") in unhealthy_ids
for name in names:
if not name:
continue
if name in unhealthy_by_name:
unhealthy_by_name[name] = unhealthy_by_name[name] and is_unhealthy
else:
unhealthy_by_name[name] = is_unhealthy
return {
name
for name, fully_unhealthy in unhealthy_by_name.items()
if fully_unhealthy
}
def _get_team_specific_model(
self, deployment: DeploymentTypedDict, team_id: Optional[str] = None
) -> Optional[str]:

View file

@ -23,6 +23,9 @@ from litellm.types.proxy.guardrails.guardrail_hooks.ibm import (
from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import (
ContentFilterCategoryConfig,
)
from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import (
OvalixGuardrailConfigModel,
)
from litellm.types.proxy.guardrails.guardrail_hooks.promptguard import (
PromptGuardConfigModel,
)
@ -97,6 +100,7 @@ class SupportedGuardrailIntegrations(Enum):
GENERIC_GUARDRAIL_API = "generic_guardrail_api"
QUALIFIRE = "qualifire"
CUSTOM_CODE = "custom_code"
OVALIX = "ovalix"
MICROSOFT_PURVIEW = "microsoft_purview"
SEMANTIC_GUARD = "semantic_guard"
MCP_END_USER_PERMISSION = "mcp_end_user_permission"
@ -852,6 +856,7 @@ class LitellmParams(
BaseLitellmParams,
EnkryptAIGuardrailConfigs,
IBMGuardrailsBaseConfigModel,
OvalixGuardrailConfigModel,
QualifireGuardrailConfigModel,
BlockCodeExecutionGuardrailConfigModel,
HiddenlayerGuardrailConfigModel,

View file

@ -0,0 +1,9 @@
from litellm.types.integrations.custom_logger import StandardCustomLoggerInitParams
class NewRelicInitParams(StandardCustomLoggerInitParams):
"""
Params for initializing a New Relic logger on litellm
"""
pass

View file

@ -291,25 +291,6 @@ class CohereToolResult(BaseModel):
outputs: List[Dict[str, Any]]
class CohereResponseFormat(BaseModel):
"""Response format for Cohere."""
type: str
class CohereResponseTextFormat(CohereResponseFormat):
"""Text response format for Cohere."""
type: Literal["text"] = "text"
class CohereResponseJSONSchemaFormat(CohereResponseFormat):
"""JSON schema response format for Cohere."""
type: Literal["json_schema"] = "json_schema"
jsonSchema: Dict[str, Any]
class CohereChatRequest(BaseModel):
"""Cohere chat request model."""
@ -336,13 +317,10 @@ class CohereChatRequest(BaseModel):
# ``OCIChatConfig.openai_to_oci_cohere_param_map`` which marks
# ``tool_choice`` as unsupported. The field is intentionally absent here
# so it isn't silently dropped or surfaced as a supported feature.
responseFormat: Optional[
Union[
CohereResponseTextFormat,
CohereResponseJSONSchemaFormat,
CohereResponseFormat,
]
] = None
# OCI Cohere responseFormat is {"type": "TEXT" | "JSON_OBJECT", "schema"?: ...};
# there is no JSON_SCHEMA type. The shape is built in
# OCIChatConfig._normalize_response_format.
responseFormat: Optional[Dict[str, Any]] = None
preambleOverride: Optional[str] = None
documents: Optional[List[Dict[str, Any]]] = None
searchQueriesOnly: Optional[bool] = None

View file

@ -802,6 +802,8 @@ ValidUserMessageContentTypes = [
"audio_url",
"document",
"guarded_text",
"grounding_source",
"query",
"video_url",
"file",
] # used for validating user messages. Prevent users from accidentally sending anthropic messages.
@ -813,6 +815,8 @@ ValidUserMessageContentTypesLiteral = Literal[
"audio_url",
"document",
"guarded_text",
"grounding_source",
"query",
"video_url",
"file",
]
@ -824,6 +828,8 @@ ValidUserMessageContentTypes = [
"audio_url",
"document",
"guarded_text",
"grounding_source",
"query",
"video_url",
"file",
] # used for validating user messages. Prevent users from accidentally sending anthropic messages.
@ -851,6 +857,8 @@ ValidChatCompletionMessageContentTypesLiteral = Literal[
"audio_url",
"document",
"guarded_text",
"grounding_source",
"query",
"video_url",
"file",
"thinking",
@ -864,6 +872,8 @@ ValidChatCompletionMessageContentTypes = [
"audio_url",
"document",
"guarded_text",
"grounding_source",
"query",
"video_url",
"file",
"thinking",

View file

@ -2,9 +2,14 @@ from typing import Any, Dict, List, Literal, Optional, Union
from typing_extensions import TypedDict
# Bedrock contextual grounding tags each content block so the guardrail knows
# which text is the reference source, the user question, and the content to grade.
BedrockGuardrailQualifier = Literal["grounding_source", "query", "guard_content"]
class BedrockTextContent(TypedDict, total=False):
text: str
qualifiers: List[BedrockGuardrailQualifier]
class BedrockContentItem(TypedDict, total=False):

View file

@ -0,0 +1,37 @@
"""Pydantic config model for the Ovalix guardrail (Tracker API, application and checkpoint IDs)."""
from typing import Optional
from pydantic import Field
from .base import GuardrailConfigModel
class OvalixGuardrailConfigModel(GuardrailConfigModel):
"""Configuration parameters for the Ovalix guardrail (pre/post call checkpoints)."""
tracker_api_base: Optional[str] = Field(
default=None,
description="Base URL for the Ovalix Tracker service.",
)
tracker_api_key: Optional[str] = Field(
default=None,
description="API key for the Ovalix Tracker service.",
)
application_id: Optional[str] = Field(
default=None,
description="Application ID for the Ovalix Tracker service.",
)
pre_checkpoint_id: Optional[str] = Field(
default=None,
description="Pre-checkpoint ID for the Ovalix Tracker service.",
)
post_checkpoint_id: Optional[str] = Field(
default=None,
description="Post-checkpoint ID for the Ovalix Tracker service.",
)
@staticmethod
def ui_friendly_name() -> str:
"""Display name for this guardrail in the proxy UI."""
return "Ovalix Guardrail"

View file

@ -115,3 +115,40 @@ class RealtimeClientSecretResponse(BaseModel):
expires_at: Optional[int] = None
value: str
session: Optional[Dict[str, Any]] = None
class RealtimeTranscriptionSessionRequest(BaseModel):
"""
Request body for POST /v1/realtime/transcription_sessions.
Mirrors OpenAI's RealtimeTranscriptionSessionCreateRequest. The model used
for routing is taken from the LiteLLM-only top-level `model` hint, falling
back to `input_audio_transcription.model`. All other fields pass through
unchanged to the provider.
"""
model_config = {"extra": "allow"}
# LiteLLM-only routing hint — stripped before forwarding upstream.
model: Optional[str] = None
input_audio_transcription: Optional[Dict[str, Any]] = None
def resolved_model(self) -> Optional[str]:
if self.model:
return self.model
if self.input_audio_transcription:
return self.input_audio_transcription.get("model")
return None
class RealtimeTranscriptionSessionResponse(BaseModel):
"""
Response from POST /v1/realtime/transcription_sessions.
`client_secret.value` contains the encrypted token instead of the raw
ephemeral key. Unknown fields pass through unchanged.
"""
model_config = {"extra": "allow"}
client_secret: Optional[Dict[str, Any]] = None

View file

@ -533,6 +533,7 @@ CallTypesLiteral = Literal[
"acreate_skill",
"acreate_realtime_client_secret",
"arealtime_calls",
"acreate_realtime_transcription_session",
]
# Mapping of API routes to their corresponding call types
@ -2493,6 +2494,7 @@ class LoggedLiteLLMParams(TypedDict, total=False):
litellm_call_id: Optional[str]
model_alias_map: Optional[dict]
metadata: Optional[dict]
litellm_metadata: Optional[dict]
model_info: Optional[dict]
proxy_server_request: Optional[dict]
acompletion: Optional[bool]

View file

@ -4409,6 +4409,23 @@
"/v1/audio/transcriptions"
]
},
"azure/gpt-realtime-whisper": {
"input_cost_per_second": 0.0002833333333333333,
"litellm_provider": "azure",
"mode": "audio_transcription",
"source": "https://learn.microsoft.com/en-us/azure/ai-foundry/openai/concepts/gpt-realtime-whisper",
"supported_endpoints": [
"/v1/realtime",
"/v1/realtime/transcription_sessions"
],
"supported_modalities": [
"audio"
],
"supported_output_modalities": [
"text"
],
"supports_audio_input": true
},
"azure/gpt-5.1-2025-11-13": {
"cache_read_input_token_cost": 1.25e-07,
"cache_read_input_token_cost_priority": 2.5e-07,
@ -7557,6 +7574,45 @@
"supports_function_calling": true,
"supports_tool_choice": true
},
"azure_ai/deepseek-v3.1": {
"input_cost_per_token": 1.23e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.94e-06,
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/deepseek/",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"azure_ai/deepseek-v4-pro": {
"input_cost_per_token": 1.74e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 1000000,
"max_output_tokens": 384000,
"max_tokens": 384000,
"mode": "chat",
"output_cost_per_token": 3.48e-06,
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/deepseek/",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"azure_ai/deepseek-v4-flash": {
"input_cost_per_token": 1.9e-07,
"litellm_provider": "azure_ai",
"max_input_tokens": 1000000,
"max_output_tokens": 384000,
"max_tokens": 384000,
"mode": "chat",
"output_cost_per_token": 5.1e-07,
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/deepseek/",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"azure_ai/embed-v-4-0": {
"input_cost_per_token": 1.2e-07,
"litellm_provider": "azure_ai",
@ -40956,6 +41012,23 @@
"supports_system_messages": true,
"supports_tool_choice": true
},
"gpt-realtime-whisper": {
"input_cost_per_second": 0.0002833333333333333,
"litellm_provider": "openai",
"mode": "audio_transcription",
"source": "https://platform.openai.com/docs/models/gpt-realtime-whisper",
"supported_endpoints": [
"/v1/realtime",
"/v1/realtime/transcription_sessions"
],
"supported_modalities": [
"audio"
],
"supported_output_modalities": [
"text"
],
"supports_audio_input": true
},
"sora-2": {
"litellm_provider": "openai",
"mode": "video_generation",

View file

@ -170,48 +170,59 @@ def _make_wrapper(
)
def drive_sync(provider_key: str, chunks_per_stream: int, n_streams: int) -> float:
@dataclass
class TimingSample:
wall_s: float
cpu_s: float
def drive_sync(
provider_key: str, chunks_per_stream: int, n_streams: int
) -> TimingSample:
provider, factory = PROVIDERS[provider_key]
# Pre-build the chunk lists; we only measure wrapper iteration cost.
chunk_lists = [factory(chunks_per_stream) for _ in range(n_streams)]
gc.collect()
gc.disable()
try:
start = time.perf_counter()
wall_start = time.perf_counter()
cpu_start = time.process_time()
for chunks in chunk_lists:
wrapper = _make_wrapper(chunks, provider, async_stream=False)
for _ in wrapper:
pass
elapsed = time.perf_counter() - start
wall_elapsed = time.perf_counter() - wall_start
cpu_elapsed = time.process_time() - cpu_start
finally:
gc.enable()
return elapsed
return TimingSample(wall_s=wall_elapsed, cpu_s=cpu_elapsed)
async def drive_async(
provider_key: str, chunks_per_stream: int, n_streams: int
) -> float:
) -> TimingSample:
provider, factory = PROVIDERS[provider_key]
chunk_lists = [factory(chunks_per_stream) for _ in range(n_streams)]
gc.collect()
gc.disable()
try:
start = time.perf_counter()
wall_start = time.perf_counter()
cpu_start = time.process_time()
for chunks in chunk_lists:
wrapper = _make_wrapper(chunks, provider, async_stream=True)
async for _ in wrapper:
pass
elapsed = time.perf_counter() - start
wall_elapsed = time.perf_counter() - wall_start
cpu_elapsed = time.process_time() - cpu_start
finally:
gc.enable()
return elapsed
return TimingSample(wall_s=wall_elapsed, cpu_s=cpu_elapsed)
# ---------------------------------------------------------------------------
# Repeat × take-min runner
# ---------------------------------------------------------------------------
@dataclass
class Result:
label: str
@ -222,7 +233,11 @@ class Result:
total_chunks: int
elapsed_min_s: float
elapsed_median_s: float
cpu_at_min_wall_s: float
cpu_median_s: float
per_chunk_us: float
cpu_per_chunk_us: float
cpu_to_wall_ratio: float
chunks_per_sec: float
streams_per_sec: float
@ -260,11 +275,16 @@ def run_case(
else:
raise ValueError(f"unknown mode {mode!r}")
elapsed_min = min(samples)
elapsed_median = statistics.median(samples)
best_sample = min(samples, key=lambda s: s.wall_s)
elapsed_min = best_sample.wall_s
elapsed_median = statistics.median(s.wall_s for s in samples)
cpu_at_min_wall = best_sample.cpu_s
cpu_median = statistics.median(s.cpu_s for s in samples)
# Each stream emits chunks_per_stream text chunks + 1 finish/usage chunk.
total_chunks = n_streams * (chunks_per_stream + 1)
per_chunk_us = (elapsed_min * 1_000_000) / total_chunks
cpu_per_chunk_us = (cpu_at_min_wall * 1_000_000) / total_chunks
cpu_to_wall_ratio = cpu_at_min_wall / elapsed_min if elapsed_min > 0 else 0.0
chunks_per_sec = total_chunks / elapsed_min if elapsed_min > 0 else 0.0
streams_per_sec = n_streams / elapsed_min if elapsed_min > 0 else 0.0
@ -277,7 +297,11 @@ def run_case(
total_chunks=total_chunks,
elapsed_min_s=elapsed_min,
elapsed_median_s=elapsed_median,
cpu_at_min_wall_s=cpu_at_min_wall,
cpu_median_s=cpu_median,
per_chunk_us=per_chunk_us,
cpu_per_chunk_us=cpu_per_chunk_us,
cpu_to_wall_ratio=cpu_to_wall_ratio,
chunks_per_sec=chunks_per_sec,
streams_per_sec=streams_per_sec,
)
@ -289,6 +313,8 @@ def format_result(r: Result) -> str:
f"min={r.elapsed_min_s*1000:8.2f} ms "
f"median={r.elapsed_median_s*1000:8.2f} ms "
f"per-chunk={r.per_chunk_us:7.2f} μs "
f"cpu/chunk={r.cpu_per_chunk_us:7.2f} μs "
f"cpu/wall={r.cpu_to_wall_ratio:5.2f}x "
f"chunks/s={r.chunks_per_sec:>10,.0f} "
f"streams/s={r.streams_per_sec:>8,.1f}"
)

View file

@ -18,6 +18,12 @@ env_keys = set()
# Terminal/environment detection variables that should not be documented
# These are internal variables used for terminal detection, not user-configurable settings
# Guard-only env vars: read solely to raise on invalid values; the only valid
# value is the default, so there is nothing meaningful to document.
EXCLUDED_GUARD_ONLY_VARS = {
"MAVVRIK_FOCUS_FREQUENCY",
}
EXCLUDED_TERMINAL_VARS = {
"TERM",
"TERM_PROGRAM",
@ -64,6 +70,7 @@ for root, dirs, files in os.walk(repo_base):
match
for match in getenv_matches
if match not in EXCLUDED_TERMINAL_VARS
and match not in EXCLUDED_GUARD_ONLY_VARS
) # Extract only the key part, excluding terminal vars
# Find all keys using litellm.get_secret()

View file

@ -1160,7 +1160,10 @@ async def test_convert_to_bedrock_format_post_call_streaming_hook():
output_call = bedrock_calls[0]
assert output_call["source"] == "OUTPUT"
assert output_call["response"] is not None
assert output_call["messages"] is None # OUTPUT calls don't need messages
# OUTPUT forwards the request messages so contextual grounding can pull
# grounding_source/query blocks from them even on streamed responses. A
# plain-text (non-grounding) request still yields the single-block payload.
assert output_call["messages"] == request_data["messages"]
# Verify that the response content was masked
# The streaming chunks should now contain the masked content

View file

@ -30,6 +30,7 @@ locally-running proxy.
from __future__ import annotations
import json
import os
import socket
import subprocess
@ -41,7 +42,6 @@ from typing import Iterator
import httpx
import pytest
# ---------------------------------------------------------------------------
# Skip gate
# ---------------------------------------------------------------------------
@ -79,7 +79,9 @@ def _wait_for_health(base_url: str, proc: subprocess.Popen, deadline: float) ->
except httpx.HTTPError:
pass
time.sleep(0.5)
raise RuntimeError(f"litellm proxy did not become ready within {STARTUP_TIMEOUT_S}s")
raise RuntimeError(
f"litellm proxy did not become ready within {STARTUP_TIMEOUT_S}s"
)
def _oci_env_from_profile() -> dict[str, str]:
@ -106,38 +108,35 @@ def _oci_env_from_profile() -> dict[str, str]:
}
@pytest.fixture(scope="module")
def proxy_url() -> Iterator[str]:
oci_env = _oci_env_from_profile()
port = _free_port()
base_url = f"http://127.0.0.1:{port}"
def _serve(config_path: str) -> Iterator[str]:
"""Boot the litellm proxy with the given config and yield its base URL."""
env = os.environ.copy()
env.update(oci_env)
env.update(_oci_env_from_profile())
# Avoid pulling in DB-backed features for this lightweight smoke run.
env.pop("DATABASE_URL", None)
env["STORE_MODEL_IN_DB"] = "False"
port = _free_port()
base_url = f"http://127.0.0.1:{port}"
# Prefer the `litellm` console script that lives next to the active
# Python so we inherit the test virtualenv. Fall back to PATH.
cli = Path(sys.executable).parent / "litellm"
if not cli.exists():
cli = "litellm"
cmd = [
str(cli),
"--config",
str(CONFIG_PATH),
"--port",
str(port),
"--host",
"127.0.0.1",
"--num_workers",
"1",
]
proc = subprocess.Popen(
cmd,
[
str(cli),
"--config",
config_path,
"--port",
str(port),
"--host",
"127.0.0.1",
"--num_workers",
"1",
],
env=env,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
@ -155,6 +154,27 @@ def proxy_url() -> Iterator[str]:
proc.wait(timeout=5)
@pytest.fixture(scope="module")
def proxy_url() -> Iterator[str]:
yield from _serve(str(CONFIG_PATH))
@pytest.fixture(scope="module")
def proxy_url_no_drop_params(tmp_path_factory) -> Iterator[str]:
"""A proxy WITHOUT drop_params, to prove benign params the proxy injects
(e.g. max_retries) don't break OCI calls."""
cfg = tmp_path_factory.mktemp("oci_nodrop") / "config.yaml"
cfg.write_text(
"model_list:\n"
" - model_name: oci-cohere-command\n"
" litellm_params:\n"
" model: oci/cohere.command-latest\n"
"general_settings:\n"
f" master_key: {MASTER_KEY}\n"
)
yield from _serve(str(cfg))
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
@ -206,9 +226,7 @@ def test_chat_completion_via_proxy(proxy_url: str, model: str) -> None:
# Reasoning models may return empty content if their budget covers only
# the thinking turn — accept either text or a non-empty reasoning field.
has_content = bool(msg.get("content"))
has_reasoning = bool(msg.get("reasoning_content")) or bool(
msg.get("reasoning")
)
has_reasoning = bool(msg.get("reasoning_content")) or bool(msg.get("reasoning"))
assert has_content or has_reasoning, f"empty assistant message for {model}: {msg}"
usage = body.get("usage") or {}
assert usage.get("total_tokens", 0) > 0
@ -232,7 +250,7 @@ def test_chat_completion_streaming_via_proxy(proxy_url: str, model: str) -> None
continue
if not line.startswith("data:"):
continue
payload = line[len("data:"):].strip()
payload = line[len("data:") :].strip()
if payload == "[DONE]":
saw_done = True
break
@ -260,7 +278,6 @@ def test_embedding_via_proxy(proxy_url: str) -> None:
assert len(embedding) >= 64
assert all(isinstance(x, (int, float)) for x in embedding)
def test_model_list_advertises_oci_models(proxy_url: str) -> None:
"""The /v1/models registry advertises every OCI alias from the config."""
r = httpx.get(
@ -272,3 +289,124 @@ def test_model_list_advertises_oci_models(proxy_url: str) -> None:
advertised = {row["id"] for row in r.json()["data"]}
for expected in CHAT_MODELS + ["oci-embed"]:
assert expected in advertised, f"{expected} missing from /v1/models: {advertised}"
def test_chat_completion_no_drop_params(proxy_url_no_drop_params: str) -> None:
"""A plain chat completion succeeds through a proxy without drop_params.
Regression for the HTTP 500 ``param `max_retries` is not supported on OCI``:
the proxy injects max_retries on every request, so without this fix any OCI
call through the proxy failed unless drop_params was set.
"""
r = httpx.post(
f"{proxy_url_no_drop_params}/v1/chat/completions",
headers=_auth_headers(),
json=_chat_payload("oci-cohere-command"),
timeout=REQUEST_TIMEOUT_S,
)
assert r.status_code == 200, f"no-drop_params -> {r.status_code}: {r.text}"
body = r.json()
assert body["object"] == "chat.completion"
assert body["choices"][0]["message"].get("content") is not None
def test_cohere_default_n_via_proxy(proxy_url: str) -> None:
"""A Cohere request carrying the default n=1 succeeds through the gateway.
Regression for the HTTP 500 ``param `n` is not supported on OCI`` that
rejected every client which always sends n=1 (e.g. the MLflow gateway),
since OCI Cohere has no numGenerations field.
"""
payload = {**_chat_payload("oci-cohere-command"), "n": 1}
r = httpx.post(
f"{proxy_url}/v1/chat/completions",
headers=_auth_headers(),
json=payload,
timeout=REQUEST_TIMEOUT_S,
)
assert r.status_code == 200, f"n=1 -> {r.status_code}: {r.text}"
body = r.json()
assert body["object"] == "chat.completion"
assert body["choices"][0]["message"].get("content") is not None
@pytest.mark.parametrize("model", ["oci-cohere-command", "oci-llama"])
def test_response_format_json_schema_via_proxy(proxy_url: str, model: str) -> None:
"""A response_format json_schema succeeds through the gateway for both a
Cohere and a generic OCI model.
Regression for the HTTP 400 ``Please pass in correct format of request``
that rejected every json_schema request (which MLflow LLM judges always
send): generic models choke on OpenAI's ``strict`` key, and Cohere has no
JSON_SCHEMA type.
"""
r = httpx.post(
f"{proxy_url}/v1/chat/completions",
headers=_auth_headers(),
json={
"model": model,
"messages": [
{
"role": "user",
"content": "Rate the answer 4 to 2+2. Give an integer score and a short rationale.",
}
],
"max_tokens": 200,
"response_format": {
"type": "json_schema",
"json_schema": {
"name": "judgment",
"strict": True,
"schema": {
"type": "object",
"properties": {
"score": {"type": "integer"},
"rationale": {"type": "string"},
},
"required": ["score", "rationale"],
"additionalProperties": False,
},
},
},
},
timeout=REQUEST_TIMEOUT_S,
)
assert r.status_code == 200, f"{model} json_schema -> {r.status_code}: {r.text}"
content = r.json()["choices"][0]["message"]["content"]
assert content is not None
assert "score" in json.loads(content)
def test_omitted_max_tokens_not_truncated(proxy_url: str) -> None:
"""A request that omits max_tokens completes instead of being cut off.
Regression for OCI's tiny server-side maxTokens default (~20 tokens): without
an injected default, a request that doesn't set max_tokens came back with
finish_reason "length" after ~19 tokens, so structured outputs (e.g. MLflow
judge JSON) arrived as unterminated strings. The OCI provider now injects a
sane default when the caller omits one.
"""
r = httpx.post(
f"{proxy_url}/v1/chat/completions",
headers=_auth_headers(),
json={
"model": "oci-cohere-command",
"messages": [
{
"role": "user",
"content": "In four or five complete sentences, explain why the sky appears blue.",
}
],
},
timeout=REQUEST_TIMEOUT_S,
)
assert r.status_code == 200, f"omitted max_tokens -> {r.status_code}: {r.text}"
body = r.json()
choice = body["choices"][0]
assert (
choice["finish_reason"] != "length"
), f"response truncated by token cap: {choice}"
assert choice["finish_reason"] == "stop"
content = choice["message"].get("content") or ""
assert content.strip(), f"empty content: {choice}"
# The ~20-token server default truncated well before this; a complete
# four-to-five sentence answer comfortably exceeds it.
assert body["usage"]["completion_tokens"] > 50, body["usage"]

View file

@ -393,3 +393,38 @@ async def test_realtime_query_params_use_normalized_model_name(monkeypatch):
called_kwargs = mock_async_realtime.call_args.kwargs
assert called_kwargs["query_params"]["model"] == "gpt-4o-realtime-preview"
assert called_kwargs["query_params"]["intent"] == "chat"
@pytest.mark.asyncio
async def test_realtime_query_params_preserve_missing_model(monkeypatch):
"""
OpenAI-compatible transcription clients can connect with only
?intent=transcription and send the model in session.update. Do not add
model= back into the upstream query params when the client omitted it.
"""
from litellm.realtime_api import main as realtime_main
mock_async_realtime = AsyncMock()
monkeypatch.setattr(
realtime_main,
"openai_realtime",
MagicMock(async_realtime=mock_async_realtime),
)
def fake_get_llm_provider(model, api_base=None, api_key=None):
return ("gpt-realtime-whisper", "openai", None, None)
monkeypatch.setattr(realtime_main, "get_llm_provider", fake_get_llm_provider)
query_params: RealtimeQueryParams = {"intent": "transcription"}
await realtime_main._arealtime(
model="gpt-realtime-whisper",
websocket=MagicMock(),
api_key="sk-test",
query_params=query_params,
litellm_logging_obj=MagicMock(),
)
called_kwargs = mock_async_realtime.call_args.kwargs
assert called_kwargs["query_params"] == {"intent": "transcription"}

View file

@ -44,10 +44,14 @@ def test_realtime_query_params_template_caches_each_pair_separately():
params_with_intent_first = _realtime_query_params_template("gpt-4o", "intent-a")
params_with_intent_second = _realtime_query_params_template("gpt-4o", "intent-a")
params_without_intent = _realtime_query_params_template("gpt-4o", None)
params_transcription_without_model = _realtime_query_params_template(
None, "transcription"
)
assert params_with_intent_first is params_with_intent_second
assert params_with_intent_first == (("model", "gpt-4o"), ("intent", "intent-a"))
assert params_without_intent == (("model", "gpt-4o"),)
assert params_transcription_without_model == (("intent", "transcription"),)
assert params_with_intent_first is not params_without_intent

View file

@ -125,6 +125,33 @@ def test_completion_skips_rewrapping_preformatted_cached_chat_stream():
assert result is stream
def test_completion_preserves_top_level_stream_flag_in_responses_request():
stream = MagicMock(spec=CustomStreamWrapper)
stream.custom_llm_provider = "cached_response"
bridge = ResponsesToCompletionBridgeHandler()
kwargs = _bridge_kwargs(stream=False)
kwargs["stream"] = True
kwargs["optional_params"].pop("stream")
with (
patch.object(
bridge.transformation_handler,
"transform_request",
return_value={"model": "gpt-5.4", "input": "hi"},
) as transform_request,
patch("litellm.responses", return_value=stream),
patch.object(
bridge,
"_apply_post_stream_processing",
side_effect=lambda s, *a, **kw: s,
),
):
result = bridge.completion(**kwargs)
assert result is stream
assert transform_request.call_args.kwargs["optional_params"]["stream"] is True
@pytest.mark.asyncio
async def test_acompletion_skips_rewrapping_preformatted_cached_chat_stream():
stream = MagicMock(spec=CustomStreamWrapper)
@ -148,3 +175,31 @@ async def test_acompletion_skips_rewrapping_preformatted_cached_chat_stream():
post.assert_called_once()
assert result is stream
@pytest.mark.asyncio
async def test_acompletion_preserves_top_level_stream_flag_in_responses_request():
stream = MagicMock(spec=CustomStreamWrapper)
stream.custom_llm_provider = "cached_response"
bridge = ResponsesToCompletionBridgeHandler()
kwargs = _bridge_kwargs(stream=False)
kwargs["stream"] = True
kwargs["optional_params"].pop("stream")
with (
patch.object(
bridge.transformation_handler,
"transform_request",
return_value={"model": "gpt-5.4", "input": "hi"},
) as transform_request,
patch("litellm.aresponses", new=AsyncMock(return_value=stream)),
patch.object(
bridge,
"_apply_post_stream_processing",
side_effect=lambda s, *a, **kw: s,
),
):
result = await bridge.acompletion(**kwargs)
assert result is stream
assert transform_request.call_args.kwargs["optional_params"]["stream"] is True

View file

@ -28,6 +28,31 @@ class TestSoftBudgetAlert:
result = alert.get_id(user_info)
assert result == "default_id"
def test_get_id_returns_team_id_for_team_event_group(self):
"""Team soft budget alerts dedupe by team, not by the calling key's token"""
alert = SoftBudgetAlert()
user_info = CallInfo(
spend=120.0,
token="test_token_123",
team_id="team_456",
event_group=Litellm_EntityType.TEAM,
)
result = alert.get_id(user_info)
assert result == "team_456"
def test_get_id_returns_default_id_for_team_event_group_without_team_id(self):
alert = SoftBudgetAlert()
user_info = CallInfo(
spend=120.0,
token="test_token_123",
team_id=None,
event_group=Litellm_EntityType.TEAM,
)
result = alert.get_id(user_info)
assert result == "default_id"
def test_get_id_with_empty_token(self):
"""Test that get_id returns 'default_id' when token is empty string"""
alert = SoftBudgetAlert()

View file

@ -0,0 +1,751 @@
"""Tests for FocusMavvrikDestination."""
from __future__ import annotations
from datetime import datetime, timezone
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.integrations.focus.destinations.base import FocusTimeWindow
from litellm.integrations.focus.destinations.mavvrik_destination import (
FocusMavvrikDestination,
_validate_api_endpoint,
)
VALID_ENDPOINT = "https://api.mavvrik.ai/tenant123"
def _make_window() -> FocusTimeWindow:
return FocusTimeWindow(
start_time=datetime(2026, 1, 1, 0, 0, 0, tzinfo=timezone.utc),
end_time=datetime(2026, 1, 2, 0, 0, 0, tzinfo=timezone.utc),
frequency="daily",
)
def _dest(**overrides) -> FocusMavvrikDestination:
config = {
"api_key": "test-key",
"api_endpoint": VALID_ENDPOINT,
"connection_id": "conn-123",
}
config.update(overrides)
return FocusMavvrikDestination(prefix="mavvrik_focus_exports", config=config)
def test_missing_api_key_raises():
with pytest.raises(ValueError, match="MAVVRIK_API_KEY"):
FocusMavvrikDestination(
prefix="p",
config={"api_endpoint": VALID_ENDPOINT, "connection_id": "c"},
)
def test_missing_api_endpoint_raises():
with pytest.raises(ValueError, match="MAVVRIK_API_ENDPOINT"):
FocusMavvrikDestination(
prefix="p",
config={"api_key": "k", "connection_id": "c"},
)
def test_missing_connection_id_raises():
with pytest.raises(ValueError, match="MAVVRIK_CONNECTION_ID"):
FocusMavvrikDestination(
prefix="p",
config={"api_key": "k", "api_endpoint": VALID_ENDPOINT},
)
def test_non_https_endpoint_raises():
with pytest.raises(ValueError, match="HTTPS"):
_validate_api_endpoint("http://api.mavvrik.ai/tenant")
def test_non_mavvrik_domain_raises():
with pytest.raises(ValueError, match="Mavvrik domain"):
_validate_api_endpoint("https://evil.com/tenant")
def test_valid_mavvrik_domains_accepted():
for domain in (
"https://api.mavvrik.ai/tenant",
"https://api.mavvrik.dev/tenant",
"https://api.mavvrik.app/tenant",
):
_validate_api_endpoint(domain) # must not raise
def test_initializes_with_not_registered():
dest = _dest()
assert dest._registered is False
@pytest.mark.asyncio
async def test_deliver_skips_empty_content():
dest = _dest()
await dest.deliver(content=b"", time_window=_make_window(), filename="usage.csv")
# _registered still False — _ensure_registered was never called
assert dest._registered is False
@pytest.mark.asyncio
async def test_large_content_uploads_in_multiple_chunks():
"""Content larger than _GCS_CHUNK_SIZE must be uploaded in multiple chunks.
GCS assembles intermediate chunks (308) + final chunk (200) into one object.
The destination must send Content-Range headers for each chunk correctly.
"""
from litellm.integrations.focus.destinations.mavvrik_destination import (
FocusMavvrikDestination,
_GCS_CHUNK_SIZE,
)
dest = FocusMavvrikDestination(
prefix="p",
config={"api_key": "k", "api_endpoint": VALID_ENDPOINT, "connection_id": "c"},
)
register_resp = MagicMock()
register_resp.status_code = 200
signed_url_resp = MagicMock()
signed_url_resp.status_code = 200
signed_url_resp.json.return_value = {
"url": "https://storage.googleapis.com/upload?sig=x"
}
init_resp = MagicMock()
init_resp.status_code = 200
init_resp.headers = {"Location": "https://storage.googleapis.com/session"}
# First chunk → 308, second (final) chunk → 200
chunk1_resp = MagicMock()
chunk1_resp.status_code = 308
chunk2_resp = MagicMock()
chunk2_resp.status_code = 200
mock_http = MagicMock()
mock_http.client = MagicMock()
mock_http.client.request = AsyncMock(
side_effect=[
register_resp,
signed_url_resp,
init_resp,
chunk1_resp,
chunk2_resp,
]
)
dest._http = mock_http
# Build content that when gzipped exceeds one chunk.
# Use incompressible random-ish bytes to ensure gzip doesn't shrink it below the chunk size.
import os as _os
raw = b"col1,col2\n" + _os.urandom(_GCS_CHUNK_SIZE + 1024)
await dest.deliver(
content=raw,
time_window=_make_window(),
filename="usage.csv",
)
# register + get_signed_url + init + 2 chunk PUTs = 5 calls
assert mock_http.client.request.call_count == 5
# Check Content-Range headers
put_calls = mock_http.client.request.call_args_list[3:]
assert "bytes" in put_calls[0].kwargs["headers"]["Content-Range"]
assert "/*" in put_calls[0].kwargs["headers"]["Content-Range"] # intermediate
assert "/*" not in put_calls[1].kwargs["headers"]["Content-Range"] # final
@pytest.mark.asyncio
async def test_deliver_calls_register_get_url_and_upload():
dest = _dest()
register_resp = MagicMock()
register_resp.status_code = 200
signed_url_resp = MagicMock()
signed_url_resp.status_code = 200
signed_url_resp.json.return_value = {"url": "https://storage.googleapis.com/signed"}
init_resp = MagicMock()
init_resp.status_code = 200
init_resp.headers = {"Location": "https://storage.googleapis.com/session-uri"}
upload_resp = MagicMock()
upload_resp.status_code = 200
mock_http = MagicMock()
mock_http.client = MagicMock()
# All 4 calls go through self._http.client.request:
# 1. register, 2. get_signed_url, 3. GCS session init POST, 4. GCS PUT
mock_http.client.request = AsyncMock(
side_effect=[register_resp, signed_url_resp, init_resp, upload_resp]
)
dest._http = mock_http
await dest.deliver(
content=b"header\nrow1\n",
time_window=_make_window(),
filename="usage.csv",
)
assert dest._registered is True
assert mock_http.client.request.call_count == 4
# Verify Content-Range header was set on the PUT
put_call = mock_http.client.request.call_args_list[3]
assert "Content-Range" in put_call.kwargs["headers"]
@pytest.mark.asyncio
async def test_register_called_only_once_across_multiple_deliveries():
dest = _dest()
register_resp = MagicMock()
register_resp.status_code = 200
def _signed_url_resp():
r = MagicMock()
r.status_code = 200
r.json.return_value = {"url": "https://storage.googleapis.com/signed"}
return r
init_resp = MagicMock()
init_resp.status_code = 200
init_resp.headers = {"Location": "https://storage.googleapis.com/session-uri"}
upload_resp = MagicMock()
upload_resp.status_code = 200
mock_http = MagicMock()
mock_http.client = MagicMock()
# First delivery: register, get_signed_url, GCS init, GCS PUT
# Second delivery: get_signed_url, GCS init, GCS PUT (register skipped)
mock_http.client.request = AsyncMock(
side_effect=[
register_resp,
_signed_url_resp(),
init_resp,
upload_resp,
_signed_url_resp(),
init_resp,
upload_resp,
]
)
dest._http = mock_http
window = _make_window()
await dest.deliver(content=b"header\nrow1\n", time_window=window, filename="1.csv")
await dest.deliver(content=b"header\nrow2\n", time_window=window, filename="2.csv")
# 7 total: register(1) + [get_url+init+put](2) × 2 deliveries
assert mock_http.client.request.call_count == 7
# First call was register
first_call = mock_http.client.request.call_args_list[0]
assert first_call.kwargs["method"] == "POST"
assert "/upload-url" not in first_call.kwargs["url"]
@pytest.mark.asyncio
async def test_deliver_raises_on_register_failure():
dest = _dest()
fail_resp = MagicMock()
fail_resp.status_code = 403
fail_resp.text = "Forbidden"
mock_http = MagicMock()
mock_http.client = MagicMock()
mock_http.client.request = AsyncMock(return_value=fail_resp)
dest._http = mock_http
with pytest.raises(RuntimeError, match="register failed"):
await dest.deliver(
content=b"data",
time_window=_make_window(),
filename="usage.csv",
)
@pytest.mark.asyncio
async def test_deliver_raises_on_signed_url_api_error():
"""_get_signed_url must raise RuntimeError when the API returns a 4xx."""
dest = _dest()
register_resp = MagicMock()
register_resp.status_code = 200
fail_resp = MagicMock()
fail_resp.status_code = 500
fail_resp.text = "Internal Server Error"
mock_http = MagicMock()
mock_http.client = MagicMock()
mock_http.client.request = AsyncMock(side_effect=[register_resp, fail_resp])
dest._http = mock_http
with pytest.raises(RuntimeError, match="failed to get signed URL"):
await dest.deliver(
content=b"data",
time_window=_make_window(),
filename="usage.csv",
)
@pytest.mark.asyncio
async def test_deliver_raises_on_missing_signed_url():
dest = _dest()
register_resp = MagicMock()
register_resp.status_code = 200
bad_url_resp = MagicMock()
bad_url_resp.status_code = 200
bad_url_resp.json.return_value = {} # no 'url' field
mock_http = MagicMock()
mock_http.client = MagicMock()
mock_http.client.request = AsyncMock(side_effect=[register_resp, bad_url_resp])
dest._http = mock_http
with pytest.raises(RuntimeError, match="missing 'url' field"):
await dest.deliver(
content=b"data",
time_window=_make_window(),
filename="usage.csv",
)
@pytest.mark.asyncio
async def test_deliver_raises_on_non_gcs_signed_url():
"""Signed URL pointing to a non-GCS host must be rejected before any upload."""
from litellm.integrations.focus.destinations.mavvrik_destination import (
_validate_gcs_url,
)
with pytest.raises(ValueError, match="GCS endpoint"):
_validate_gcs_url("https://evil.com/upload?token=abc", "signed URL")
@pytest.mark.asyncio
async def test_deliver_raises_on_non_gcs_session_uri():
"""Session URI from Location header pointing to a non-GCS host must be rejected."""
dest = _dest()
register_resp = MagicMock()
register_resp.status_code = 200
signed_url_resp = MagicMock()
signed_url_resp.status_code = 200
# signed URL is valid GCS
signed_url_resp.json.return_value = {
"url": "https://storage.googleapis.com/upload?sig=abc"
}
# Location header points to a non-GCS host
init_resp = MagicMock()
init_resp.status_code = 200
init_resp.headers = {"Location": "https://evil.com/session-uri"}
mock_http = MagicMock()
mock_http.client = MagicMock()
# register, get_signed_url, GCS session init (returns bad Location)
mock_http.client.request = AsyncMock(
side_effect=[register_resp, signed_url_resp, init_resp]
)
dest._http = mock_http
with pytest.raises(ValueError, match="GCS endpoint"):
await dest.deliver(
content=b"data",
time_window=_make_window(),
filename="usage.csv",
)
def test_factory_creates_mavvrik_destination(monkeypatch):
monkeypatch.setenv("MAVVRIK_API_KEY", "k")
monkeypatch.setenv("MAVVRIK_API_ENDPOINT", VALID_ENDPOINT)
monkeypatch.setenv("MAVVRIK_CONNECTION_ID", "c")
from litellm.integrations.focus.destinations.factory import FocusDestinationFactory
dest = FocusDestinationFactory.create(provider="mavvrik", prefix="p")
assert isinstance(dest, FocusMavvrikDestination)
assert dest.api_key == "k"
assert dest.connection_id == "c"
def test_only_daily_frequency_is_supported():
"""MavvrikFocusLogger must raise ValueError for non-daily frequencies."""
import importlib
for freq in ("hourly", "interval"):
def _make(f=freq, monkeypatch=None):
import os
old = os.environ.get("MAVVRIK_FOCUS_FREQUENCY")
os.environ["MAVVRIK_FOCUS_FREQUENCY"] = f
try:
from litellm.integrations.mavvrik_focus import mavvrik_focus_logger
importlib.reload(mavvrik_focus_logger)
with pytest.raises(ValueError, match="Only 'daily' is allowed"):
mavvrik_focus_logger.MavvrikFocusLogger()
finally:
if old is None:
os.environ.pop("MAVVRIK_FOCUS_FREQUENCY", None)
else:
os.environ["MAVVRIK_FOCUS_FREQUENCY"] = old
_make()
def test_max_rows_defaults_to_500k():
"""MAVVRIK_FOCUS_MAX_ROWS defaults to 500_000 when not set."""
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import (
MavvrikFocusLogger,
)
logger = MavvrikFocusLogger()
assert logger._max_rows == 500_000
def test_max_rows_reads_from_env(monkeypatch):
"""MAVVRIK_FOCUS_MAX_ROWS env var is respected."""
monkeypatch.setenv("MAVVRIK_FOCUS_MAX_ROWS", "100000")
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import (
MavvrikFocusLogger,
)
logger = MavvrikFocusLogger()
assert logger._max_rows == 100_000
@pytest.mark.asyncio
async def test_export_window_passes_max_rows_as_limit(monkeypatch):
"""_export_window must pass _max_rows as limit to get_usage_data."""
monkeypatch.setenv("MAVVRIK_FOCUS_MAX_ROWS", "1000")
import polars as pl
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import (
MavvrikFocusLogger,
)
from litellm.integrations.focus.destinations.base import FocusTimeWindow
from datetime import datetime, timezone
logger = MavvrikFocusLogger()
assert logger._max_rows == 1000
# Mock the engine internals so _export_window runs through our new code path
db_mock = MagicMock()
db_mock.get_usage_data = AsyncMock(return_value=pl.DataFrame()) # empty → no upload
engine_mock = MagicMock()
engine_mock._database = db_mock
logger._engine = engine_mock
window = FocusTimeWindow(
start_time=datetime(2026, 1, 1, tzinfo=timezone.utc),
end_time=datetime(2026, 1, 2, tzinfo=timezone.utc),
frequency="daily",
)
await logger._export_window(window=window, limit=None)
db_mock.get_usage_data.assert_called_once_with(
limit=1000,
start_time_utc=window.start_time,
end_time_utc=window.end_time,
)
@pytest.mark.asyncio
async def test_run_scheduled_export_catches_up_missed_dates():
"""If metricsMarker is 2 days behind, _run_scheduled_export exports missed dates first."""
import polars as pl
from datetime import datetime, timedelta, timezone
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import (
MavvrikFocusLogger,
)
from litellm.integrations.focus.destinations.mavvrik_destination import (
FocusMavvrikDestination,
)
logger = MavvrikFocusLogger()
# metricsMarker = 3 days ago → 2 missed dates (day-2 and day-1) + today's run
now = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0)
yesterday = now - timedelta(days=1)
two_days_ago = now - timedelta(days=2)
three_days_ago = now - timedelta(days=3)
marker_ts = int(three_days_ago.timestamp())
# Mock destination
dest_mock = MagicMock(spec=FocusMavvrikDestination)
dest_mock.get_metrics_marker = AsyncMock(return_value=marker_ts)
# Mock engine
db_mock = MagicMock()
db_mock.get_usage_data = AsyncMock(return_value=pl.DataFrame())
engine_mock = MagicMock()
engine_mock._database = db_mock
engine_mock._destination = dest_mock
logger._engine = engine_mock
await logger._run_scheduled_export()
# Should have queried DB 3 times: day-2, day-1 (yesterday), and the normal yesterday window
# Actually: catch-up covers [three_days_ago+1 .. yesterday) = [two_days_ago, yesterday)
# = two_days_ago only (1 missed), then normal yesterday = 2 total calls
calls = db_mock.get_usage_data.call_args_list
assert len(calls) == 2
# First call is the catch-up (two_days_ago)
assert calls[0].kwargs["start_time_utc"].date() == two_days_ago.date()
# Second call is yesterday's normal daily run
assert calls[1].kwargs["start_time_utc"].date() == yesterday.date()
@pytest.mark.asyncio
async def test_run_scheduled_export_no_catchup_when_marker_is_current():
"""If metricsMarker = yesterday, no catch-up needed — just export yesterday."""
import polars as pl
from datetime import datetime, timedelta, timezone
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import (
MavvrikFocusLogger,
)
from litellm.integrations.focus.destinations.mavvrik_destination import (
FocusMavvrikDestination,
)
logger = MavvrikFocusLogger()
now = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0)
yesterday = now - timedelta(days=1)
marker_ts = int(yesterday.timestamp())
dest_mock = MagicMock(spec=FocusMavvrikDestination)
dest_mock.get_metrics_marker = AsyncMock(return_value=marker_ts)
db_mock = MagicMock()
db_mock.get_usage_data = AsyncMock(return_value=pl.DataFrame())
engine_mock = MagicMock()
engine_mock._database = db_mock
engine_mock._destination = dest_mock
logger._engine = engine_mock
await logger._run_scheduled_export()
# Only one call — yesterday's normal run, no catch-up
assert db_mock.get_usage_data.call_count == 1
assert (
db_mock.get_usage_data.call_args.kwargs["start_time_utc"].date()
== yesterday.date()
)
@pytest.mark.asyncio
async def test_metrics_marker_always_calls_api():
"""get_metrics_marker must call the register API every time to get a fresh marker.
This is the key difference from deliver() — catch-up requires the current
metricsMarker on every scheduled run, not just the first one.
"""
dest = _dest()
register_resp = MagicMock()
register_resp.status_code = 200
register_resp.json.return_value = {
"id": "litellm-conn-123",
"metricsMarker": 1749340800,
}
mock_http = MagicMock()
mock_http.client = MagicMock()
mock_http.client.request = AsyncMock(return_value=register_resp)
dest._http = mock_http
# First call
marker = await dest.get_metrics_marker()
assert marker == 1749340800
assert dest._registered is True
# Second call — must call API again to get fresh marker (not return None)
marker2 = await dest.get_metrics_marker()
assert marker2 == 1749340800
assert mock_http.client.request.call_count == 2 # API called both times
def test_parse_metrics_marker_handles_unix_timestamp():
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import (
_parse_metrics_marker,
)
from datetime import datetime, timezone
# Use a known date and compute its timestamp to avoid hardcoding
known_date = datetime(2026, 6, 9, 0, 0, 0, tzinfo=timezone.utc)
ts = int(known_date.timestamp())
result = _parse_metrics_marker(ts)
assert result is not None
assert result.date().isoformat() == "2026-06-09"
assert result.tzinfo == timezone.utc
def test_parse_metrics_marker_handles_iso_date_string():
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import (
_parse_metrics_marker,
)
result = _parse_metrics_marker("2026-06-09")
assert result is not None
assert result.date().isoformat() == "2026-06-09"
def test_parse_metrics_marker_handles_iso_datetime_string():
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import (
_parse_metrics_marker,
)
result = _parse_metrics_marker("2026-06-09T00:00:00Z")
assert result is not None
assert result.date().isoformat() == "2026-06-09"
def test_parse_metrics_marker_returns_none_for_zero():
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import (
_parse_metrics_marker,
)
assert _parse_metrics_marker(0) is None
assert _parse_metrics_marker(None) is None
assert _parse_metrics_marker("") is None
def test_parse_metrics_marker_returns_none_for_garbage():
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import (
_parse_metrics_marker,
)
# Should not raise — logs warning and returns None
assert _parse_metrics_marker("not-a-date") is None
@pytest.mark.asyncio
async def test_catchup_capped_at_max_catchup_days():
"""Catch-up must not go further back than _MAX_CATCHUP_DAYS."""
import polars as pl
from datetime import datetime, timedelta, timezone
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import (
MavvrikFocusLogger,
)
from litellm.integrations.focus.destinations.mavvrik_destination import (
FocusMavvrikDestination,
)
logger = MavvrikFocusLogger()
max_days = MavvrikFocusLogger._MAX_CATCHUP_DAYS
now = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0)
yesterday = now - timedelta(days=1)
# Marker is 30 days ago — well beyond the cap
thirty_days_ago = now - timedelta(days=30)
marker_ts = int(thirty_days_ago.timestamp())
dest_mock = MagicMock(spec=FocusMavvrikDestination)
dest_mock.get_metrics_marker = AsyncMock(return_value=marker_ts)
db_mock = MagicMock()
db_mock.get_usage_data = AsyncMock(return_value=pl.DataFrame())
engine_mock = MagicMock()
engine_mock._database = db_mock
engine_mock._destination = dest_mock
logger._engine = engine_mock
await logger._run_scheduled_export()
# Should have queried at most _MAX_CATCHUP_DAYS times
# (max_days - 1 catch-up dates + 1 yesterday = max_days total)
assert db_mock.get_usage_data.call_count <= max_days
# First catch-up date must not be earlier than (yesterday - max_days + 1)
earliest_allowed = yesterday - timedelta(days=max_days - 1)
first_call_start = db_mock.get_usage_data.call_args_list[0].kwargs["start_time_utc"]
assert first_call_start.date() >= earliest_allowed.date()
@pytest.mark.asyncio
async def test_register_resets_on_410():
"""_registered flag must be False after a 410 so next run re-registers."""
dest = _dest()
resp_410 = MagicMock()
resp_410.status_code = 410
resp_410.text = "Gone"
mock_http = MagicMock()
mock_http.client = MagicMock()
mock_http.client.request = AsyncMock(return_value=resp_410)
dest._http = mock_http
dest._registered = False # not yet registered — trigger the call
with pytest.raises(RuntimeError, match="disconnected"):
await dest._ensure_registered()
assert dest._registered is False
@pytest.mark.asyncio
async def test_gcs_session_cancelled_on_chunk_failure():
"""GCS session must be cancelled (DELETE) when a chunk PUT fails."""
dest = _dest()
register_resp = MagicMock()
register_resp.status_code = 200
signed_url_resp = MagicMock()
signed_url_resp.status_code = 200
signed_url_resp.json.return_value = {
"url": "https://storage.googleapis.com/upload?sig=x"
}
init_resp = MagicMock()
init_resp.status_code = 200
init_resp.headers = {"Location": "https://storage.googleapis.com/session"}
# Chunk PUT fails with 500
fail_resp = MagicMock()
fail_resp.status_code = 500
fail_resp.text = "Internal Server Error"
# DELETE (session cancel)
delete_resp = MagicMock()
delete_resp.status_code = 200
mock_http = MagicMock()
mock_http.client = MagicMock()
mock_http.client.request = AsyncMock(
side_effect=[register_resp, signed_url_resp, init_resp, fail_resp, delete_resp]
)
dest._http = mock_http
with pytest.raises(RuntimeError, match="GCS chunk upload failed"):
await dest.deliver(
content=b"header\nrow1\n",
time_window=_make_window(),
filename="usage.csv",
)
# Verify DELETE was called to cancel the session
calls = mock_http.client.request.call_args_list
delete_call = calls[4]
assert delete_call.kwargs["method"] == "DELETE"
assert "storage.googleapis.com/session" in delete_call.kwargs["url"]

File diff suppressed because it is too large Load diff

View file

@ -34,6 +34,23 @@ class TestStandardizedResetTime(unittest.TestCase):
custom_day_result = get_next_standardized_reset_time("3d", base_time, "UTC")
self.assertEqual(custom_day_result, custom_day_expected)
def test_week_based_resets(self):
"""Test week-based reset durations (1w, 2w).
1w snaps to the next Monday at midnight (same as 7d).
2w advances exactly 14 days from the current date at midnight.
"""
# 1w from a Wednesday -> next Monday (5 days away, not 7)
wednesday = datetime(2023, 5, 17, 15, 45, 0, tzinfo=timezone.utc)
weekly_expected = datetime(2023, 5, 22, 0, 0, 0, tzinfo=timezone.utc)
weekly_result = get_next_standardized_reset_time("1w", wednesday, "UTC")
self.assertEqual(weekly_result, weekly_expected)
# 2w from a Wednesday -> exactly 14 days out (lands on a Wednesday, not Monday)
base_time = datetime(2023, 5, 17, 10, 30, 0, tzinfo=timezone.utc)
two_week_expected = datetime(2023, 5, 31, 0, 0, 0, tzinfo=timezone.utc)
two_week_result = get_next_standardized_reset_time("2w", base_time, "UTC")
self.assertEqual(two_week_result, two_week_expected)
def test_hour_minute_second_resets(self):
"""Test hour, minute, and second based reset durations"""
# Base time: 2023-05-15 15:20:30 UTC (3:20:30 PM)

View file

@ -3151,6 +3151,41 @@ def test_get_error_information_prefers_message_attribute_over_empty_str():
assert info["error_code"] == "401"
def _anthropic_messages_logging_obj():
return LitellmLogging(
model="openai/my-local",
messages=[{"role": "user", "content": "hi"}],
stream=True,
call_type="anthropic_messages",
start_time=time.time(),
litellm_call_id="28595",
function_id="28595",
)
def _responses_api_response_with_text(text="hello world"):
from openai.types.responses import ResponseOutputMessage, ResponseOutputText
from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse
return ResponsesAPIResponse(
id="resp-28595",
created_at=1700000000,
output=[
ResponseOutputMessage(
id="msg-1",
type="message",
role="assistant",
status="completed",
content=[
ResponseOutputText(annotations=[], text=text, type="output_text")
],
)
],
usage=ResponseAPIUsage(input_tokens=11, output_tokens=7, total_tokens=18),
)
@pytest.mark.parametrize(
"event_cls, event_type",
[
@ -3159,34 +3194,68 @@ def test_get_error_information_prefers_message_attribute_over_empty_str():
("ResponseFailedEvent", "response.failed"),
],
)
def test_handle_anthropic_messages_response_logging_with_terminal_responses_api_events(
def test_handle_anthropic_messages_response_logging_translates_terminal_responses_api_event(
event_cls, event_type
):
"""Regression test for #28943: when anthropic_messages routes to OpenAI Responses
API and stream=True, success_handler receives a terminal ResponsesAPI event instead
of a ModelResponse. The handler must return the inner ResponsesAPIResponse rather
than crashing with AnthropicResponse.model_validate."""
"""Regression for #28595 / #28943. When anthropic_messages routes to the OpenAI
Responses backend and stream=True, success_handler receives a terminal Responses
API event. The handler must translate it to a ModelResponse whose choices carry
the assistant text, so the proxy UI Logs tab (which reads response.choices[0])
renders the response content instead of "No response data available"."""
import importlib
openai_types = importlib.import_module("litellm.types.llms.openai")
EventClass = getattr(openai_types, event_cls)
from litellm.types.llms.openai import ResponsesAPIResponse
logging_obj = LitellmLogging(
model="gpt-4o",
messages=[{"role": "user", "content": "hello"}],
stream=True,
call_type="anthropic_messages",
start_time=time.time(),
litellm_call_id="test-rce-123",
function_id="test-fn",
)
inner_response = ResponsesAPIResponse(
id="resp_test", created_at=1700000000, output=[]
)
logging_obj = _anthropic_messages_logging_obj()
inner_response = _responses_api_response_with_text("hello world")
event = EventClass(type=event_type, response=inner_response)
result = logging_obj._handle_anthropic_messages_response_logging(result=event)
assert result is inner_response
assert isinstance(result, ModelResponse)
assert result.choices[0].message.content == "hello world" # type: ignore[union-attr]
assert result.usage.prompt_tokens == 11 # type: ignore[attr-defined]
assert result.usage.completion_tokens == 7 # type: ignore[attr-defined]
def test_handle_anthropic_messages_response_logging_translates_bare_responses_api_response():
"""Non-streaming bridge path: result is a bare ResponsesAPIResponse (no event wrap)."""
logging_obj = _anthropic_messages_logging_obj()
result = logging_obj._handle_anthropic_messages_response_logging(
result=_responses_api_response_with_text("hi there")
)
assert isinstance(result, ModelResponse)
assert result.choices[0].message.content == "hi there" # type: ignore[union-attr]
assert result.usage.total_tokens == 18 # type: ignore[attr-defined]
def test_handle_anthropic_messages_response_logging_passes_model_response_through():
"""Anthropic-native path already yields a ModelResponse; it must be returned unchanged."""
logging_obj = _anthropic_messages_logging_obj()
model_response = ModelResponse()
assert (
logging_obj._handle_anthropic_messages_response_logging(result=model_response)
is model_response
)
def test_handle_anthropic_messages_response_logging_degrades_on_unparseable_responses_payload():
"""If the Responses translation raises (eg. empty output on an incomplete response),
the row must still land: a minimal ModelResponse with model + usage is returned."""
from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse
logging_obj = _anthropic_messages_logging_obj()
empty = ResponsesAPIResponse(
id="resp-empty",
created_at=1700000000,
output=[],
usage=ResponseAPIUsage(input_tokens=4, output_tokens=0, total_tokens=4),
)
result = logging_obj._handle_anthropic_messages_response_logging(result=empty)
assert isinstance(result, ModelResponse)
assert result.model == "openai/my-local"
assert result.usage.prompt_tokens == 4 # type: ignore[attr-defined]

View file

@ -521,6 +521,290 @@ async def test_transcription_captured_in_backend_to_client():
assert logging_obj.model_call_details["messages"] == streaming.input_messages
@pytest.mark.asyncio
async def test_transcription_session_captures_usage_and_skips_response_create():
"""
For a transcription-only session (session.type == "transcription", e.g.
gpt-realtime-whisper), the completed event's audio-duration usage must be
captured for cost and response.create must NOT be sent to the backend.
"""
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
session_created = json.dumps(
{
"type": "session.created",
"session": {
"type": "transcription",
"audio": {
"input": {"transcription": {"model": "gpt-realtime-whisper"}}
},
},
}
).encode()
completed = json.dumps(
{
"type": "conversation.item.input_audio_transcription.completed",
"transcript": "hello world",
"item_id": "item_1",
"usage": {"type": "duration", "seconds": 12.0},
}
).encode()
backend_ws = MagicMock()
backend_ws.recv = AsyncMock(
side_effect=[session_created, completed, ConnectionClosed(None, None)]
)
backend_ws.send = AsyncMock()
logging_obj = MagicMock()
logging_obj.model_call_details = {}
logging_obj.async_success_handler = AsyncMock()
logging_obj.success_handler = MagicMock()
streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj)
await streaming.backend_to_client_send_messages()
assert streaming._is_transcription_session is True
captured = [
m
for m in streaming.messages
if m.get("type") == "conversation.item.input_audio_transcription.completed"
]
assert len(captured) == 1, "completed usage event must be captured for cost"
assert captured[0]["usage"]["seconds"] == 12.0
# Transcript still forwarded to the client.
client_ws.send_text.assert_any_call(completed.decode())
# No response.create — transcription sessions have no assistant turn.
sent_to_backend = [
json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args
]
assert all(
e.get("type") != "response.create" for e in sent_to_backend
), f"transcription session must not trigger response.create, got: {sent_to_backend}"
@pytest.mark.asyncio
async def test_non_transcription_completed_event_still_triggers_response_create():
"""
Regression guard: a normal (non-transcription) session with no guardrails must
keep triggering response.create on a completed transcription event.
"""
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
completed = json.dumps(
{
"type": "conversation.item.input_audio_transcription.completed",
"transcript": "hi",
"item_id": "item_1",
}
).encode()
backend_ws = MagicMock()
backend_ws.recv = AsyncMock(side_effect=[completed, ConnectionClosed(None, None)])
backend_ws.send = AsyncMock()
logging_obj = MagicMock()
logging_obj.async_success_handler = AsyncMock()
logging_obj.success_handler = MagicMock()
streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj)
await streaming.backend_to_client_send_messages()
assert streaming._is_transcription_session is False
sent_to_backend = [
json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args
]
assert any(e.get("type") == "response.create" for e in sent_to_backend)
def test_client_session_update_marks_transcription_session():
"""A client session.update with type=transcription flags the session."""
streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock())
assert streaming._is_transcription_session is False
streaming._collect_user_input_from_client_event(
json.dumps({"type": "session.update", "session": {"type": "transcription"}})
)
assert streaming._is_transcription_session is True
@pytest.mark.asyncio
async def test_transcription_session_update_enforces_authorized_flat_model():
backend_ws = MagicMock()
backend_ws.send = AsyncMock()
streaming = RealTimeStreaming(
MagicMock(),
backend_ws,
MagicMock(),
model="gpt-realtime-whisper",
force_transcription_model="gpt-realtime-whisper",
)
await streaming._send_to_backend(
json.dumps(
{
"type": "session.update",
"session": {
"type": "transcription",
"input_audio_transcription": {
"model": "restricted-transcription-model",
"language": "en",
},
},
}
)
)
sent = json.loads(backend_ws.send.await_args.args[0])
assert sent["session"]["input_audio_transcription"] == {
"model": "gpt-realtime-whisper",
"language": "en",
}
assert streaming._is_transcription_session is True
@pytest.mark.asyncio
async def test_transcription_session_update_enforces_authorized_nested_model():
backend_ws = MagicMock()
backend_ws.send = AsyncMock()
streaming = RealTimeStreaming(
MagicMock(),
backend_ws,
MagicMock(),
model="gpt-realtime-whisper",
force_transcription_model="gpt-realtime-whisper",
)
await streaming._send_to_backend(
json.dumps(
{
"type": "session.update",
"session": {
"type": "transcription",
"audio": {
"input": {
"transcription": {
"model": "restricted-transcription-model",
"prompt": "domain words",
},
"format": {"type": "audio/pcm", "rate": 24000},
}
},
},
}
)
)
sent = json.loads(backend_ws.send.await_args.args[0])
assert sent["session"]["audio"]["input"]["transcription"] == {
"model": "gpt-realtime-whisper",
"prompt": "domain words",
}
assert sent["session"]["audio"]["input"]["format"] == {
"type": "audio/pcm",
"rate": 24000,
}
assert streaming._is_transcription_session is True
@pytest.mark.asyncio
async def test_normal_realtime_session_keeps_nested_transcription_model():
backend_ws = MagicMock()
backend_ws.send = AsyncMock()
streaming = RealTimeStreaming(
MagicMock(),
backend_ws,
MagicMock(),
model="gpt-4o-realtime-preview",
)
await streaming._send_to_backend(
json.dumps(
{
"type": "session.update",
"session": {
"type": "realtime",
"audio": {
"input": {
"transcription": {
"model": "whisper-1",
"language": "en",
}
}
},
},
}
)
)
sent = json.loads(backend_ws.send.await_args.args[0])
assert sent["session"]["audio"]["input"]["transcription"] == {
"model": "whisper-1",
"language": "en",
}
assert streaming._is_transcription_session is False
def test_detect_transcription_session_from_backend_transcription_session_events():
"""Backend transcription_session.created/updated events flag the session."""
streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock())
assert streaming._is_transcription_session is False
streaming._detect_transcription_session_from_backend(
{"type": "transcription_session.created"}
)
assert streaming._is_transcription_session is True
streaming2 = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock())
streaming2._detect_transcription_session_from_backend(
{"type": "transcription_session.updated"}
)
assert streaming2._is_transcription_session is True
def test_detect_transcription_session_from_backend_session_created_with_type():
"""Backend session.created with type=transcription flags the session."""
streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock())
streaming._detect_transcription_session_from_backend(
{"type": "session.created", "session": {"type": "transcription"}}
)
assert streaming._is_transcription_session is True
def test_detect_transcription_session_from_backend_ignores_non_transcription():
"""Backend session.created without type=transcription does not flag the session."""
streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock())
streaming._detect_transcription_session_from_backend(
{"type": "session.created", "session": {"model": "gpt-4o-realtime-preview"}}
)
assert streaming._is_transcription_session is False
def test_capture_transcription_usage_deduplicates_when_already_stored():
"""
When the event is already in messages (logged via store_message), it must not
be appended a second time by _capture_transcription_usage.
"""
import litellm
streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock())
# Add the event type to the default logged list so _should_store_message returns True.
streaming.logged_real_time_event_types = [
"conversation.item.input_audio_transcription.completed"
]
event = {
"type": "conversation.item.input_audio_transcription.completed",
"usage": {"type": "duration", "seconds": 5.0},
}
streaming.store_message(json.dumps(event))
initial_count = len(streaming.messages)
streaming._capture_transcription_usage(event)
assert len(streaming.messages) == initial_count # no duplicate
@pytest.mark.asyncio
async def test_client_ack_caches_setup_to_prevent_duplicate_session_update_setup():
websocket = MagicMock()

View file

@ -0,0 +1,312 @@
"""
Regression tests for issue #30014.
When LiteLLM proxies ``client -> /v1/messages -> /v1/chat/completions`` and a
streaming chunk both *triggers* a new Anthropic content block (its type differs
from the active block) and *carries* the first delta of that new block, the
trigger chunk's delta must be re-emitted as a ``content_block_delta``.
The synthesized ``content_block_start`` always carries an empty body, so before
the fix the first non-empty ``text_delta`` of every transitioned block was
silently dropped — e.g. text resuming after a tool call started from the second
token ("The weather is nice." was lost, "Hi" rendered as ""). Bundled
``input_json_delta`` tool arguments were already preserved and must stay
preserved, and empty trigger deltas must not produce spurious events.
"""
import os
import sys
from typing import List, Optional
from unittest.mock import MagicMock
import pytest
sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.llms.anthropic.experimental_pass_through.adapters.streaming_iterator import (
AnthropicStreamWrapper,
)
from litellm.types.utils import (
ChatCompletionDeltaToolCall,
Delta,
Function,
StreamingChoices,
)
def _make_chunk(delta: Delta, finish_reason: Optional[str] = None) -> MagicMock:
chunk = MagicMock()
chunk.choices = [
StreamingChoices(
finish_reason=finish_reason,
index=0,
delta=delta,
logprobs=None,
)
]
chunk.usage = None
chunk._hidden_params = {}
return chunk
def _tool_chunk(
call_id: str, name: Optional[str], arguments: Optional[str]
) -> MagicMock:
return _make_chunk(
Delta(
content=None,
tool_calls=[
ChatCompletionDeltaToolCall(
id=call_id,
function=Function(name=name, arguments=arguments),
type="function",
index=0,
)
],
)
)
class _AsyncStream:
def __init__(self, items: List[MagicMock]):
self._it = iter(items)
def __aiter__(self):
return self
async def __anext__(self):
try:
return next(self._it)
except StopIteration:
raise StopAsyncIteration
def _drain_sync(wrapper: AnthropicStreamWrapper) -> List[dict]:
return list(wrapper)
async def _drain_async(wrapper: AnthropicStreamWrapper) -> List[dict]:
return [event async for event in wrapper]
def _text_deltas(events: List[dict]) -> List[str]:
return [
e["delta"]["text"]
for e in events
if e.get("type") == "content_block_delta"
and e["delta"].get("type") == "text_delta"
]
def _input_json_deltas(events: List[dict]) -> List[str]:
return [
e["delta"]["partial_json"]
for e in events
if e.get("type") == "content_block_delta"
and e["delta"].get("type") == "input_json_delta"
]
def test_first_text_delta_after_tool_use_is_not_dropped_sync():
"""A tool_use -> text transition (text resuming after a tool call) carries
the resumed text's first token in the trigger chunk. Without the fix it was
dropped, so "The weather is nice." vanished and the answer began at " Bye.".
"""
chunks = [
_make_chunk(Delta(content="Let me check.")),
_tool_chunk("call_1", "get_weather", '{"city":'),
_tool_chunk("call_1", None, ' "NY"}'),
_make_chunk(Delta(content="The weather is nice.")),
_make_chunk(Delta(content=" Bye.")),
_make_chunk(Delta(content=None), finish_reason="stop"),
]
wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="claude-x")
events = _drain_sync(wrapper)
assert _input_json_deltas(events) == ['{"city":', ' "NY"}']
assert _text_deltas(events) == [
"Let me check.",
"The weather is nice.",
" Bye.",
]
@pytest.mark.asyncio
async def test_first_text_delta_after_tool_use_is_not_dropped_async():
"""Async path mirrors the sync regression — the proxy serves the async
iterator, so it must preserve the first resumed text delta too.
"""
chunks = [
_make_chunk(Delta(content="Let me check.")),
_tool_chunk("call_1", "get_weather", '{"city": "NY"}'),
_make_chunk(Delta(content="The weather is nice.")),
_make_chunk(Delta(content=" Bye.")),
_make_chunk(Delta(content=None), finish_reason="stop"),
]
wrapper = AnthropicStreamWrapper(
completion_stream=_AsyncStream(chunks), model="claude-x"
)
events = await _drain_async(wrapper)
assert _input_json_deltas(events) == ['{"city": "NY"}']
assert _text_deltas(events) == [
"Let me check.",
"The weather is nice.",
" Bye.",
]
def test_single_first_text_token_after_tool_use_preserved_sync():
"""Minimal reproduction of the issue's example: a single short text token
("Hi") resuming after a tool call. Without the fix the whole answer is
dropped because its only delta sits in the transition trigger chunk.
"""
chunks = [
_tool_chunk("call_1", "get_weather", '{"city": "NY"}'),
_make_chunk(Delta(content="Hi")),
_make_chunk(Delta(content=None), finish_reason="stop"),
]
wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="claude-x")
events = _drain_sync(wrapper)
assert _text_deltas(events) == ["Hi"]
def test_multiple_text_deltas_after_tool_use_preserved_sync():
"""Multiple-delta edge case: only the *first* text delta sits in the
transition trigger chunk; the rest stream normally. All of them — leading
one included — must reach the client in order.
"""
chunks = [
_tool_chunk("call_1", "get_weather", '{"city": "NY"}'),
_make_chunk(Delta(content="Hi")),
_make_chunk(Delta(content=", how ")),
_make_chunk(Delta(content="can I help ")),
_make_chunk(Delta(content="you?")),
_make_chunk(Delta(content=None), finish_reason="stop"),
]
wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="claude-x")
events = _drain_sync(wrapper)
assert _text_deltas(events) == ["Hi", ", how ", "can I help ", "you?"]
assert "".join(_text_deltas(events)) == "Hi, how can I help you?"
def test_empty_trigger_delta_is_not_re_emitted_sync():
"""A transition whose trigger chunk carries no content (empty text) must
NOT produce a spurious empty ``content_block_delta`` — only the synthesized
``content_block_start`` is emitted for the new block. Here a ``tool_use ->
text`` transition is triggered by an empty-content chunk; the re-emit guard
must reject it so the new text block opens without a leading empty delta.
"""
chunks = [
_tool_chunk("call_1", "get_weather", '{"city": "NY"}'),
# tool_use -> text transition triggered by an empty content chunk; the
# real text arrives in the following chunk.
_make_chunk(Delta(content="")),
_make_chunk(Delta(content="real text")),
_make_chunk(Delta(content=None), finish_reason="stop"),
]
wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="claude-x")
events = _drain_sync(wrapper)
# No empty-string text_delta should be present.
assert "" not in _text_deltas(events)
assert "".join(_text_deltas(events)) == "real text"
def test_bundled_tool_args_on_transition_still_preserved_sync():
"""Existing behavior guard: when the trigger chunk that opens a tool_use
block also carries arguments (xAI/Gemini style), the ``input_json_delta``
must still be emitted after ``content_block_start``.
"""
chunks = [
_make_chunk(Delta(content="Calling a tool.")),
_tool_chunk("call_1", "get_weather", '{"city": "NY"}'),
_make_chunk(Delta(content=None), finish_reason="tool_calls"),
]
wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="claude-x")
events = _drain_sync(wrapper)
assert _text_deltas(events) == ["Calling a tool."]
assert _input_json_deltas(events) == ['{"city": "NY"}']
@pytest.mark.parametrize(
"processed_chunk, expected",
[
# Non-empty deltas of every type must be re-emitted.
(
{
"type": "content_block_delta",
"delta": {"type": "text_delta", "text": "x"},
},
True,
),
(
{
"type": "content_block_delta",
"delta": {"type": "input_json_delta", "partial_json": "{}"},
},
True,
),
(
{
"type": "content_block_delta",
"delta": {"type": "thinking_delta", "thinking": "t"},
},
True,
),
(
{
"type": "content_block_delta",
"delta": {"type": "signature_delta", "signature": "s"},
},
True,
),
# Empty deltas must NOT be re-emitted (no spurious events).
(
{
"type": "content_block_delta",
"delta": {"type": "text_delta", "text": ""},
},
False,
),
(
{
"type": "content_block_delta",
"delta": {"type": "input_json_delta", "partial_json": ""},
},
False,
),
(
{
"type": "content_block_delta",
"delta": {"type": "thinking_delta", "thinking": ""},
},
False,
),
(
{
"type": "content_block_delta",
"delta": {"type": "signature_delta", "signature": ""},
},
False,
),
# Unknown delta type / non-content_block_delta / malformed delta.
(
{"type": "content_block_delta", "delta": {"type": "other_delta"}},
False,
),
({"type": "message_delta", "delta": {"stop_reason": "stop"}}, False),
({"type": "content_block_delta", "delta": None}, False),
],
)
def test_trigger_delta_has_content_branches(processed_chunk, expected):
"""Directly exercise the re-emit predicate across all delta types and the
empty/malformed guards, so the helper's behavior is pinned independently of
upstream chunk-translation details.
"""
assert (
AnthropicStreamWrapper._trigger_delta_has_content(processed_chunk) is expected
)

View file

@ -278,7 +278,8 @@ def test_anthropic_stream_wrapper_interleaved_tool_calls_and_text():
"content_block_delta", # {"city":
"content_block_delta", # "NY"}
"content_block_stop", # End of first tool_use content block
"content_block_start", # "The weather is nice today"
"content_block_start", # "The weather is nice today" text block
"content_block_delta", # "The weather is nice today." text_delta
"content_block_stop",
"content_block_start", # Start of second tool_use content block
"content_block_delta", # {"city":
@ -288,7 +289,8 @@ def test_anthropic_stream_wrapper_interleaved_tool_calls_and_text():
"content_block_delta", # {"city":
"content_block_delta", # " CHI"}
"content_block_stop", # End of third tool_use content block
"content_block_start", # "The weather is not so nice today"
"content_block_start", # "The weather is not so nice today" text block
"content_block_delta", # "The weather is not so nice today." text_delta
"content_block_stop",
"message_delta", # Stop reason with merged usage
"message_stop", # Final message stop
@ -296,6 +298,20 @@ def test_anthropic_stream_wrapper_interleaved_tool_calls_and_text():
assert expected_types == chunk_types
# Regression: the first (and only) text delta of each text block sits in
# the chunk that *triggered* the tool_use -> text transition. It must be
# re-emitted as a content_block_delta instead of being silently dropped.
text_deltas = [
chunk["delta"]["text"]
for chunk in chunks
if chunk.get("type") == "content_block_delta"
and chunk["delta"].get("type") == "text_delta"
]
assert text_deltas == [
"The weather is nice today.",
"The weather is not so nice today.",
]
get_weather_calls = 0
for chunk in chunks:

View file

@ -147,6 +147,103 @@ async def test_construct_url_ga_protocol():
assert "deployment" not in url
@pytest.mark.asyncio
async def test_construct_url_forwards_transcription_intent_ga():
"""
Transcription sessions connect with intent=transcription. The Azure handler
must forward that query param so gpt-realtime-whisper opens a transcription
session instead of a normal realtime session.
"""
from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime
handler = AzureOpenAIRealtime()
url = handler._construct_url(
api_base="https://my-endpoint.openai.azure.com",
model="gpt-realtime-whisper",
api_version="2025-04-01-preview",
realtime_protocol="GA",
query_params={"model": "gpt-realtime-whisper", "intent": "transcription"},
)
assert "/openai/v1/realtime?" in url
assert "intent=transcription" in url
assert "model=" not in url
@pytest.mark.asyncio
async def test_construct_url_forwards_transcription_intent_ga_without_model_query():
"""
OpenAI-compatible transcription clients may connect with only
intent=transcription and send the transcription model in session.update.
Preserve that query shape instead of forcing model= into the upstream URL.
"""
from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime
handler = AzureOpenAIRealtime()
url = handler._construct_url(
api_base="https://my-endpoint.openai.azure.com",
model="gpt-realtime-whisper",
api_version="2025-04-01-preview",
realtime_protocol="GA",
query_params={"intent": "transcription"},
)
assert url == (
"wss://my-endpoint.openai.azure.com/openai/v1/realtime"
"?intent=transcription"
)
@pytest.mark.asyncio
async def test_construct_url_forwards_transcription_intent_beta():
from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime
handler = AzureOpenAIRealtime()
url = handler._construct_url(
api_base="https://my-endpoint.openai.azure.com",
model="whisper-deploy",
api_version="2024-10-01-preview",
query_params={"intent": "transcription"},
)
assert "/openai/realtime?" in url
assert "deployment=whisper-deploy" in url
assert "intent=transcription" in url
@pytest.mark.asyncio
async def test_construct_url_encodes_intent_value():
"""A crafted intent value must be URL-encoded, not injected as raw query params."""
from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime
handler = AzureOpenAIRealtime()
url = handler._construct_url(
api_base="https://my-endpoint.openai.azure.com",
model="gpt-realtime-whisper",
api_version="2025-04-01-preview",
realtime_protocol="GA",
query_params={"intent": "transcription&foo=bar"},
)
assert "intent=transcription%26foo%3Dbar" in url
assert "&foo=bar" not in url
@pytest.mark.asyncio
async def test_construct_url_no_intent_when_absent():
"""No intent param leaks into the URL when not provided."""
from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime
handler = AzureOpenAIRealtime()
url = handler._construct_url(
api_base="https://my-endpoint.openai.azure.com",
model="gpt-4o-realtime-preview",
api_version="2024-10-01-preview",
realtime_protocol="GA",
query_params={"model": "gpt-4o-realtime-preview"},
)
assert "intent=" not in url
@pytest.mark.asyncio
async def test_construct_url_v1_protocol():
"""
@ -368,6 +465,45 @@ async def test_realtime_protocol_from_litellm_params():
assert litellm_params.get("realtime_protocol") == "GA"
@pytest.mark.asyncio
async def test_arealtime_transcription_intent_defaults_to_ga(monkeypatch):
"""
Azure gpt-realtime-whisper transcription connects on the GA /openai/v1/realtime
path. If the DB model lacks realtime_protocol, infer GA from intent=transcription.
"""
from litellm.realtime_api import main as realtime_main
mock_async_realtime = AsyncMock()
monkeypatch.setattr(
realtime_main,
"azure_realtime",
MagicMock(async_realtime=mock_async_realtime),
)
def fake_get_llm_provider(model, api_base=None, api_key=None):
return (
"gpt-realtime-whisper",
"azure",
"test-key",
"https://my-endpoint.openai.azure.com",
)
monkeypatch.setattr(realtime_main, "get_llm_provider", fake_get_llm_provider)
await realtime_main._arealtime(
model="azure/gpt-realtime-whisper",
websocket=MagicMock(),
api_key="test-key",
api_version="2025-04-01-preview",
query_params={"intent": "transcription"},
litellm_logging_obj=MagicMock(),
)
called_kwargs = mock_async_realtime.call_args.kwargs
assert called_kwargs["realtime_protocol"] == "GA"
assert called_kwargs["query_params"] == {"intent": "transcription"}
@pytest.mark.asyncio
async def test_async_realtime_default_maintains_backwards_compatibility():
"""

View file

@ -5386,3 +5386,45 @@ def test_converse_top_k_zero_forwarded_on_models_that_accept_it():
)
assert result["additionalModelRequestFields"]["top_k"] == 0
@pytest.mark.asyncio
async def test_grounding_source_and_query_rendered_as_text():
"""grounding_source / query content blocks must render as plain text on the
generate path (the model needs to see the RAG context + question). The bedrock
converse dispatch silently drops unrecognised content types, so these would
otherwise vanish from the prompt."""
from litellm.litellm_core_utils.prompt_templates.factory import (
BedrockConverseMessagesProcessor,
_bedrock_converse_messages_pt,
)
messages = [
{
"role": "user",
"content": [
{"type": "grounding_source", "text": "Tokyo is the capital of Japan."},
{"type": "query", "text": "What is the capital of Japan?"},
],
}
]
result = _bedrock_converse_messages_pt(
messages=messages,
model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
llm_provider="bedrock_converse",
)
async_result = (
await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
messages=messages,
model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
llm_provider="bedrock_converse",
)
)
assert result == async_result
assert len(result) == 1
assert result[0]["role"] == "user"
user_content = result[0]["content"]
assert {"text": "Tokyo is the capital of Japan."} in user_content
assert {"text": "What is the capital of Japan?"} in user_content

View file

@ -81,6 +81,77 @@ def test_prepare_fake_stream_request():
assert result_data["messages"] == [{"role": "user", "content": "Hello"}]
def test_response_api_handler_streams_when_provider_transform_adds_stream():
handler = BaseLLMHTTPHandler()
config = Mock()
config.validate_environment.return_value = {}
config.get_complete_url.return_value = "https://chatgpt.example.com/responses"
config.transform_responses_api_request.return_value = {
"model": "gpt-5.3-codex",
"input": "hi",
"stream": True,
}
config.sign_request.return_value = ({}, None)
client = HTTPHandler(client=httpx.Client())
client.post = Mock(
return_value=httpx.Response(
200,
request=httpx.Request("POST", "https://chatgpt.example.com/responses"),
)
)
logging_obj = Mock()
handler.response_api_handler(
model="gpt-5.3-codex",
input="hi",
responses_api_provider_config=config,
response_api_optional_request_params={},
custom_llm_provider="chatgpt",
litellm_params=GenericLiteLLMParams(),
logging_obj=logging_obj,
client=client,
)
assert client.post.call_args.kwargs["stream"] is True
assert client.post.call_args.kwargs["json"]["stream"] is True
@pytest.mark.asyncio
async def test_async_response_api_handler_streams_when_provider_transform_adds_stream():
handler = BaseLLMHTTPHandler()
config = Mock()
config.validate_environment.return_value = {}
config.get_complete_url.return_value = "https://chatgpt.example.com/responses"
config.transform_responses_api_request.return_value = {
"model": "gpt-5.3-codex",
"input": "hi",
"stream": True,
}
config.sign_request.return_value = ({}, None)
client = AsyncHTTPHandler()
client.post = AsyncMock(
return_value=httpx.Response(
200,
request=httpx.Request("POST", "https://chatgpt.example.com/responses"),
)
)
logging_obj = Mock()
await handler.async_response_api_handler(
model="gpt-5.3-codex",
input="hi",
responses_api_provider_config=config,
response_api_optional_request_params={},
custom_llm_provider="chatgpt",
litellm_params=GenericLiteLLMParams(),
logging_obj=logging_obj,
client=client,
)
assert client.post.call_args.kwargs["stream"] is True
assert client.post.call_args.kwargs["json"]["stream"] is True
def test_get_agentic_loop_settings_defaults_and_overrides():
handler = BaseLLMHTTPHandler()

View file

@ -585,3 +585,221 @@ class TestGithubCopilotResponsesAPIRouting:
provider=LlmProviders.GITHUB_COPILOT,
)
assert isinstance(config, GithubCopilotResponsesAPIConfig)
class TestGithubCopilotReasoningStreamItemIdNormalization:
"""GitHub Copilot's native /responses stream tags every reasoning-summary
event with a different item_id (and the reasoning output_item.added /
output_item.done ids also differ). Strict clients (Vercel ai-sdk) key
reasoning state by item_id and crash when a summary delta references an
unregistered id. The config normalizes every reasoning event in an
output_index group to the id from its output_item.added."""
def _config(self):
with patch(
"litellm.llms.github_copilot.responses.transformation.Authenticator"
):
return GithubCopilotResponsesAPIConfig()
def _transform(self, config, chunk):
return config.transform_streaming_response(
model="github_copilot/gpt-5.5",
parsed_chunk=chunk,
logging_obj=MagicMock(),
)
def test_summary_events_normalized_to_output_item_added_id(self):
config = self._config()
self._transform(
config,
{
"type": "response.output_item.added",
"output_index": 0,
"item": {"id": "stable_rs_id", "type": "reasoning"},
},
)
summary_chunks = [
{
"type": "response.reasoning_summary_part.added",
"output_index": 0,
"summary_index": 0,
"item_id": "bad_part_added",
"part": {"type": "summary_text", "text": ""},
},
{
"type": "response.reasoning_summary_text.delta",
"output_index": 0,
"summary_index": 0,
"item_id": "bad_delta_1",
"delta": "Hello",
},
{
"type": "response.reasoning_summary_text.delta",
"output_index": 0,
"summary_index": 0,
"item_id": "bad_delta_2",
"delta": " world",
},
{
"type": "response.reasoning_summary_text.done",
"output_index": 0,
"summary_index": 0,
"item_id": "bad_text_done",
"text": "Hello world",
},
{
"type": "response.reasoning_summary_part.done",
"output_index": 0,
"summary_index": 0,
"item_id": "bad_part_done",
"part": {"type": "summary_text", "text": "Hello world"},
},
]
for chunk in summary_chunks:
event = self._transform(config, chunk)
assert event.item_id == "stable_rs_id"
def test_reasoning_output_item_done_normalized_to_added_id(self):
config = self._config()
self._transform(
config,
{
"type": "response.output_item.added",
"output_index": 0,
"item": {"id": "stable_rs_id", "type": "reasoning"},
},
)
event = self._transform(
config,
{
"type": "response.output_item.done",
"output_index": 0,
"item": {
"id": "different_done_id",
"type": "reasoning",
"encrypted_content": "ENC",
},
},
)
assert event.item.id == "stable_rs_id"
def test_interleaved_message_item_does_not_corrupt_mapping(self):
config = self._config()
self._transform(
config,
{
"type": "response.output_item.added",
"output_index": 0,
"item": {"id": "stable_rs_id", "type": "reasoning"},
},
)
self._transform(
config,
{
"type": "response.output_item.added",
"output_index": 1,
"item": {"id": "msg_id", "type": "message"},
},
)
event = self._transform(
config,
{
"type": "response.reasoning_summary_text.delta",
"output_index": 0,
"summary_index": 0,
"item_id": "bad_delta",
"delta": "x",
},
)
assert event.item_id == "stable_rs_id"
def test_event_without_registered_item_passes_through_unchanged(self):
config = self._config()
event = self._transform(
config,
{
"type": "response.output_text.delta",
"output_index": 0,
"item_id": "msg_native_id",
"delta": "hi",
},
)
assert event.item_id == "msg_native_id"
def test_message_text_events_normalized_to_output_item_added_id(self):
config = self._config()
self._transform(
config,
{
"type": "response.output_item.added",
"output_index": 0,
"item": {"id": "stable_msg_id", "type": "message"},
},
)
for chunk in [
{
"type": "response.content_part.added",
"output_index": 0,
"content_index": 0,
"item_id": "bad_cp_added",
"part": {"type": "output_text", "text": ""},
},
{
"type": "response.output_text.delta",
"output_index": 0,
"content_index": 0,
"item_id": "bad_text_delta",
"delta": "Paris",
},
{
"type": "response.output_text.done",
"output_index": 0,
"content_index": 0,
"item_id": "bad_text_done",
"text": "Paris",
},
]:
event = self._transform(config, chunk)
assert event.item_id == "stable_msg_id"
def test_event_without_output_index_passes_through_unchanged(self):
config = self._config()
event = self._transform(
config,
{
"type": "response.completed",
"response": {"id": "resp_1", "status": "completed", "output": []},
},
)
assert event.type == "response.completed"
def test_normalization_continues_after_a_terminal_event(self):
config = self._config()
self._transform(
config,
{
"type": "response.output_item.added",
"output_index": 0,
"item": {"id": "stable_rs_id", "type": "reasoning"},
},
)
self._transform(
config,
{
"type": "response.completed",
"response": {"id": "resp_1", "status": "completed", "output": []},
},
)
event = self._transform(
config,
{
"type": "response.reasoning_summary_text.delta",
"output_index": 0,
"summary_index": 0,
"item_id": "bad_delta",
"delta": "x",
},
)
assert event.item_id == "stable_rs_id"

View file

@ -11,6 +11,7 @@ import litellm
sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm import ModelResponse
from litellm.constants import DEFAULT_OCI_CHAT_MAX_TOKENS
from litellm.llms.oci.chat.transformation import (
OCIChatConfig,
OCIRequestWrapper,
@ -104,6 +105,7 @@ class TestOCIChatConfig:
"chatRequest": {
"apiFormat": "GENERIC",
"isStream": False,
"maxTokens": DEFAULT_OCI_CHAT_MAX_TOKENS,
"messages": [
{
"role": "USER",
@ -362,6 +364,137 @@ class TestOCIChatConfig:
rf = transformed_request["chatRequest"]["responseFormat"]
assert rf["type"] == "JSON_OBJECT"
def test_transform_request_response_format_json_schema_generic(self):
"""A GENERIC json_schema must become OCI's JSON_SCHEMA shape with the
OpenAI ``strict`` key renamed to ``isStrict``.
OCI's ResponseJsonSchema rejects ``strict`` (and any other extra key)
with HTTP 400 "Please pass in correct format of request", so the raw
OpenAI body must not be forwarded.
"""
config = OCIChatConfig()
optional_params = {
"oci_compartment_id": TEST_COMPARTMENT_ID,
"response_format": {
"type": "json_schema",
"json_schema": {
"name": "judgment",
"description": "a score and rationale",
"strict": True,
"schema": {
"type": "object",
"properties": {"score": {"type": "integer"}},
"required": ["score"],
},
},
},
}
transformed_request = config.transform_request(
model=TEST_MODEL_NAME, # xai.grok-4 -> GENERIC
messages=TEST_MESSAGES, # type: ignore
optional_params=optional_params,
litellm_params={},
headers={},
)
rf = transformed_request["chatRequest"]["responseFormat"]
assert rf["type"] == "JSON_SCHEMA"
assert "strict" not in rf["jsonSchema"]
assert rf["jsonSchema"]["isStrict"] is True
assert rf["jsonSchema"]["name"] == "judgment"
assert rf["jsonSchema"]["description"] == "a score and rationale"
assert rf["jsonSchema"]["schema"]["properties"]["score"]["type"] == "integer"
def test_transform_request_response_format_json_schema_generic_no_strict(self):
"""A GENERIC json_schema without ``strict`` must omit ``isStrict``."""
config = OCIChatConfig()
optional_params = {
"oci_compartment_id": TEST_COMPARTMENT_ID,
"response_format": {
"type": "json_schema",
"json_schema": {"name": "j", "schema": {"type": "object"}},
},
}
transformed_request = config.transform_request(
model=TEST_MODEL_NAME,
messages=TEST_MESSAGES, # type: ignore
optional_params=optional_params,
litellm_params={},
headers={},
)
rf = transformed_request["chatRequest"]["responseFormat"]
assert rf["type"] == "JSON_SCHEMA"
assert "isStrict" not in rf["jsonSchema"]
def test_transform_request_response_format_json_schema_cohere(self):
"""A Cohere json_schema must fold the schema onto JSON_OBJECT.
OCI Cohere has no JSON_SCHEMA type; sending one yields HTTP 400.
"""
config = OCIChatConfig()
optional_params = {
"oci_compartment_id": TEST_COMPARTMENT_ID,
"response_format": {
"type": "json_schema",
"json_schema": {
"name": "judgment",
"strict": True,
"schema": {
"type": "object",
"properties": {"score": {"type": "integer"}},
},
},
},
}
transformed_request = config.transform_request(
model="cohere.command-latest",
messages=TEST_MESSAGES, # type: ignore
optional_params=optional_params,
litellm_params={},
headers={},
)
rf = transformed_request["chatRequest"]["responseFormat"]
assert rf["type"] == "JSON_OBJECT"
assert "jsonSchema" not in rf
assert rf["schema"]["properties"]["score"]["type"] == "integer"
def test_transform_request_response_format_cohere_json_object(self):
"""Cohere json_object without a schema stays a bare JSON_OBJECT."""
config = OCIChatConfig()
optional_params = {
"oci_compartment_id": TEST_COMPARTMENT_ID,
"response_format": {"type": "json_object"},
}
transformed_request = config.transform_request(
model="cohere.command-latest",
messages=TEST_MESSAGES, # type: ignore
optional_params=optional_params,
litellm_params={},
headers={},
)
rf = transformed_request["chatRequest"]["responseFormat"]
assert rf == {"type": "JSON_OBJECT"}
def test_transform_request_json_schema_without_body_raises_generic(self):
"""A GENERIC json_schema with no ``json_schema`` body must raise an early
400, not silently emit {"type": "JSON_SCHEMA"} (which OCI rejects)."""
from litellm.llms.oci.common_utils import OCIError
config = OCIChatConfig()
optional_params = {
"oci_compartment_id": TEST_COMPARTMENT_ID,
"response_format": {"type": "json_schema"},
}
with pytest.raises(OCIError) as exc_info:
config.transform_request(
model=TEST_MODEL_NAME, # GENERIC
messages=TEST_MESSAGES, # type: ignore
optional_params=optional_params,
litellm_params={},
headers={},
)
assert exc_info.value.status_code == 400
assert "json_schema" in str(exc_info.value)
def test_transform_response_without_token_details(self):
"""
Tests that responses missing completionTokensDetails and promptTokensDetails
@ -956,6 +1089,44 @@ class TestOCICohereParamMapping:
assert result.get("temperature") == 0.5
class TestOCIDefaultMaxTokens:
"""Regression for OCI's tiny server-side token cap (~20 tokens), which
silently truncated responses mid-string whenever the caller omitted
max_tokens (MLflow judges never send it, so their JSON came back cut off).
transform_request injects DEFAULT_OCI_CHAT_MAX_TOKENS when no limit is
supplied, and leaves an explicit limit untouched."""
def _chat_request(self, model: str, optional_params: dict) -> dict:
config = OCIChatConfig()
body = config.transform_request(
model=model,
messages=[{"role": "user", "content": "hi"}],
optional_params={**BASE_OCI_PARAMS, **optional_params},
litellm_params={},
headers={},
)
return body["chatRequest"]
@pytest.mark.parametrize(
"model", ["cohere.command-latest", "meta.llama-3.3-70b-instruct"]
)
def test_default_injected_when_max_tokens_omitted(self, model):
chat_request = self._chat_request(model, {})
assert chat_request["maxTokens"] == DEFAULT_OCI_CHAT_MAX_TOKENS
@pytest.mark.parametrize(
"model", ["cohere.command-latest", "meta.llama-3.3-70b-instruct"]
)
def test_explicit_max_tokens_not_overridden(self, model):
chat_request = self._chat_request(model, {"max_tokens": 256})
assert chat_request["maxTokens"] == 256
def test_reasoning_model_defaults_max_completion_tokens(self):
chat_request = self._chat_request("openai.gpt-5", {})
assert chat_request["maxCompletionTokens"] == DEFAULT_OCI_CHAT_MAX_TOKENS
assert "maxTokens" not in chat_request
class TestOCIReasoningEffort:
"""
Reasoning-effort handling for GENERIC reasoning models:
@ -1133,8 +1304,7 @@ class TestOCIStreamingSignedBody:
When signed_json_body is provided, the POST must use that exact bytes object,
not json.dumps(data) — otherwise the RSA-SHA256 signature is invalid.
"""
import httpx
from unittest.mock import MagicMock, patch
from unittest.mock import MagicMock
config = OCIChatConfig()
signed_bytes = b'{"signed": true}'
@ -1293,6 +1463,68 @@ class TestOCIChatConfigErrorPaths:
)
assert "audio" not in result
@pytest.mark.parametrize("model", ["cohere.command-latest", "xai.grok-4"])
def test_map_openai_params_max_retries_dropped_without_drop_params(self, model):
"""max_retries is a litellm control param, not a generation param. It
must be dropped silently (no raise) even when drop_params is False, so
the litellm proxy (which injects max_retries on every request) does not
500 every OCI call.
"""
config = OCIChatConfig()
result = config.map_openai_params(
non_default_params={"max_retries": 3},
optional_params={},
model=model,
drop_params=False,
)
assert "max_retries" not in result
def test_map_openai_params_cohere_n_default_dropped(self):
"""Cohere has no numGenerations field, but n=1 (and None) is the OpenAI
default single-generation request. It must be dropped silently rather
than raising, so standard clients that always send n=1 (e.g. the MLflow
gateway) are not rejected."""
config = OCIChatConfig()
for n in (1, None):
result = config.map_openai_params(
non_default_params={"n": n},
optional_params={},
model="cohere.command-latest",
drop_params=False,
)
assert "n" not in result and "numGenerations" not in result
def test_map_openai_params_cohere_n_gt_1_raises_without_drop(self):
"""n>1 is genuinely unsupported on Cohere and must raise without drop."""
config = OCIChatConfig()
with pytest.raises(Exception, match="not supported on OCI"):
config.map_openai_params(
non_default_params={"n": 3},
optional_params={},
model="cohere.command-latest",
drop_params=False,
)
def test_map_openai_params_cohere_n_gt_1_dropped_with_drop(self):
config = OCIChatConfig()
result = config.map_openai_params(
non_default_params={"n": 3},
optional_params={},
model="cohere.command-latest",
drop_params=True,
)
assert "n" not in result and "numGenerations" not in result
def test_map_openai_params_generic_n_maps_to_num_generations(self):
"""Generic models keep numGenerations, including n>1."""
config = OCIChatConfig()
result = config.map_openai_params(
non_default_params={"n": 2},
optional_params={},
model=TEST_MODEL_NAME,
drop_params=False,
)
assert result["numGenerations"] == 2
def test_transform_request_tool_choice_string_mapped(self):
config = OCIChatConfig()
result = config.transform_request(

View file

@ -5,6 +5,7 @@ import json
from unittest.mock import patch, MagicMock
from litellm import ModelResponse
from litellm.constants import DEFAULT_OCI_CHAT_MAX_TOKENS
from litellm.llms.oci.chat.cohere import (
adapt_messages_to_cohere_standard,
adapt_tool_definitions_to_cohere_standard,
@ -236,25 +237,30 @@ class TestOCICohereToolCalls:
assert result.usage.completion_tokens == 22
assert result.usage.total_tokens == 48
def test_cohere_request_preserves_json_schema_response_format(self):
"""Ensure Cohere requests retain JSON schema payloads in responseFormat."""
def test_cohere_request_folds_json_schema_into_json_object(self):
"""A Cohere json_schema must fold the schema onto JSON_OBJECT.
OCI Cohere has no JSON_SCHEMA type; sending {"type": "JSON_SCHEMA", ...}
(or the raw lowercase "json_schema" with a jsonSchema body) is rejected
with HTTP 400. The schema rides on JSON_OBJECT instead.
"""
config = OCIChatConfig()
messages = [{"role": "user", "content": "Return structured info"}]
response_format = {
"type": "json_schema",
"json_schema": {
"name": "test_schema",
"strict": True,
"schema": {
"type": "object",
"properties": {"foo": {"type": "string"}},
"required": ["foo"],
},
},
schema = {
"type": "object",
"properties": {"foo": {"type": "string"}},
"required": ["foo"],
}
optional_params = {
"oci_compartment_id": TEST_COMPARTMENT_ID,
"response_format": response_format,
"response_format": {
"type": "json_schema",
"json_schema": {
"name": "test_schema",
"strict": True,
"schema": schema,
},
},
}
transformed_request = config.transform_request(
@ -265,18 +271,14 @@ class TestOCICohereToolCalls:
headers={},
)
chat_request = transformed_request["chatRequest"]
assert chat_request["apiFormat"] == "COHERE"
assert "responseFormat" in chat_request
cohere_response_format = chat_request["responseFormat"]
assert cohere_response_format["type"] == "json_schema"
cohere_response_format = transformed_request["chatRequest"]["responseFormat"]
assert cohere_response_format["type"] == "JSON_OBJECT"
assert "jsonSchema" not in cohere_response_format
assert "json_schema" not in cohere_response_format
assert "jsonSchema" in cohere_response_format
assert cohere_response_format["jsonSchema"] == response_format["json_schema"]
assert cohere_response_format["schema"] == schema
def test_cohere_request_response_format_text_stays_lowercase(self):
"""Ensure Cohere keeps response_format type lowercase (e.g. 'text' not 'TEXT')."""
def test_cohere_request_response_format_text_is_uppercased(self):
"""Cohere response_format type 'text' maps to OCI's canonical 'TEXT'."""
config = OCIChatConfig()
messages = [{"role": "user", "content": "Hello"}]
optional_params = {
@ -292,10 +294,7 @@ class TestOCICohereToolCalls:
headers={},
)
chat_request = transformed_request["chatRequest"]
assert chat_request["apiFormat"] == "COHERE"
assert "responseFormat" in chat_request
assert chat_request["responseFormat"]["type"] == "text"
assert transformed_request["chatRequest"]["responseFormat"] == {"type": "TEXT"}
def test_cohere_tool_call_only_message_no_text(self):
"""Test chat history with an assistant message that has tool calls but no text content."""
@ -462,7 +461,8 @@ class TestOCICohereToolCalls:
assert "tool_choice" not in supported_params
def test_cohere_default_parameters(self):
"""Test that Cohere requests do not inject hardcoded defaults — caller supplies all params."""
"""maxTokens is defaulted (OCI's server default truncates at ~20 tokens);
every other param is still pass-through with no hardcoded default."""
config = OCIChatConfig()
messages = [{"role": "user", "content": "Hello"}]
optional_params = {"oci_compartment_id": TEST_COMPARTMENT_ID}
@ -477,8 +477,7 @@ class TestOCICohereToolCalls:
chat_request = transformed_request["chatRequest"]
# No hardcoded defaults injected — only pass through what the user supplies
assert "maxTokens" not in chat_request
assert chat_request["maxTokens"] == DEFAULT_OCI_CHAT_MAX_TOKENS
assert "topK" not in chat_request
assert "topP" not in chat_request
assert "frequencyPenalty" not in chat_request

View file

@ -2,7 +2,6 @@
Unit tests for litellm/llms/oci/chat/generic.py — error paths and stream handling.
"""
import json
import pytest
from unittest.mock import MagicMock
@ -16,7 +15,12 @@ from litellm.llms.oci.chat.generic import (
handle_generic_response,
handle_generic_stream_chunk,
)
from litellm.llms.oci.chat.transformation import OCIChatConfig, OCIStreamWrapper
from litellm.llms.oci.chat.transformation import (
OCIChatConfig,
OCIStreamWrapper,
OCIVendors,
_model_uses_max_completion_tokens,
)
from litellm.llms.oci.common_utils import OCIError
# ---------------------------------------------------------------------------
@ -271,7 +275,6 @@ class TestHandleGenericStreamChunk:
assert result.choices[0].index == 0
def test_image_content_in_stream_raises(self):
from litellm.types.llms.oci import OCIImageContentPart, OCIImageUrl, OCIMessage
chunk = {
"apiFormat": "GENERIC",
@ -368,10 +371,6 @@ def _register_oci_gpt5_in_catalog():
class TestGpt5MaxCompletionTokens:
def test_helper_detects_gpt5_family(self, _register_oci_gpt5_in_catalog):
from litellm.llms.oci.chat.transformation import (
_model_uses_max_completion_tokens,
)
assert _model_uses_max_completion_tokens("openai.gpt-5") is True
assert _model_uses_max_completion_tokens("openai.gpt-5-mini") is True
assert _model_uses_max_completion_tokens("openai.gpt-5-nano") is True
@ -382,11 +381,40 @@ class TestGpt5MaxCompletionTokens:
assert _model_uses_max_completion_tokens("cohere.command-latest") is False
assert _model_uses_max_completion_tokens("") is False
def test_helper_covers_openai_models_absent_from_catalog(self):
"""OCI keeps adding OpenAI models (gpt-4.1, gpt-5.1..5.5, o-series)
faster than the litellm catalog tracks them. The vendor-prefix rule
must route them to maxCompletionTokens even with no catalog entry,
since OpenAI accepts max_completion_tokens on every chat model while
the reasoning families hard-reject max_tokens."""
import litellm
for name in (
"openai.gpt-5.2",
"openai.gpt-4.1",
"openai.o3",
"oci/openai.gpt-5.1-codex",
):
assert f"oci/{name.removeprefix('oci/')}" not in litellm.model_cost
assert _model_uses_max_completion_tokens(name) is True
assert _model_uses_max_completion_tokens("openai.gpt-oss-20b") is False
def test_default_injection_uses_max_completion_tokens_for_uncataloged_gpt(self):
"""Regression: with the injected default maxTokens, a GPT model absent
from the catalog got "maxTokens" on every request and OCI returned 400
("Use 'max_completion_tokens' instead") even when the caller never set
max_tokens."""
from litellm.constants import DEFAULT_OCI_CHAT_MAX_TOKENS
cfg = OCIChatConfig()
out = cfg._get_optional_params(OCIVendors.GENERIC, {}, model="openai.gpt-5.2")
assert out.get("maxCompletionTokens") == DEFAULT_OCI_CHAT_MAX_TOKENS
assert "maxTokens" not in out
def test_gpt5_routes_max_tokens_to_max_completion_tokens(
self, _register_oci_gpt5_in_catalog
):
from litellm.llms.oci.chat.transformation import OCIChatConfig, OCIVendors
cfg = OCIChatConfig()
# Both shapes optional_params can take after upstream map_openai_params:
# 1. openai-side key still present
@ -404,8 +432,6 @@ class TestGpt5MaxCompletionTokens:
assert "maxTokens" not in out_b
def test_non_gpt5_keeps_max_tokens(self):
from litellm.llms.oci.chat.transformation import OCIChatConfig, OCIVendors
cfg = OCIChatConfig()
out = cfg._get_optional_params(
OCIVendors.GENERIC,
@ -416,8 +442,6 @@ class TestGpt5MaxCompletionTokens:
assert "maxCompletionTokens" not in out
def test_cohere_reasoning_model_keeps_max_tokens(self):
from litellm.llms.oci.chat.transformation import OCIChatConfig, OCIVendors
cfg = OCIChatConfig()
out = cfg._get_optional_params(
OCIVendors.COHERE,

View file

@ -0,0 +1,226 @@
"""
Tests for the Realtime transcription_sessions surface used by gpt-realtime-whisper:
- OpenAI / Azure URL construction (POST /v1/realtime/transcription_sessions)
- RealtimeTranscriptionSessionRequest model-resolution + passthrough
- BaseLLMHTTPHandler.async_realtime_transcription_session_handler targeting
"""
import os
import sys
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.llms.azure.realtime.http_transformation import AzureRealtimeHTTPConfig
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.llms.openai.realtime.http_transformation import OpenAIRealtimeHTTPConfig
from litellm.types.realtime import RealtimeTranscriptionSessionRequest
def test_openai_transcription_session_url():
cfg = OpenAIRealtimeHTTPConfig()
assert (
cfg.get_transcription_session_url(
api_base="https://api.openai.com", model="gpt-realtime-whisper"
)
== "https://api.openai.com/v1/realtime/transcription_sessions"
)
def test_openai_transcription_session_url_strips_trailing_v1():
"""A /v1 suffix must not be duplicated in the path."""
cfg = OpenAIRealtimeHTTPConfig()
assert (
cfg.get_transcription_session_url(
api_base="https://api.openai.com/v1", model="gpt-realtime-whisper"
)
== "https://api.openai.com/v1/realtime/transcription_sessions"
)
def test_azure_transcription_session_url_uses_deployment_and_api_version():
cfg = AzureRealtimeHTTPConfig()
url = cfg.get_transcription_session_url(
api_base="https://my.openai.azure.com",
model="whisper-deploy",
api_version="2025-04-01-preview",
)
assert (
url
== "https://my.openai.azure.com/openai/realtime/transcription_sessions?api-version=2025-04-01-preview"
)
def test_request_resolves_model_returns_none_when_both_absent():
req = RealtimeTranscriptionSessionRequest(input_audio_format="pcm16")
assert req.resolved_model() is None
req = RealtimeTranscriptionSessionRequest(
model="openai/gpt-realtime-whisper",
input_audio_transcription={"model": "gpt-realtime-whisper"},
)
assert req.resolved_model() == "openai/gpt-realtime-whisper"
def test_request_resolves_model_from_input_audio_transcription():
req = RealtimeTranscriptionSessionRequest(
input_audio_transcription={"model": "gpt-realtime-whisper", "language": "en"},
)
assert req.resolved_model() == "gpt-realtime-whisper"
def test_request_passthrough_excludes_routing_hint():
"""Unknown fields pass through; the litellm-only `model` hint is not forwarded."""
req = RealtimeTranscriptionSessionRequest(
model="openai/gpt-realtime-whisper",
input_audio_format="pcm16",
input_audio_transcription={"model": "gpt-realtime-whisper"},
turn_detection=None,
)
forwarded = req.model_dump(exclude_none=True, exclude={"model"})
assert "model" not in forwarded
assert forwarded["input_audio_format"] == "pcm16"
assert forwarded["input_audio_transcription"] == {"model": "gpt-realtime-whisper"}
@pytest.mark.asyncio
async def test_handler_posts_to_transcription_sessions_url():
handler = BaseLLMHTTPHandler()
mock_response = MagicMock(spec=httpx.Response)
mock_client = MagicMock(spec=AsyncHTTPHandler)
mock_client.post = AsyncMock(return_value=mock_response)
logging_obj = MagicMock()
logging_obj.pre_call = MagicMock()
request_body = {"input_audio_transcription": {"model": "gpt-realtime-whisper"}}
result = await handler.async_realtime_transcription_session_handler(
api_base="https://api.openai.com",
api_key="sk-test",
request_data=request_body,
logging_obj=logging_obj,
timeout=10.0,
provider_config=OpenAIRealtimeHTTPConfig(),
model="gpt-realtime-whisper",
client=mock_client,
)
assert result is mock_response
_, kwargs = mock_client.post.call_args
assert kwargs["url"] == "https://api.openai.com/v1/realtime/transcription_sessions"
assert kwargs["json"] == request_body
assert kwargs["headers"]["Authorization"] == "Bearer sk-test"
@pytest.mark.asyncio
async def test_client_secret_handler_still_targets_client_secrets_url():
"""Refactor regression: the client_secrets handler must keep its own URL."""
handler = BaseLLMHTTPHandler()
mock_response = MagicMock(spec=httpx.Response)
mock_client = MagicMock(spec=AsyncHTTPHandler)
mock_client.post = AsyncMock(return_value=mock_response)
logging_obj = MagicMock()
logging_obj.pre_call = MagicMock()
await handler.async_realtime_client_secret_handler(
api_base="https://api.openai.com",
api_key="sk-test",
request_data={"session": {"type": "realtime"}},
logging_obj=logging_obj,
timeout=10.0,
provider_config=OpenAIRealtimeHTTPConfig(),
model="gpt-4o-realtime-preview",
client=mock_client,
)
_, kwargs = mock_client.post.call_args
assert kwargs["url"] == "https://api.openai.com/v1/realtime/client_secrets"
@pytest.mark.asyncio
async def test_sdk_fn_routes_openai_transcription_session(monkeypatch):
"""
litellm.acreate_realtime_transcription_session resolves the OpenAI provider
from the transcription model and POSTs to the OpenAI transcription_sessions URL.
"""
import litellm
monkeypatch.setenv("OPENAI_API_KEY", "sk-unit-test")
mock_response = MagicMock(spec=httpx.Response)
mock_client = MagicMock(spec=AsyncHTTPHandler)
mock_client.post = AsyncMock(return_value=mock_response)
result = await litellm.acreate_realtime_transcription_session(
model="openai/gpt-realtime-whisper",
transcription_session={
"input_audio_format": "pcm16",
"input_audio_transcription": {"model": "gpt-realtime-whisper"},
},
client=mock_client,
)
assert result is mock_response
_, kwargs = mock_client.post.call_args
assert kwargs["url"].endswith("/v1/realtime/transcription_sessions")
# The litellm-only routing hint must not be forwarded upstream.
assert "model" not in kwargs["json"]
assert kwargs["json"]["input_audio_transcription"] == {
"model": "gpt-realtime-whisper"
}
def test_append_query_params_skips_existing_keys():
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
url = "wss://example.com/v1/realtime?model=gpt-4o"
result = BaseLLMHTTPHandler._append_query_params(
url, {"model": "ignored", "intent": "transcription"}
)
assert "model=ignored" not in result
assert "intent=transcription" in result
def test_append_query_params_no_params_returns_unchanged():
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
url = "wss://example.com/v1/realtime?model=gpt-4o"
assert BaseLLMHTTPHandler._append_query_params(url, None) == url
assert BaseLLMHTTPHandler._append_query_params(url, {}) == url
def test_append_query_params_encodes_special_chars():
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
url = "wss://example.com/v1/realtime"
result = BaseLLMHTTPHandler._append_query_params(url, {"intent": "a&b=c"})
assert "intent=a%26b%3Dc" in result
assert "&b=c" not in result
def test_azure_construct_url_encodes_model_and_api_version():
"""model and api-version must be URL-encoded to prevent query-string injection."""
from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime
h = AzureOpenAIRealtime()
url = h._construct_url(
"https://x.openai.azure.com",
"deploy&evil=1",
"2024-10-01-preview",
)
assert "evil=1" not in url.split("?", 1)[1]
url_ga = h._construct_url(
"https://x.openai.azure.com",
"deploy&evil=1",
None,
realtime_protocol="GA",
)
assert "evil=1" not in url_ga.split("?", 1)[1]

View file

@ -0,0 +1,324 @@
"""
Tests for Parallel AI Search API integration (v1 endpoint).
"""
import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
sys.path.insert(0, os.path.abspath("../../../.."))
import litellm
MOCK_V1_RESPONSE = {
"search_id": "search_abc123",
"session_id": "session_xyz",
"results": [
{
"url": "https://example.com/1",
"title": "Test Result 1",
"publish_date": "2026-01-15",
"excerpts": ["First excerpt.", "Second excerpt."],
},
{
"url": "https://example.com/2",
"title": None,
"publish_date": None,
"excerpts": ["Only excerpt."],
},
],
"usage": [{"name": "search_advanced", "count": 1}],
}
def _mock_response():
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = MOCK_V1_RESPONSE
return mock_response
class TestParallelAISearch:
@pytest.fixture(autouse=True)
def _set_api_key(self, monkeypatch):
monkeypatch.setenv("PARALLEL_API_KEY", "test-api-key")
monkeypatch.delenv("PARALLEL_AI_API_BASE", raising=False)
@pytest.mark.asyncio
async def test_v1_endpoint_and_headers(self):
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = _mock_response()
await litellm.asearch(
query="latest developments in AI",
search_provider="parallel_ai",
)
call_args = mock_post.call_args
assert call_args.kwargs["url"] == "https://api.parallel.ai/v1/search"
headers = call_args.kwargs.get("headers", {})
assert headers["x-api-key"] == "test-api-key"
assert headers["Content-Type"] == "application/json"
assert "parallel-beta" not in headers
@pytest.mark.asyncio
async def test_string_query_maps_to_search_queries_and_objective(self):
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = _mock_response()
await litellm.asearch(
query="latest developments in AI",
search_provider="parallel_ai",
)
json_data = mock_post.call_args.kwargs.get("json")
assert json_data["search_queries"] == ["latest developments in AI"]
assert json_data["objective"] == "latest developments in AI"
@pytest.mark.asyncio
async def test_list_query_maps_to_search_queries(self):
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = _mock_response()
await litellm.asearch(
query=["AI developments", "machine learning trends"],
search_provider="parallel_ai",
)
json_data = mock_post.call_args.kwargs.get("json")
assert json_data["search_queries"] == [
"AI developments",
"machine learning trends",
]
assert "objective" not in json_data
@pytest.mark.asyncio
async def test_mode_param_passthrough(self):
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = _mock_response()
await litellm.asearch(
query="AI developments",
search_provider="parallel_ai",
mode="turbo",
)
json_data = mock_post.call_args.kwargs.get("json")
assert json_data["mode"] == "turbo"
@pytest.mark.asyncio
async def test_default_mode_is_basic(self):
"""v1 defaults to 'advanced' server-side; litellm must send 'basic' to keep v1beta's default tier and cost tracking accurate."""
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = _mock_response()
await litellm.asearch(
query="AI developments",
search_provider="parallel_ai",
)
json_data = mock_post.call_args.kwargs.get("json")
assert json_data["mode"] == "basic"
@pytest.mark.parametrize(
"processor,expected_mode", [("base", "basic"), ("pro", "advanced")]
)
@pytest.mark.asyncio
async def test_legacy_processor_maps_to_mode(self, processor, expected_mode):
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = _mock_response()
await litellm.asearch(
query="AI developments",
search_provider="parallel_ai",
processor=processor,
)
json_data = mock_post.call_args.kwargs.get("json")
assert json_data["mode"] == expected_mode
assert "processor" not in json_data
@pytest.mark.asyncio
async def test_explicit_mode_wins_over_processor(self):
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = _mock_response()
await litellm.asearch(
query="AI developments",
search_provider="parallel_ai",
mode="turbo",
processor="pro",
)
json_data = mock_post.call_args.kwargs.get("json")
assert json_data["mode"] == "turbo"
assert "processor" not in json_data
@pytest.mark.asyncio
async def test_top_level_v1_params_pass_through(self):
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = _mock_response()
await litellm.asearch(
query="AI developments",
search_provider="parallel_ai",
session_id="session_123",
max_chars_total=4000,
max_tokens_per_page=1024,
)
json_data = mock_post.call_args.kwargs.get("json")
assert json_data["session_id"] == "session_123"
assert json_data["max_chars_total"] == 4000
assert "max_tokens_per_page" not in json_data
@pytest.mark.asyncio
async def test_optional_params_nest_under_advanced_settings(self):
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = _mock_response()
await litellm.asearch(
query="AI developments",
search_provider="parallel_ai",
max_results=5,
country="US",
search_domain_filter=["arxiv.org", "nature.com"],
exclude_domains=["reddit.com"],
max_chars_per_result=1500,
)
json_data = mock_post.call_args.kwargs.get("json")
advanced_settings = json_data["advanced_settings"]
assert advanced_settings["max_results"] == 5
assert advanced_settings["location"] == "US"
assert advanced_settings["source_policy"]["include_domains"] == [
"arxiv.org",
"nature.com",
]
assert advanced_settings["source_policy"]["exclude_domains"] == [
"reddit.com"
]
assert advanced_settings["excerpt_settings"]["max_chars_per_result"] == 1500
assert "max_results" not in json_data
assert "source_policy" not in json_data
assert "search_domain_filter" not in json_data
assert "exclude_domains" not in json_data
assert "max_chars_per_result" not in json_data
assert "country" not in json_data
@pytest.mark.asyncio
async def test_explicit_advanced_settings_take_precedence(self):
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = _mock_response()
await litellm.asearch(
query="AI developments",
search_provider="parallel_ai",
max_results=5,
advanced_settings={"max_results": 7},
)
json_data = mock_post.call_args.kwargs.get("json")
assert json_data["advanced_settings"]["max_results"] == 7
@pytest.mark.asyncio
async def test_response_transformation(self):
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = _mock_response()
response = await litellm.asearch(
query="AI developments",
search_provider="parallel_ai",
)
assert response.object == "search"
assert len(response.results) == 2
first = response.results[0]
assert first.title == "Test Result 1"
assert first.url == "https://example.com/1"
assert first.snippet == "First excerpt. ... Second excerpt."
assert first.date == "2026-01-15"
second = response.results[1]
assert second.title == ""
assert second.snippet == "Only excerpt."
assert second.date is None
@pytest.mark.parametrize(
"api_base",
[
"https://proxy.internal.example.com",
"https://proxy.internal.example.com/",
"https://proxy.internal.example.com/v1",
"https://proxy.internal.example.com/v1/",
"https://proxy.internal.example.com/v1/search",
],
)
@pytest.mark.asyncio
async def test_custom_api_base_appends_v1_search(self, api_base):
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = _mock_response()
await litellm.asearch(
query="AI developments",
search_provider="parallel_ai",
api_base=api_base,
)
call_args = mock_post.call_args
assert (
call_args.kwargs["url"]
== "https://proxy.internal.example.com/v1/search"
)
@pytest.mark.asyncio
async def test_missing_api_key_raises(self, monkeypatch):
monkeypatch.delenv("PARALLEL_API_KEY", raising=False)
monkeypatch.delenv("PARALLEL_AI_API_KEY", raising=False)
with pytest.raises(Exception, match="PARALLEL_API_KEY"):
await litellm.asearch(
query="AI developments",
search_provider="parallel_ai",
)

View file

@ -302,3 +302,39 @@ def test_streaming_tool_call_finish_reason_with_empty_content_in_final_chunk():
assert len(response2.choices) == 1
# Must be "tool_calls", NOT "stop"
assert response2.choices[0].finish_reason == "tool_calls"
def test_streaming_metadata_only_chunk_does_not_yield_empty_choices():
"""
web_search + reasoning makes Gemini emit mid-stream chunks that carry only
grounding/thought metadata — no content part and no finishReason.
_process_candidates skips content-less candidates, so without a fallback
`choices` is empty and the downstream streaming handler hits
`IndexError: list index out of range` on choices[0].
Ref: https://github.com/BerriAI/litellm/issues/28884
"""
logging_obj = _make_logging_obj()
iterator = ModelResponseIterator(
streaming_response=iter([]),
sync_stream=True,
logging_obj=logging_obj,
)
# Grounding-only chunk: a candidate with groundingMetadata but no content
# part and no finishReason (what web_search + reasoning produces mid-stream).
metadata_only_chunk = {
"candidates": [
{
"index": 0,
"groundingMetadata": {"webSearchQueries": ["weather boston"]},
}
]
}
response = iterator.chunk_parser(metadata_only_chunk)
assert response is not None
# Must expose at least one choice so downstream choices[0] is safe.
assert len(response.choices) == 1
assert response.choices[0].finish_reason is None
assert response.choices[0].delta.content is None

View file

@ -88,6 +88,76 @@ async def test_fetch_tools_from_passthrough_raises_on_upstream_401():
mock_client.list_tools.assert_awaited_with(raise_on_error=True)
@pytest.mark.asyncio
async def test_fetch_tools_from_delegated_oauth2_raises_on_upstream_401():
manager = MCPServerManager()
delegated_server = MCPServer(
server_id="oauth1",
name="delegated_docs",
url="https://upstream/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=True,
)
response = httpx.Response(
status_code=401,
headers={"www-authenticate": 'Bearer resource_metadata="https://upstream"'},
request=httpx.Request("GET", "https://upstream/mcp"),
)
upstream_error = httpx.HTTPStatusError(
"401", request=response.request, response=response
)
mock_client = MagicMock()
mock_client.list_tools = AsyncMock(side_effect=upstream_error)
with pytest.raises(MCPUpstreamAuthError) as exc_info:
await manager._fetch_tools_with_timeout(
mock_client, delegated_server.name, server=delegated_server
)
assert exc_info.value.status_code == 401
assert exc_info.value.www_authenticate == (
'Bearer resource_metadata="https://upstream"'
)
assert exc_info.value.server_name == "delegated_docs"
mock_client.list_tools.assert_awaited_with(raise_on_error=True)
@pytest.mark.asyncio
async def test_fetch_tools_from_client_credentials_oauth2_keeps_swallow_behavior():
manager = MCPServerManager()
m2m_server = MCPServer(
server_id="oauth-m2m",
name="m2m_docs",
url="https://upstream/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=True,
oauth2_flow="client_credentials",
)
response = httpx.Response(
status_code=401,
headers={"www-authenticate": 'Bearer resource_metadata="https://upstream"'},
request=httpx.Request("GET", "https://upstream/mcp"),
)
upstream_error = httpx.HTTPStatusError(
"401", request=response.request, response=response
)
mock_client = MagicMock()
mock_client.list_tools = AsyncMock(side_effect=upstream_error)
tools = await manager._fetch_tools_with_timeout(
mock_client, m2m_server.name, server=m2m_server
)
assert tools == []
mock_client.list_tools.assert_awaited_with(raise_on_error=False)
@pytest.mark.asyncio
async def test_fetch_tools_from_passthrough_returns_tools_on_success():
manager = MCPServerManager()

View file

@ -5489,3 +5489,84 @@ async def test_create_mcp_client_sampling_enabled():
client = await manager._create_mcp_client(server=server)
assert client._sampling_callback is not None
@pytest.mark.asyncio
async def test_execute_mcp_tool_sets_model_in_model_call_details():
"""Regression test: MCP tools/call spend logs persisted with model="".
execute_mcp_tool set logging_obj.model only; the spend-log writer reads
model_call_details["model"], which stays None when function_setup builds
the logging object without a "model" kwarg.
"""
import uuid
from datetime import timezone
from litellm.proxy._experimental.mcp_server import server as mcp_module
from litellm.proxy._types import LitellmUserRoles
from litellm.utils import Rules, function_setup
user = UserAPIKeyAuth(
api_key="sk-user",
user_id="alice",
user_role=LitellmUserRoles.INTERNAL_USER.value,
)
fake_server = MagicMock()
fake_server.name = "openapi-petstore"
fake_server.is_byok = False
fake_server.auth_type = None
fake_server.mcp_info = None
fake_server.server_id = "srv-1"
fake_server.server_name = "openapi-petstore"
fake_tool = MagicMock()
fake_tool.name = "list_pets"
start_time = datetime.now(timezone.utc)
litellm_logging_obj, _ = function_setup(
original_function="call_mcp_tool",
rules_obj=Rules(),
start_time=start_time,
litellm_call_id=str(uuid.uuid4()),
name="list_pets",
arguments={"limit": 10},
)
assert litellm_logging_obj.model_call_details.get("model") is None
with (
patch.object(
mcp_module.global_mcp_server_manager,
"_get_mcp_server_from_tool_name",
return_value=fake_server,
),
patch.object(
mcp_module.global_mcp_server_manager,
"pre_call_tool_check",
new=AsyncMock(return_value={}),
),
patch.object(
mcp_module.global_mcp_tool_registry,
"get_tool",
return_value=fake_tool,
),
patch(
"litellm.proxy._experimental.mcp_server.server._handle_local_mcp_tool",
new=AsyncMock(return_value=[]),
),
patch(
"litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed",
return_value=True,
),
):
await mcp_module.execute_mcp_tool(
name="list_pets",
arguments={"limit": 10},
allowed_mcp_servers=[fake_server],
start_time=start_time,
user_api_key_auth=user,
litellm_logging_obj=litellm_logging_obj,
)
assert litellm_logging_obj.model_call_details["model"] == "MCP: list_pets"
assert litellm_logging_obj.model == "MCP: list_pets"

View file

@ -673,6 +673,110 @@ async def test_per_user_oauth_missing_stored_token_returns_preemptive_401():
assert "Bearer authorization_uri=" in exc_info.value.headers["www-authenticate"]
@pytest.mark.asyncio
async def test_handle_streamable_http_mcp_delegated_server_surfaces_upstream_challenge():
"""
OAuth2 server with ``delegate_auth_to_upstream=True`` should let the
upstream MCP server's RFC 9728 challenge reach the client instead of
pre-emptively returning LiteLLM's gateway authorization_uri challenge.
"""
from fastapi import HTTPException
try:
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPUpstreamAuthError,
)
from litellm.proxy._experimental.mcp_server.server import (
handle_streamable_http_mcp,
session_manager_stateful,
)
except ImportError:
pytest.skip("MCP server not available")
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/delegated_oauth_server",
"scheme": "https",
"query_string": b"",
"root_path": "",
"server": ("litellm.example.com", 443),
"headers": [
(b"content-type", b"application/json"),
(b"host", b"litellm.example.com"),
],
}
receive = AsyncMock(
return_value={
"type": "http.request",
"body": b'{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}',
"more_body": False,
}
)
send = AsyncMock()
user_auth = MagicMock()
user_auth.user_id = None
delegated_server = MagicMock()
delegated_server.auth_type = MCPAuth.oauth2
delegated_server.delegate_auth_to_upstream = True
delegated_server.needs_user_oauth_token = True
delegated_server.server_id = "delegated-oauth-server"
upstream_challenge = (
'Bearer resource_metadata="https://upstream.example.com/.well-known/oauth-protected-resource"'
)
with (
patch(
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
new_callable=AsyncMock,
return_value=(
user_auth,
None,
["delegated_oauth_server"],
None,
None,
None,
),
),
patch("litellm.proxy._experimental.mcp_server.server.set_auth_context"),
patch(
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
True,
),
patch(
"litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session",
new_callable=AsyncMock,
return_value=False,
),
patch(
"litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db",
new_callable=AsyncMock,
return_value=None,
),
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name",
return_value=delegated_server,
),
patch.object(
session_manager_stateful,
"handle_request",
new_callable=AsyncMock,
side_effect=MCPUpstreamAuthError(
status_code=401,
www_authenticate=upstream_challenge,
server_name="delegated_oauth_server",
),
) as mock_handle_request,
):
with pytest.raises(HTTPException) as exc_info:
await handle_streamable_http_mcp(scope, receive, send)
assert mock_handle_request.await_count == 1
assert exc_info.value.status_code == 401
assert exc_info.value.headers == {"www-authenticate": upstream_challenge}
@pytest.mark.asyncio
async def test_per_user_oauth_with_stored_token_skips_preemptive_401():
"""
@ -759,19 +863,16 @@ async def test_per_user_oauth_with_stored_token_skips_preemptive_401():
@pytest.mark.asyncio
async def test_handle_streamable_http_mcp_emits_401_for_delegated_server_without_token():
async def test_handle_streamable_http_mcp_delegated_server_without_token_reaches_session_manager():
"""
OAuth2 server with ``delegate_auth_to_upstream=True`` and no Authorization
header must still emit a pre-emptive 401 with WWW-Authenticate so the
client kicks off PKCE. The 401 points at LiteLLM's discovery shim, which
in turn delegates to the upstream OAuth issuer.
OAuth2 server with ``delegate_auth_to_upstream=True`` and no stored token
should not receive LiteLLM's gateway authorization_uri challenge. The
request continues so the upstream MCP server can emit its RFC 9728 challenge.
"""
from fastapi import HTTPException
try:
from litellm.proxy._experimental.mcp_server.server import (
handle_streamable_http_mcp,
session_manager,
session_manager_stateless,
)
except ImportError:
pytest.skip("MCP server not available")
@ -785,7 +886,13 @@ async def test_handle_streamable_http_mcp_emits_401_for_delegated_server_without
(b"host", b"litellm.example.com"),
],
}
receive = AsyncMock()
receive = AsyncMock(
return_value={
"type": "http.request",
"body": b'{"jsonrpc":"2.0","id":1,"method":"tools/list","params":{}}',
"more_body": False,
}
)
send = AsyncMock()
user_auth = MagicMock()
user_auth.user_id = None
@ -819,19 +926,22 @@ async def test_handle_streamable_http_mcp_emits_401_for_delegated_server_without
new_callable=AsyncMock,
return_value=False,
),
patch(
"litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db",
new_callable=AsyncMock,
return_value=None,
) as mock_get_stored_token,
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name",
return_value=delegated_server,
),
patch.object(
session_manager,
session_manager_stateless,
"handle_request",
new_callable=AsyncMock,
) as mock_handle_request,
):
with pytest.raises(HTTPException) as exc_info:
await handle_streamable_http_mcp(scope, receive, send)
await handle_streamable_http_mcp(scope, receive, send)
assert exc_info.value.status_code == 401
assert "www-authenticate" in exc_info.value.headers
assert mock_handle_request.await_count == 0
assert mock_get_stored_token.await_count == 1
assert mock_handle_request.await_count == 1

View file

@ -2503,3 +2503,267 @@ async def test_post_call_success_hook_only_runs_output_scan():
mock_make.call_args.kwargs.get("logging_event_type")
== GuardrailEventHooks.post_call
)
# ---------------------------------------------------------------------------
# Contextual grounding: request-side qualifiers
# ---------------------------------------------------------------------------
#
# Bedrock contextual grounding tags each ApplyGuardrail content block with a
# `qualifiers` array (grounding_source / query / guard_content). A caller marks
# message content blocks `{"type": "grounding_source", ...}` / `{"type": "query", ...}`;
# at post_call the hook assembles one source="OUTPUT" call carrying the source +
# query + the model response (as guard_content). A request without these tags
# produces the plain-text payload with no qualifiers.
_GROUNDING_SOURCE_TEXT = "Tokyo is the capital of Japan."
_GROUNDING_QUERY_TEXT = "What is the capital of Japan?"
_GROUNDING_RESPONSE_TEXT = "The capital of Japan is Tokyo."
def _grounding_guardrail() -> BedrockGuardrail:
return BedrockGuardrail(
guardrailIdentifier="test-guardrail", guardrailVersion="DRAFT"
)
def _grounding_messages() -> list:
return [
{
"role": "system",
"content": [{"type": "grounding_source", "text": _GROUNDING_SOURCE_TEXT}],
},
{
"role": "user",
"content": [{"type": "query", "text": _GROUNDING_QUERY_TEXT}],
},
]
def _model_response(content: str) -> ModelResponse:
from litellm.types.utils import Choices, Message, ModelResponse
return ModelResponse(
choices=[
Choices(
index=0,
message=Message(role="assistant", content=content),
finish_reason="stop",
)
]
)
# Expected OUTPUT content blocks, keyed by their grounding qualifier, so the
# per-test assertions read as the block sequence they expect.
_GROUNDING_SOURCE_BLOCK = {
"text": {"text": _GROUNDING_SOURCE_TEXT, "qualifiers": ["grounding_source"]}
}
_QUERY_BLOCK = {"text": {"text": _GROUNDING_QUERY_TEXT, "qualifiers": ["query"]}}
_GUARD_BLOCK = {
"text": {"text": _GROUNDING_RESPONSE_TEXT, "qualifiers": ["guard_content"]}
}
def _input_request(messages: list) -> dict:
"""Arrange a guardrail and act: build the Bedrock INPUT payload."""
return _grounding_guardrail().convert_to_bedrock_format(
source="INPUT", messages=messages
)
def _output_request(messages: list, response=None) -> dict:
"""Arrange a guardrail and act: build the Bedrock OUTPUT payload."""
return _grounding_guardrail().convert_to_bedrock_format(
source="OUTPUT", response=response, messages=messages
)
def test_grounding_input_strips_grounding_and_query_qualifiers():
"""Grounding is OUTPUT-only: tagged source/query reach Bedrock as plain text on an
INPUT scan, so a tag cannot change how input-safety policies scan content (no bypass).
"""
expected_request = {
"source": "INPUT",
"content": [
{"text": {"text": _GROUNDING_SOURCE_TEXT}},
{"text": {"text": _GROUNDING_QUERY_TEXT}},
],
}
actual_request = _input_request(_grounding_messages())
assert actual_request == expected_request
def test_grounding_input_leaves_existing_guarded_text_unqualified():
"""An existing guarded_text input block keeps its legacy unqualified payload."""
expected_request = {"source": "INPUT", "content": [{"text": {"text": "policy"}}]}
actual_request = _input_request(
[{"role": "user", "content": [{"type": "guarded_text", "text": "policy"}]}]
)
assert actual_request == expected_request
def test_grounding_output_assembles_source_query_and_response():
"""OUTPUT emits grounding_source + query (from the request) then the response as
guard_content, so Bedrock can grade the response against the source and query."""
expected_request = {
"source": "OUTPUT",
"content": [_GROUNDING_SOURCE_BLOCK, _QUERY_BLOCK, _GUARD_BLOCK],
}
actual_request = _output_request(
_grounding_messages(), _model_response(_GROUNDING_RESPONSE_TEXT)
)
assert actual_request == expected_request
def test_grounding_output_keeps_legacy_payload_without_tags():
"""Without grounding tags the OUTPUT payload is the legacy single response block."""
expected_request = {
"source": "OUTPUT",
"content": [{"text": {"text": "Hi there."}}],
}
actual_request = _output_request(
[{"role": "user", "content": "hello"}], _model_response("Hi there.")
)
assert actual_request == expected_request
def test_grounding_output_combines_multiple_sources():
"""Every grounding_source block is emitted; Bedrock combines them into one corpus."""
uk_source_text = "London is the capital of UK."
uk_source_block = {
"text": {"text": uk_source_text, "qualifiers": ["grounding_source"]}
}
messages = [
{
"role": "system",
"content": [
{"type": "grounding_source", "text": uk_source_text},
{"type": "grounding_source", "text": _GROUNDING_SOURCE_TEXT},
],
},
{"role": "user", "content": [{"type": "query", "text": _GROUNDING_QUERY_TEXT}]},
]
expected_request = {
"source": "OUTPUT",
"content": [
uk_source_block,
_GROUNDING_SOURCE_BLOCK,
_QUERY_BLOCK,
_GUARD_BLOCK,
],
}
actual_request = _output_request(
messages, _model_response(_GROUNDING_RESPONSE_TEXT)
)
assert actual_request == expected_request
def test_grounding_output_keeps_grounding_for_non_model_response():
"""Harvested grounding blocks survive a non-ModelResponse output instead of being
silently dropped (regression guard for the unconditional content assignment)."""
expected_request = {
"source": "OUTPUT",
"content": [_GROUNDING_SOURCE_BLOCK, _QUERY_BLOCK],
}
actual_request = _output_request(_grounding_messages(), response=None)
assert actual_request == expected_request
@pytest.mark.parametrize(
"role, is_trusted",
[
("system", True),
("developer", True),
("tool", False),
("function", False),
("user", False),
("assistant", False),
],
)
def test_grounding_source_trusted_only_from_app_roles(role, is_trusted):
"""grounding_source is honored only from app-authored roles (system/developer). A
tag on a user, tool, function or assistant message is ignored, so neither a forwarded
end user nor an externally-influenced tool result can supply fake evidence for the
grounding check to grade the response against; query is always collected."""
messages = [
{
"role": role,
"content": [{"type": "grounding_source", "text": _GROUNDING_SOURCE_TEXT}],
},
{"role": "user", "content": [{"type": "query", "text": _GROUNDING_QUERY_TEXT}]},
]
expected_content = [_QUERY_BLOCK, _GUARD_BLOCK]
if is_trusted:
expected_content = [_GROUNDING_SOURCE_BLOCK, *expected_content]
actual_request = _output_request(
messages, _model_response(_GROUNDING_RESPONSE_TEXT)
)
assert actual_request == {"source": "OUTPUT", "content": expected_content}
@pytest.mark.asyncio
async def test_grounding_output_blocked_raises_400():
"""A BLOCKED contextualGroundingPolicy filter raises HTTP 400."""
guardrail = _grounding_guardrail()
mock_bedrock_response = MagicMock()
mock_bedrock_response.status_code = 200
mock_bedrock_response.json.return_value = {
"action": "GUARDRAIL_INTERVENED",
"assessments": [
{
"contextualGroundingPolicy": {
"filters": [
{
"type": "GROUNDING",
"threshold": 0.7,
"score": 0.1,
"action": "BLOCKED",
}
]
}
}
],
"outputs": [{"text": "Response blocked: not grounded in the provided source."}],
}
mock_credentials = MagicMock()
mock_credentials.access_key = "test-access-key"
mock_credentials.secret_key = "test-secret-key"
mock_credentials.token = None
with (
patch.object(
guardrail.async_handler, "post", new_callable=AsyncMock
) as mock_post,
patch.object(
guardrail, "_load_credentials", return_value=(mock_credentials, "us-east-1")
),
patch.object(guardrail, "_prepare_request", return_value=MagicMock()),
):
mock_post.return_value = mock_bedrock_response
with pytest.raises(HTTPException) as exc_info:
await guardrail.make_bedrock_api_request(
source="OUTPUT",
response=_model_response("The capital of Japan is Paris."),
messages=_grounding_messages(),
request_data={"messages": _grounding_messages()},
)
assert exc_info.value.status_code == 400

View file

@ -0,0 +1,653 @@
"""
Unit tests for Ovalix guardrail: config resolution and apply_guardrail behavior
with mocked Tracker service responses (allow, anonymize, block).
"""
import os
from typing import Any, List
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from litellm.exceptions import GuardrailRaisedException
from litellm.proxy.guardrails.guardrail_hooks.ovalix.ovalix import (
OvalixGuardrail,
OvalixGuardrailBlockedException,
OvalixGuardrailMissingSecrets,
)
from litellm.types.utils import GenericGuardrailAPIInputs
# Example Tracker responses (as returned by the checkpoint API)
TRACKER_RESPONSE_ALLOW = {
"action_type": "allow",
"data_type": "TEXT",
"original_data": {"content": "how are you?"},
"modified_data": {"content": "how are you?"},
"alerts": [],
}
TRACKER_RESPONSE_ANONYMIZE = {
"action_type": "anonymize",
"data_type": "TEXT",
"original_data": {"content": "Hello, my name is David."},
"modified_data": {"content": "Hello, my name is {Name}. How are you?"},
"alerts": [
{
"title": "Sensitive Data Alert",
"subtitle": "We've identified that you were trying to share sensitive information",
"alerts": ["Name:\tDavid\nRedacted to:\t{Name}"],
}
],
}
TRACKER_RESPONSE_BLOCK = {
"action_type": "block",
"data_type": "TEXT",
"original_data": {"content": "I am 15 YO"},
"modified_data": {"content": "This message was blocked by Ovalix"},
"alerts": [
{
"title": "Sensitive Data Alert",
"subtitle": "We've identified that you were trying to share sensitive information",
"alerts": ["Age:\t15\nBlocked"],
}
],
}
def _ovalix_env():
return {
"OVALIX_TRACKER_API_BASE": "https://tracker.test",
"OVALIX_TRACKER_API_KEY": "key",
"OVALIX_APPLICATION_ID": "app-1",
"OVALIX_PRE_CHECKPOINT_ID": "pre-1",
"OVALIX_POST_CHECKPOINT_ID": "post-1",
}
def _guardrail_kwargs():
return {
"guardrail_name": "ovalix-test",
"event_hook": "pre_call",
"default_on": True,
}
class TestOvalixGuardrailConfigModel:
"""Minimal config model tests: wiring only."""
def test_get_config_model_returns_ovalix_config_model(self):
"""get_config_model returns OvalixGuardrailConfigModel for proxy/config wiring."""
config_model = OvalixGuardrail.get_config_model()
assert config_model is not None
assert config_model.__name__ == "OvalixGuardrailConfigModel"
assert config_model.ui_friendly_name() == "Ovalix Guardrail"
class TestOvalixGuardrail:
"""Behavioral tests with mocked Tracker checkpoint API."""
def setup_method(self):
for key in list(os.environ.keys()):
if key.startswith("OVALIX_"):
del os.environ[key]
def teardown_method(self):
for key in list(os.environ.keys()):
if key.startswith("OVALIX_"):
del os.environ[key]
@pytest.fixture
def guardrail_with_env(self):
"""Guardrail with OVALIX_* env set; cleans up in teardown."""
for k, v in _ovalix_env().items():
os.environ[k] = v
try:
yield OvalixGuardrail(**_guardrail_kwargs())
finally:
for k in _ovalix_env():
if k in os.environ:
del os.environ[k]
def test_initialization_requires_secrets(self):
"""Initialization raises when required Tracker/application/checkpoint config is missing."""
with pytest.raises(OvalixGuardrailMissingSecrets):
OvalixGuardrail(
guardrail_name="ovalix-test",
event_hook="pre_call",
default_on=True,
)
def test_initialization_with_explicit_params(self):
"""Guardrail initializes with explicit tracker base, key, app and checkpoint IDs."""
guardrail = OvalixGuardrail(
tracker_api_base="https://tracker.example",
tracker_api_key="secret",
application_id="app-x",
pre_checkpoint_id="pre-x",
post_checkpoint_id="post-x",
**_guardrail_kwargs(),
)
assert guardrail._tracker_api_base == "https://tracker.example"
assert guardrail._application_id == "app-x"
assert guardrail._pre_checkpoint_id == "pre-x"
assert guardrail._post_checkpoint_id == "post-x"
def test_initialization_with_env_vars(self):
"""Guardrail picks up OVALIX_* env vars when params not passed."""
for k, v in _ovalix_env().items():
os.environ[k] = v
try:
guardrail = OvalixGuardrail(**_guardrail_kwargs())
assert guardrail._tracker_api_base == "https://tracker.test"
assert guardrail._tracker_api_key == "key"
assert guardrail._application_id == "app-1"
assert guardrail._pre_checkpoint_id == "pre-1"
assert guardrail._post_checkpoint_id == "post-1"
finally:
for k in _ovalix_env():
if k in os.environ:
del os.environ[k]
@pytest.mark.asyncio
async def test_call_checkpoint_sends_correct_payload_and_returns_json(self):
"""_call_checkpoint POSTs to tracker with application_id, checkpoint_id, actor, session_id, data."""
for k, v in _ovalix_env().items():
os.environ[k] = v
try:
guardrail = OvalixGuardrail(**_guardrail_kwargs())
mock_response = MagicMock()
mock_response.json.return_value = TRACKER_RESPONSE_ALLOW
mock_response.raise_for_status = MagicMock()
with patch.object(
guardrail._async_handler, "post", new_callable=AsyncMock
) as mock_post:
mock_post.return_value = mock_response
result = await guardrail._call_checkpoint(
content="hello",
checkpoint_id="pre-1",
actor="a1b2c3d4",
session_id="session-1",
)
assert result == TRACKER_RESPONSE_ALLOW
mock_post.assert_called_once()
call_args = mock_post.call_args
assert call_args.args[0] == (
"https://tracker.test/tracking/custom_application/checkpoint"
)
body = call_args.kwargs["json"]
assert body["application_id"] == "app-1"
assert body["checkpoint_id"] == "pre-1"
assert body["actor"] == "a1b2c3d4"
assert body["session_id"] == "session-1"
assert body["data_type"] == "TEXT"
assert body["data"] == {"content": "hello"}
finally:
for k in _ovalix_env():
if k in os.environ:
del os.environ[k]
@pytest.mark.asyncio
async def test_apply_guardrail_request_allow_passes_through(self):
"""When Tracker returns allow, apply_guardrail returns inputs with texts set to modified_data content."""
for k, v in _ovalix_env().items():
os.environ[k] = v
try:
guardrail = OvalixGuardrail(**_guardrail_kwargs())
inputs = GenericGuardrailAPIInputs(
structured_messages=[{"role": "user", "content": "how are you?"}],
texts=["how are you?"],
)
request_data = {}
mock_response = MagicMock()
mock_response.json.return_value = TRACKER_RESPONSE_ALLOW
mock_response.raise_for_status = MagicMock()
with patch.object(
guardrail._async_handler, "post", new_callable=AsyncMock
) as mock_post:
mock_post.return_value = mock_response
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="request",
logging_obj=None,
)
assert result.get("texts") == ["how are you?"]
assert mock_post.call_count == 1
finally:
for k in _ovalix_env():
if k in os.environ:
del os.environ[k]
@pytest.mark.asyncio
async def test_apply_guardrail_request_anonymize_returns_modified_text(self):
"""When Tracker returns anonymize, apply_guardrail returns texts with modified_data content."""
for k, v in _ovalix_env().items():
os.environ[k] = v
try:
guardrail = OvalixGuardrail(**_guardrail_kwargs())
inputs = GenericGuardrailAPIInputs(
structured_messages=[
{"role": "user", "content": "Hello, my name is David."}
],
texts=["Hello, my name is David."],
)
request_data = {}
mock_response = MagicMock()
mock_response.json.return_value = TRACKER_RESPONSE_ANONYMIZE
mock_response.raise_for_status = MagicMock()
with patch.object(
guardrail._async_handler, "post", new_callable=AsyncMock
) as mock_post:
mock_post.return_value = mock_response
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="request",
logging_obj=None,
)
assert result.get("texts") == ["Hello, my name is {Name}. How are you?"]
assert mock_post.call_count == 1
finally:
for k in _ovalix_env():
if k in os.environ:
del os.environ[k]
@pytest.mark.asyncio
async def test_apply_guardrail_request_block_raises_with_tracker_message(self):
"""When Tracker returns block on the (chronologically) last user message, OvalixGuardrailBlockedException is raised."""
for k, v in _ovalix_env().items():
os.environ[k] = v
try:
guardrail = OvalixGuardrail(**_guardrail_kwargs())
inputs = GenericGuardrailAPIInputs(
structured_messages=[{"role": "user", "content": "I am 15 YO"}],
texts=["I am 15 YO"],
)
request_data = {}
mock_response = MagicMock()
mock_response.json.return_value = TRACKER_RESPONSE_BLOCK
mock_response.raise_for_status = MagicMock()
with patch.object(
guardrail._async_handler, "post", new_callable=AsyncMock
) as mock_post:
mock_post.return_value = mock_response
with pytest.raises(OvalixGuardrailBlockedException) as exc_info:
await guardrail.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="request",
logging_obj=None,
)
assert "This message was blocked by Ovalix" in str(exc_info.value.message)
assert exc_info.value.status_code == 400
assert mock_post.call_count == 1
finally:
for k in _ovalix_env():
if k in os.environ:
del os.environ[k]
@pytest.mark.asyncio
async def test_apply_guardrail_request_block_non_last_replaced_in_texts(self):
"""When Tracker returns block on a non-last user message, that message is replaced in texts and no exception is raised."""
for k, v in _ovalix_env().items():
os.environ[k] = v
try:
guardrail = OvalixGuardrail(**_guardrail_kwargs())
inputs = GenericGuardrailAPIInputs(
structured_messages=[
{"role": "user", "content": "I am 15 YO"},
{"role": "user", "content": "how are you?"},
],
texts=["I am 15 YO", "how are you?"],
)
request_data = {}
def side_effect(*args, **kwargs):
body = kwargs.get("json", {})
content = (body.get("data") or {}).get("content", "")
resp = MagicMock()
if "15" in content:
resp.json.return_value = TRACKER_RESPONSE_BLOCK
else:
resp.json.return_value = TRACKER_RESPONSE_ALLOW
resp.raise_for_status = MagicMock()
return resp
with patch.object(
guardrail._async_handler, "post", new_callable=AsyncMock
) as mock_post:
mock_post.side_effect = side_effect
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="request",
logging_obj=None,
)
assert result.get("texts") == [
"This message was blocked by Ovalix",
"how are you?",
]
assert mock_post.call_count == 2
finally:
for k in _ovalix_env():
if k in os.environ:
del os.environ[k]
@pytest.mark.asyncio
async def test_apply_guardrail_response_allow_returns_inputs(self):
"""When input_type is response and Tracker allows, apply_guardrail returns inputs with texts updated from Tracker."""
for k, v in _ovalix_env().items():
os.environ[k] = v
try:
guardrail = OvalixGuardrail(**_guardrail_kwargs())
inputs = GenericGuardrailAPIInputs(
structured_messages=[
{"role": "assistant", "content": "Safe assistant reply"}
],
texts=["Safe assistant reply"],
)
request_data = {}
mock_response = MagicMock()
mock_response.json.return_value = TRACKER_RESPONSE_ALLOW
mock_response.raise_for_status = MagicMock()
with patch.object(
guardrail._async_handler, "post", new_callable=AsyncMock
) as mock_post:
mock_post.return_value = mock_response
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="response",
logging_obj=None,
)
assert result.get("texts") == ["how are you?"]
assert mock_post.call_count == 1
finally:
for k in _ovalix_env():
if k in os.environ:
del os.environ[k]
@pytest.mark.asyncio
async def test_apply_guardrail_response_block_raises(self, guardrail_with_env):
"""When Tracker blocks on response, apply_guardrail raises OvalixGuardrailBlockedException."""
guardrail = guardrail_with_env
inputs = GenericGuardrailAPIInputs(
structured_messages=[{"role": "user", "content": "I am 15 YO"}],
texts=["I am 15 YO"],
)
request_data = {}
mock_response = MagicMock()
mock_response.json.return_value = TRACKER_RESPONSE_BLOCK
mock_response.raise_for_status = MagicMock()
with patch.object(
guardrail._async_handler, "post", new_callable=AsyncMock
) as mock_post:
mock_post.return_value = mock_response
with pytest.raises(OvalixGuardrailBlockedException) as exc_info:
await guardrail.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="response",
logging_obj=None,
)
assert "This message was blocked by Ovalix" in str(exc_info.value.message)
assert exc_info.value.status_code == 400
assert mock_post.call_count == 1
@pytest.mark.asyncio
async def test_apply_guardrail_request_missing_modified_data_uses_original_content(
self, guardrail_with_env
):
"""When Tracker response has no modified_data.content, original content is used."""
guardrail = guardrail_with_env
inputs = GenericGuardrailAPIInputs(
structured_messages=[{"role": "user", "content": "original text"}],
texts=["original text"],
)
request_data = {}
mock_response = MagicMock()
mock_response.json.return_value = {
"action_type": "allow",
"data_type": "TEXT",
"original_data": {"content": "original text"},
"modified_data": {},
"alerts": [],
}
mock_response.raise_for_status = MagicMock()
with patch.object(
guardrail._async_handler, "post", new_callable=AsyncMock
) as mock_post:
mock_post.return_value = mock_response
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="request",
logging_obj=None,
)
assert result.get("texts") == ["original text"]
assert mock_post.call_count == 1
@pytest.mark.asyncio
async def test_apply_guardrail_tracker_http_error_raises_guardrail_exception(
self, guardrail_with_env
):
"""When Tracker returns HTTP error (e.g. 400), GuardrailRaisedException is raised."""
guardrail = guardrail_with_env
inputs = GenericGuardrailAPIInputs(
structured_messages=[{"role": "user", "content": "hello"}],
texts=["hello"],
)
request_data = {}
mock_response = MagicMock()
mock_response.raise_for_status.side_effect = httpx.HTTPStatusError(
"Bad Request",
request=MagicMock(),
response=MagicMock(status_code=400),
)
with patch.object(
guardrail._async_handler, "post", new_callable=AsyncMock
) as mock_post:
mock_post.return_value = mock_response
with pytest.raises(GuardrailRaisedException):
await guardrail.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="request",
logging_obj=None,
)
assert mock_post.call_count == 1
@pytest.mark.asyncio
async def test_apply_guardrail_checkpoint_error_raises_guardrail_exception(self):
"""When Tracker checkpoint call fails, GuardrailRaisedException is raised."""
for k, v in _ovalix_env().items():
os.environ[k] = v
try:
guardrail = OvalixGuardrail(**_guardrail_kwargs())
inputs = GenericGuardrailAPIInputs(
structured_messages=[{"role": "user", "content": "hello"}],
texts=["hello"],
)
request_data = {}
with patch.object(
guardrail._async_handler,
"post",
new_callable=AsyncMock,
side_effect=httpx.ConnectError("Connection refused"),
):
with pytest.raises(GuardrailRaisedException):
await guardrail.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="request",
logging_obj=None,
)
finally:
for k in _ovalix_env():
if k in os.environ:
del os.environ[k]
@pytest.mark.asyncio
async def test_apply_guardrail_request_empty_messages_returns_inputs(self):
"""When request has no messages, apply_guardrail returns inputs without calling Tracker."""
for k, v in _ovalix_env().items():
os.environ[k] = v
try:
guardrail = OvalixGuardrail(**_guardrail_kwargs())
inputs = GenericGuardrailAPIInputs(structured_messages=[], texts=[])
request_data = {}
with patch.object(
guardrail._async_handler, "post", new_callable=AsyncMock
) as mock_post:
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="request",
logging_obj=None,
)
assert result == inputs
mock_post.assert_not_called()
finally:
for k in _ovalix_env():
if k in os.environ:
del os.environ[k]
def test_get_actor_from_metadata(self):
"""Actor is taken from metadata.user_api_key_user_email or user_api_key_user_id."""
for k, v in _ovalix_env().items():
os.environ[k] = v
try:
guardrail = OvalixGuardrail(**_guardrail_kwargs())
assert (
guardrail._get_actor(
{"metadata": {"user_api_key_user_email": "a@b.com"}}
)
== "a@b.com"
)
assert (
guardrail._get_actor({"metadata": {"user_api_key_user_id": "uid-1"}})
== "uid-1"
)
assert (
guardrail._get_actor(
{"litellm_metadata": {"user_api_key_user_id": "uid-2"}}
)
== "uid-2"
)
assert guardrail._get_actor({}) == "unknown"
finally:
for k in _ovalix_env():
if k in os.environ:
del os.environ[k]
def test_get_actor_prefers_email_over_id(self, guardrail_with_env):
"""When both user_api_key_user_email and user_api_key_user_id exist, email is used."""
guardrail = guardrail_with_env
data = {
"metadata": {
"user_api_key_user_email": "primary@test.com",
"user_api_key_user_id": "uid-99",
}
}
assert guardrail._get_actor(data) == "primary@test.com"
def test_get_tracker_actor_id_is_hash_not_raw_pii(self, guardrail_with_env):
"""Tracker API actor field uses a short hash of _get_actor, not email/user id."""
guardrail = guardrail_with_env
data = {"metadata": {"user_api_key_user_email": "user@example.com"}}
raw = guardrail._get_actor(data)
hashed = guardrail._get_tracker_actor_id(data)
assert raw == "user@example.com"
assert hashed != raw
assert len(hashed) == 8
assert all(c in "0123456789abcdef" for c in hashed)
def test_get_session_id_deterministic_and_includes_app_id(self, guardrail_with_env):
"""Session ID is stable for same actor/day and includes application_id."""
guardrail = guardrail_with_env
data = {"metadata": {"user_api_key_user_id": "user-1"}}
session_id_1 = guardrail._get_session_id(data)
session_id_2 = guardrail._get_session_id(data)
assert session_id_1 == session_id_2
assert "app-1" in session_id_1
def test_block_current_message_raises_ovalix_blocked_exception(
self, guardrail_with_env
):
"""_block_current_message raises OvalixGuardrailBlockedException with status_code 400."""
guardrail = guardrail_with_env
with pytest.raises(OvalixGuardrailBlockedException) as exc_info:
guardrail._block_current_message("Custom block reason")
assert "Custom block reason" in str(exc_info.value.message)
assert exc_info.value.status_code == 400
def test_get_trackers_corrected_message(self, guardrail_with_env):
"""_get_trackers_corrected_message returns modified_data.content or None."""
guardrail = guardrail_with_env
assert (
guardrail._get_trackers_corrected_message(
{"modified_data": {"content": "corrected text"}}
)
== "corrected text"
)
assert guardrail._get_trackers_corrected_message({"modified_data": {}}) is None
assert (
guardrail._get_trackers_corrected_message({"modified_data": "not-a-dict"})
is None
)
@pytest.mark.asyncio
async def test_apply_guardrail_response_no_texts_returns_unchanged(self):
"""When input_type is response and inputs have no texts, apply_guardrail returns inputs without calling Tracker."""
for k, v in _ovalix_env().items():
os.environ[k] = v
try:
guardrail = OvalixGuardrail(**_guardrail_kwargs())
inputs = GenericGuardrailAPIInputs()
request_data = {}
with patch.object(
guardrail._async_handler, "post", new_callable=AsyncMock
) as mock_post:
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="response",
logging_obj=None,
)
assert result == inputs
mock_post.assert_not_called()
finally:
for k in _ovalix_env():
if k in os.environ:
del os.environ[k]

View file

@ -756,6 +756,85 @@ async def test_health_services_endpoint_rejects_unknown_service():
await health_services_endpoint(service="totally_unknown_service_xyz")
@pytest.mark.asyncio
@pytest.mark.parametrize(
"role",
[
None,
LitellmUserRoles.INTERNAL_USER,
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
LitellmUserRoles.TEAM,
LitellmUserRoles.CUSTOMER,
],
)
async def test_health_services_endpoint_newrelic_blocks_non_admin(role):
"""
/health/services?service=newrelic emits a real LiteLLMConnectionTest event
to the configured New Relic account. Only proxy admins (full or view-only)
should be able to trigger it; every other caller must be rejected before
the external event is recorded.
"""
from litellm.proxy._types import ProxyException
user_api_key_dict = UserAPIKeyAuth(
token="non-admin-token",
user_id="non-admin-user",
user_role=role,
)
with patch(
"litellm.integrations.newrelic.newrelic.NewRelicLogger"
) as MockNewRelicLogger:
mock_instance = MagicMock()
mock_instance.async_health_check = AsyncMock(
return_value={"status": "healthy", "error_message": ""}
)
MockNewRelicLogger.return_value = mock_instance
with pytest.raises(ProxyException) as exc_info:
await health_services_endpoint(
user_api_key_dict=user_api_key_dict,
service="newrelic",
)
assert str(exc_info.value.code) == "403"
mock_instance.async_health_check.assert_not_awaited()
MockNewRelicLogger.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"admin_role",
[LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY],
)
async def test_health_services_endpoint_newrelic_allows_proxy_admin(admin_role):
"""
Proxy admins (full and view-only) can trigger the New Relic test event.
"""
user_api_key_dict = UserAPIKeyAuth(
token="admin-token",
user_id="admin-user",
user_role=admin_role,
)
with patch(
"litellm.integrations.newrelic.newrelic.NewRelicLogger"
) as MockNewRelicLogger:
mock_instance = MagicMock()
mock_instance.async_health_check = AsyncMock(
return_value={"status": "healthy", "error_message": ""}
)
MockNewRelicLogger.return_value = mock_instance
result = await health_services_endpoint(
user_api_key_dict=user_api_key_dict,
service="newrelic",
)
assert result["status"] == "healthy"
mock_instance.async_health_check.assert_awaited_once()
@pytest.fixture(scope="function")
def proxy_client(monkeypatch):
"""

View file

@ -6,6 +6,7 @@ import asyncio
import os
import sys
import time
from contextlib import contextmanager
from datetime import datetime, timedelta
from typing import Any, Dict, List, Optional
@ -20,6 +21,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler,
)
from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.utils import (
EmbeddingResponse,
ModelResponse,
@ -3189,6 +3191,246 @@ def test_get_key_mcp_rpm_limit_precedence():
assert get_team_mcp_rpm_limit(none_set) is None
async def _seed_max_parallel_requests_counter(
dual_cache: DualCache, counter_key: str, window_size: int
) -> None:
await dual_cache.async_increment_cache_pipeline(
increment_list=[
RedisPipelineIncrementOperation(
key=counter_key, increment_value=1, ttl=window_size
)
]
)
async def _build_seeded_limiter():
"""Build a v3 limiter whose api-key counter already holds the pre-call +1."""
api_key = hash_token("sk-disconnect")
cache = DualCache()
limiter = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(cache)
)
counter_key = f"{{api_key:{api_key}}}:max_parallel_requests"
await _seed_max_parallel_requests_counter(cache, counter_key, limiter.window_size)
user_api_key_dict = UserAPIKeyAuth(api_key=api_key, max_parallel_requests=2)
return limiter, cache, counter_key, user_api_key_dict
@contextmanager
def _override_litellm_callbacks(new_callbacks):
"""Swap litellm.callbacks so _callback_capabilities recomputes deterministically."""
saved = litellm.callbacks
litellm.callbacks = new_callbacks
try:
yield
finally:
litellm.callbacks = saved
async def _drain_release_task():
# The disconnect release is scheduled fire-and-forget via create_task.
for _ in range(5):
await asyncio.sleep(0)
@pytest.mark.asyncio
async def test_release_max_parallel_requests_on_disconnect_v3():
"""
Regression for issue #27955: a stream cancelled mid-flight must release the
pre-call +1 reservation. The success/failure logging callbacks never fire
on cancellation, so without an explicit release the api-key counter climbs
by one per cancelled request until the key wedges at its limit. The release
must decrement the api-key max_parallel_requests counter by exactly one.
"""
_api_key = hash_token("sk-12345")
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=2)
counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests"
await _seed_max_parallel_requests_counter(
local_cache, counter_key, handler.window_size
)
assert await local_cache.async_get_cache(key=counter_key) == 1
await handler.async_release_max_parallel_requests_on_disconnect(user_api_key_dict)
assert await local_cache.async_get_cache(key=counter_key) == 0
@pytest.mark.asyncio
async def test_release_max_parallel_requests_on_disconnect_noop_v3():
"""
The release must be a no-op when the key never reserved a parallel slot
(no api_key, or max_parallel_requests unset). Otherwise a cancelled
no-limit request would drive an unrelated counter negative.
"""
_api_key = hash_token("sk-12345")
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests"
await handler.async_release_max_parallel_requests_on_disconnect(
UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=None)
)
assert await local_cache.async_get_cache(key=counter_key) is None
await handler.async_release_max_parallel_requests_on_disconnect(
UserAPIKeyAuth(api_key=None, max_parallel_requests=5)
)
assert await local_cache.async_get_cache(key=counter_key) is None
@pytest.mark.parametrize("disconnect", ["cancel", "aclose"])
@pytest.mark.asyncio
async def test_async_streaming_data_generator_releases_counter_on_disconnect_v3(
disconnect,
):
"""
Regression for issue #27955 on the outer SSE generator (used by /v1/messages
and other event-stream routes). A client that disconnects mid-stream raises
GeneratorExit (aclose) or CancelledError into async_streaming_data_generator;
both are BaseException and bypass the success/failure logging callbacks, so
the generator itself must refund the pre-call max_parallel_requests +1.
Releasing inside the nested iterator hook does not work because that
generator is only closed on garbage collection, which is non-deterministic.
"""
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
limiter, cache, counter_key, user_api_key_dict = await _build_seeded_limiter()
assert await cache.async_get_cache(key=counter_key) == 1
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = limiter
async def upstream():
yield ModelResponse()
if disconnect == "cancel":
raise asyncio.CancelledError()
while True:
yield ModelResponse()
with _override_litellm_callbacks([]):
gen = ProxyBaseLLMRequestProcessing.async_sse_data_generator(
response=upstream(),
user_api_key_dict=user_api_key_dict,
request_data={"model": "claude-test"},
proxy_logging_obj=proxy_logging_obj,
)
await gen.__anext__()
if disconnect == "cancel":
with pytest.raises(asyncio.CancelledError):
await gen.__anext__()
else:
await gen.aclose()
await _drain_release_task()
assert await cache.async_get_cache(key=counter_key) == 0
@pytest.mark.parametrize("disconnect", ["cancel", "aclose"])
@pytest.mark.asyncio
async def test_async_data_generator_releases_counter_on_disconnect_v3(disconnect):
"""
Regression for issue #27955 on the chat-completions outer generator
(proxy_server.async_data_generator). With only the v3 parallel limiter
enabled, needs_iterator_wrap() is False, so this generator iterates the
upstream response directly and the iterator hook is bypassed entirely -- the
gap that let a disconnect leak the slot in the default limiter-only config.
A mid-stream disconnect must still refund the pre-call +1.
"""
import litellm.proxy.proxy_server as proxy_server
limiter, cache, counter_key, user_api_key_dict = await _build_seeded_limiter()
proxy_logging_obj = proxy_server.proxy_logging_obj
saved_hook = proxy_logging_obj.proxy_hook_mapping.get("parallel_request_limiter")
proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = limiter
async def upstream():
yield ModelResponse()
if disconnect == "cancel":
raise asyncio.CancelledError()
while True:
yield ModelResponse()
try:
with _override_litellm_callbacks([]):
assert proxy_logging_obj.needs_iterator_wrap() is False
gen = proxy_server.async_data_generator(
response=upstream(),
user_api_key_dict=user_api_key_dict,
request_data={"model": "gpt-test"},
)
await gen.__anext__()
if disconnect == "cancel":
with pytest.raises(asyncio.CancelledError):
await gen.__anext__()
else:
await gen.aclose()
await _drain_release_task()
assert await cache.async_get_cache(key=counter_key) == 0
finally:
if saved_hook is not None:
proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = (
saved_hook
)
else:
proxy_logging_obj.proxy_hook_mapping.pop("parallel_request_limiter", None)
@pytest.mark.asyncio
async def test_async_data_generator_releases_counter_when_wrapped_v3():
"""
Companion to the no-wrap case for issue #27955. With an iterator-override
callback active, needs_iterator_wrap() is True and async_data_generator
drives the chained iterator hook. The refund must still fire exactly once
from the outer generator: the counter returns to 0 (not -1), proving the
nested hook does not also refund and there is no double decrement.
"""
from litellm.integrations.custom_logger import CustomLogger
import litellm.proxy.proxy_server as proxy_server
class _PassthroughIteratorOverride(CustomLogger):
async def async_post_call_streaming_iterator_hook(
self, user_api_key_dict, response, request_data
):
async for chunk in response:
yield chunk
limiter, cache, counter_key, user_api_key_dict = await _build_seeded_limiter()
proxy_logging_obj = proxy_server.proxy_logging_obj
saved_hook = proxy_logging_obj.proxy_hook_mapping.get("parallel_request_limiter")
proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = limiter
async def upstream():
while True:
yield ModelResponse()
try:
with _override_litellm_callbacks([_PassthroughIteratorOverride()]):
assert proxy_logging_obj.needs_iterator_wrap() is True
gen = proxy_server.async_data_generator(
response=upstream(),
user_api_key_dict=user_api_key_dict,
request_data={"model": "gpt-test"},
)
await gen.__anext__()
await gen.aclose()
await _drain_release_task()
assert await cache.async_get_cache(key=counter_key) == 0
finally:
if saved_hook is not None:
proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = (
saved_hook
)
else:
proxy_logging_obj.proxy_hook_mapping.pop("parallel_request_limiter", None)
def test_tpm_reservation_enabled_by_default(monkeypatch):
"""Upfront TPM reservation is on unless explicitly disabled via env."""
monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False)

View file

@ -57,6 +57,51 @@ async def test_get_daily_activity_empty_entity_id_list():
assert where_conditions["team_id"] == {"in": []}
@pytest.mark.asyncio
async def test_get_daily_activity_order_has_id_tiebreaker():
"""Regression for #30164.
``date`` alone is not a unique sort key for either
``LiteLLM_DailyUserSpend`` or ``LiteLLM_DailyTeamSpend`` -- a busy
tenant has many rows per date (one per api_key, model, model_group,
provider, endpoint, ...). Offset pagination over a non-unique sort
landed on arbitrary page boundaries between queries, so summing
per-page totals across pages produced non-deterministic results
(sometimes inflated, sometimes deflated). The tiebreaker on the
UUID primary key pins the row order so a client paging through all
results gets the correct total.
"""
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_table = MagicMock()
mock_table.count = AsyncMock(return_value=0)
mock_table.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_verificationtoken = MagicMock()
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_dailyspend = mock_table
await get_daily_activity(
prisma_client=mock_prisma,
table_name="litellm_dailyspend",
entity_id_field="team_id",
entity_id="team-1",
entity_metadata_field=None,
start_date="2024-01-01",
end_date="2024-01-02",
model=None,
api_key=None,
page=1,
page_size=10,
)
mock_table.find_many.assert_called_once()
order = mock_table.find_many.call_args[1]["order"]
assert order == [{"date": "desc"}, {"id": "asc"}], (
f"order must include the id tiebreaker after date for stable offset "
f"pagination (see #30164); got {order!r}"
)
def test_is_user_agent_tag():
"""Test _is_user_agent_tag function."""
# Test None and empty string

View file

@ -1488,6 +1488,65 @@ async def test_prepare_key_update_data_duration_none_never_expires():
assert result["expires"] is None
@pytest.mark.asyncio
@pytest.mark.parametrize("cleared_value", [[], None])
async def test_prepare_key_update_data_budget_limits_clears_field(cleared_value):
"""budget_limits=[] / None must serialize to JSON null, never reach Prisma raw."""
from litellm.proxy._types import UpdateKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import (
prepare_key_update_data,
)
existing_key = LiteLLM_VerificationToken(
token="test-token",
key_alias="test-key",
models=["gpt-3.5-turbo"],
user_id="test-user",
team_id=None,
metadata={},
)
update_request = UpdateKeyRequest(key="test-token", budget_limits=cleared_value)
result = await prepare_key_update_data(
data=update_request, existing_key_row=existing_key
)
assert result["budget_limits"] == json.dumps(None)
@pytest.mark.asyncio
async def test_prepare_key_update_data_budget_limits_serializes_windows():
"""Non-empty budget_limits stay JSON-encoded with reset_at initialized."""
from litellm.proxy._types import UpdateKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import (
prepare_key_update_data,
)
existing_key = LiteLLM_VerificationToken(
token="test-token",
key_alias="test-key",
models=["gpt-3.5-turbo"],
user_id="test-user",
team_id=None,
metadata={},
)
update_request = UpdateKeyRequest(
key="test-token",
budget_limits=[{"budget_duration": "1d", "max_budget": 10.0}],
)
result = await prepare_key_update_data(
data=update_request, existing_key_row=existing_key
)
windows = json.loads(result["budget_limits"])
assert isinstance(result["budget_limits"], str)
assert windows[0]["max_budget"] == 10.0
assert windows[0]["reset_at"] is not None
@pytest.mark.asyncio
async def test_validate_team_id_used_in_service_account_request_requires_team_id():
"""
@ -9685,6 +9744,58 @@ class TestKeyOwnerPrivilegeEscalation:
)
mock_check.assert_called_once()
@pytest.mark.asyncio
@pytest.mark.parametrize("cleared_value", [[], None])
async def test_creator_cannot_clear_own_budget_limits(self, cleared_value):
"""Clearing budget_limits is a budget change and requires admin."""
data = UpdateKeyRequest(key="sk-test", budget_limits=cleared_value)
existing = self._make_existing_key(created_by="creator-123")
auth = self._make_auth(user_id="creator-123")
mock_check = AsyncMock(
side_effect=HTTPException(status_code=403, detail="Not authorized")
)
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints._check_key_admin_access",
mock_check,
):
with pytest.raises(HTTPException):
await _validate_update_key_data(
data=data,
existing_key_row=existing,
user_api_key_dict=auth,
llm_router=None,
premium_user=False,
prisma_client=AsyncMock(),
user_api_key_cache=MagicMock(),
)
mock_check.assert_called_once()
@pytest.mark.asyncio
async def test_admin_can_clear_budget_limits(self):
data = UpdateKeyRequest(key="sk-test", budget_limits=[])
existing = self._make_existing_key(created_by="someone-else")
auth = UserAPIKeyAuth(
user_id="admin-user",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
mock_check = AsyncMock()
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints._check_key_admin_access",
mock_check,
):
await _validate_update_key_data(
data=data,
existing_key_row=existing,
user_api_key_dict=auth,
llm_router=None,
premium_user=False,
prisma_client=AsyncMock(),
user_api_key_cache=MagicMock(),
)
mock_check.assert_not_called()
@pytest.mark.asyncio
async def test_admin_can_update_any_field(self):
data = UpdateKeyRequest(key="sk-test", models=["gpt-4"], max_budget=999.0)

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