mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
9ddf0535b1
commit
cfcdf8714a
112 changed files with 12140 additions and 446 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
345
litellm/integrations/focus/destinations/mavvrik_destination.py
Normal file
345
litellm/integrations/focus/destinations/mavvrik_destination.py
Normal 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
|
||||
)
|
||||
0
litellm/integrations/mavvrik_focus/__init__.py
Normal file
0
litellm/integrations/mavvrik_focus/__init__.py
Normal file
272
litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py
Normal file
272
litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py
Normal 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
|
||||
)
|
||||
10
litellm/integrations/newrelic/__init__.py
Normal file
10
litellm/integrations/newrelic/__init__.py
Normal 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"]
|
||||
926
litellm/integrations/newrelic/newrelic.py
Normal file
926
litellm/integrations/newrelic/newrelic.py
Normal 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")
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)),
|
||||
|
|
|
|||
BIN
litellm/proxy/_experimental/out/assets/logos/newrelic.png
Normal file
BIN
litellm/proxy/_experimental/out/assets/logos/newrelic.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 862 B |
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
46
litellm/proxy/guardrails/guardrail_hooks/ovalix/__init__.py
Normal file
46
litellm/proxy/guardrails/guardrail_hooks/ovalix/__init__.py
Normal 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,
|
||||
}
|
||||
330
litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py
Normal file
330
litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py
Normal 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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 #################################
|
||||
####################################################################################
|
||||
####################################################################################
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 ###
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
9
litellm/types/integrations/newrelic.py
Normal file
9
litellm/types/integrations/newrelic.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
37
litellm/types/proxy/guardrails/guardrail_hooks/ovalix.py
Normal file
37
litellm/types/proxy/guardrails/guardrail_hooks/ovalix.py
Normal 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"
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
1351
tests/test_litellm/integrations/newrelic/test_newrelic.py
Normal file
1351
tests/test_litellm/integrations/newrelic/test_newrelic.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
0
tests/test_litellm/llms/parallel_ai/__init__.py
Normal file
0
tests/test_litellm/llms/parallel_ai/__init__.py
Normal file
324
tests/test_litellm/llms/parallel_ai/test_parallel_ai_search.py
Normal file
324
tests/test_litellm/llms/parallel_ai/test_parallel_ai_search.py
Normal 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",
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue