mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
* fix(proxy): bump health-check max_tokens default to 16 for GPT-5 compatibility (#30708) OpenAI GPT-5 models require max_completion_tokens >= 16. Health checks were using 5 (proxy/health_check.py) and 10 (health_check_helpers.py), causing failures on GPT-5 models. Fixes #23836 * fix: increase health check max_tokens from 5 to 16 (#23836) (#26610) GPT-5 models enforce a minimum of 16 for max_output_tokens. The current default of 5 still causes health checks to fail for these models. Bump the non-wildcard default to 16 — the smallest value that satisfies all known provider minimums while keeping health checks lightweight. Also tightens the wildcard test assertion from a weak disjunctive check to strict key-absence. Co-authored-by: Sameer Kankute <sameer@berri.ai> * fix: ensure checks show gemini-3-flash-preview supports responseJsonS… (#30696) * fix: ensure checks show gemini-3-flash-preview supports responseJsonSchema. * fix: remove async keyword from test. * fix: make Bedrock Mantle Responses routing data-driven per model (#30700) * Make Bedrock Mantle Responses routing data-driven per model Route Bedrock Mantle models to the native Responses API based on each model's price-map capability signal instead of a hardcoded model-name heuristic, and derive the OpenAI-compatible base path segment per model. Responses dispatch now selects the native config when the model advertises responses support (/v1/responses in supported_endpoints, or mode=responses), both overridable via register_model and proxy model_info. This enables native Responses for gpt-oss-120b/20b and the gemma-4 family while keeping chat-only models (gpt-oss safeguard, nvidia, mistral, ...) on the existing chat-completions emulation. Capability is per-model, so gpt-oss-120b routes natively while gpt-oss-safeguard-120b does not despite sharing the gpt-oss substring. The wire path is a separate concern, driven by the existing use_openai_responses_path flag rather than a model-name match: gpt-5.x and gemma-4-* on /openai/v1, everything else (incl. gpt-oss) on /v1. The chat config now derives its base from the same flag, fixing gemma-4 chat-completions requests that previously went to /v1 instead of /openai/v1. Cost maps: add supported_endpoints to the gpt-oss entries (responses for the non-safeguard variants, chat-only for safeguard) and supported_endpoints + use_openai_responses_path to all three gemma-4 entries. Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com> * Address review: move capability helper into bedrock_mantle package Move the Responses capability check out of utils.py into litellm/llms/bedrock_mantle/common_utils.py as mantle_supports_responses, alongside its companion wire-path helper mantle_base_segment. Both are now pure functions of (model, model_cost): the price-map mode/supported_endpoints read replaces the get_model_info call, so the rules are unit-testable without patching global state and the Bedrock Mantle package is self-contained. Use str | None instead of Optional[str] on the new signatures to satisfy the ruff UP045 strict-rule gate. Add direct unit tests for both helpers. Fix test_register_model_restore_undoes_existing_key_overwrite: gpt-oss-120b now legitimately supports Responses, so it can no longer be the "None after restore" vehicle; use the chat-only safeguard variant, which isolates the register/restore effect from the model's own capability. Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com> --------- Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com> Co-authored-by: Sameer Kankute <sameer@berri.ai> * fix(proxy): fail fast on non-PostgreSQL DATABASE_URL instead of hanging on startup (#30366) * fix(proxy): fail fast on non-PostgreSQL DATABASE_URL instead of hanging on startup LiteLLM's Prisma datasource is pinned to provider = 'postgresql', so a sqlite:// or mysql:// DATABASE_URL can never connect. Today that surfaces as an opaque startup stall where the port never binds, and a separate 'DB not connected' 500 on /key/generate when no DATABASE_URL is set at all leaves operators guessing what to configure. Validate the DATABASE_URL / DIRECT_URL scheme in run_server before any Prisma call and exit with an actionable message naming the unsupported scheme. Also reword CommonProxyErrors.db_not_connected_error to tell the operator to set DATABASE_URL to a postgresql:// connection string. Add regression tests covering postgres acceptance and sqlite/mysql/mssql rejection. * fix: resolve CI failures and proxy DB URL typing issue * fix(dashscope): treat an explicit 0.0 tier cost as a real price, not missing (#30653) The tiered cost calculator resolved a tier's per-token cost with `tier.get(cost_key) or tier.get(fallback_cost_key, 0)`. Because `or` short-circuits on any falsy value, a tier that legitimately prices a component at 0.0 (e.g. a free-cache-read tier with cache_read_input_token_cost: 0.0, or a free-reasoning tier) is treated as missing and silently billed at the full fallback rate (input_cost_per_token / output_cost_per_token). The flat-pricing path in the same module already handles this correctly with an `is None` guard. Resolve tier costs through a small helper that mirrors it, so 0.0 is honored at both the in-range and overflow sites. No shipped model currently has a 0.0 tier cost, so this is a latent defect; the fix makes the tiered path consistent with the flat path and prevents over-charging the first time such a tier appears. Adds unit tests covering the in-range and overflow paths, and drops an unused import flagged by ruff in the touched test file. * feat(proxy): show session-aggregate cost and duration in request logs (#25708) (#30507) * fix(anthropic): don't leak tool 'type' into OpenAI function parameters schema (#30618) In the messages->chat/completions bridge, translate_anthropic_tools_to_openai merged every non-mapped tool key into the function parameters dict. The Anthropic tool 'type' (e.g. 'custom') thus overwrote parameters.type ('object' -> 'custom'), and providers reject it ('custom' is not a valid JSON-Schema type). Exclude 'type' from the passthrough. Fixes #30557. * fix(proxy): stop IAM-refresh engine restart from cascading reconnects (#29176) (#30183) An RDS IAM token refresh recreates the Prisma client, which SIGKILLs the running query-engine and spawns a new one. That planned kill was indistinguishable from a crash, and three reconnect paths used two uncoordinated locks, so a single refresh triggered a cascade of engine kill/respawn cycles: 1. `_safe_refresh_token` (holds `_reconnection_lock`) -> recreate -> kill old engine, spawn new one. 2. The engine-death watcher sees that kill, assumes a crash, and calls `attempt_db_reconnect(force=True)` (a different lock, `_db_reconnect_lock`) -> recreate again -> kills the fresh engine. 3. In-flight queries failing during the swap are classified as transport errors and trigger their own `attempt_db_reconnect` -> recreate again. Fix coordinates planned restarts across the wrapper and the watcher: - PrismaWrapper records the old engine PID in `_expected_engine_deaths` before killing it; all four watcher death-detectors (waitpid thread, pidfd, already-dead probe, os.kill poll) consume that PID and skip the reconnect instead of treating it as a crash. - `recreate_prisma_client` now serializes through `_reconnection_lock` and bumps a monotonic `_engine_generation`. Callers pass `expected_generation` as an optimistic-lock token, so racing/cascading recreates collapse into a single restart (losers no-op). This closes the two-lock gap. - The direct reconnect path probes the writer with SELECT 1 before recreating; a healthy connection (e.g. engine already replaced by a refresh) skips the recreate entirely. - `_safe_refresh_token` coalesces: it skips when the current token still has more than the refresh buffer of runway, so stacked triggers (proactive loop + __getattr__ fallback) don't each restart the engine. An `on_engine_replaced` hook re-arms the watcher on the new PID. RoutingPrismaWrapper forwards `expected_generation` and skips recreating the reader when the writer recreate was skipped. * feat(bedrock): support file content retrieval for batch output files (#30595) Implements transform_file_content_request and transform_file_content_response in BedrockFilesConfig so GET /v1/files/{id}/content works for Bedrock batch files. The request transform resolves the file id (direct s3:// URI or base64 unified id) to its S3 object, validates bucket and key prefix against the server-configured bucket, and SigV4-signs an S3 GetObject using the same credential and region resolution as the existing upload path. The credential and region params are validated into a typed model at the boundary, so the only untyped values left are the botocore signing primitives. Also fixes the proxy managed-files path: CredentialLiteLLMParams now carries s3_bucket_name (previously dropped when building deployment credentials) and the managed-files hook passes the deployment credential snapshot when routing afile_content, so unified-id content retrieval works with per-model bucket config instead of only the AWS_S3_BUCKET_NAME env var. Preserves managed-file access control: the proxy file-content endpoint now rejects raw cloud-storage ids (s3://, gs://), which would otherwise skip the owner/team check that only runs for unified ids and let a caller read another tenant's batch output by its object key. Managed outputs are reachable only through their unified file id. The afile_content "not found" error now reports the caller's unified id rather than the resolved internal S3 URI. Fixes #16186, #15563 * fix(oci): make Cohere {{trace}} judges work (tool param types + agentic tool-calling continuation) (#30646) * fix(oci): map Cohere tool array/object params to lowercase builtins OCI's Cohere backend returns HTTP 500 on a tool parameter typed as a bare "List", which is what OCI_JSON_TO_PYTHON_TYPES produced for JSON-schema arrays. MLflow {{trace}} judges trip this: their tools (get_root_span, get_span) take an attributes_to_fetch array. The lowercase builtins list/dict are accepted; only the bare "List" 500s ("Dict" happens to be tolerated, but both are lowercased for consistency). Verified live against us-chicago-1 (cohere.command-a-03-2025 and command-latest). Adds a unit regression on the transformed parameterDefinitions plus a gated integration test exercising an array-param tool end to end. * fix(oci): make Cohere agentic tool-calling continuation work Two bugs broke the OCI Cohere tool-calling loop that MLflow {{trace}} judges drive once a tool has been executed and its result is fed back. Request side: litellm pulled the last user message into the top-level `message` and emitted the tool result as a TOOL entry in chatHistory. OCI rejects that ("cannot specify message if the last entry in chat history contains tool results"), and an empty message alone is rejected too ("message must be at least 1 token long or tool results must be specified"). OCI carries the current turn's results in a dedicated top-level `toolResults` field. The Cohere transform now sends an empty message, keeps the user turn in chatHistory, and puts the results in `toolResults`, matching the langchain-oracle reference. Tool results are no longer represented as chatHistory entries. Response side: tool-grounded answers come back with citations carrying `documentIds` (camelCase) and no `document_ids`, which made the required `CohereCitation.document_ids` field fail validation and sink the whole response parse. Those citations are never surfaced, so the field (and CohereSearchQuery's generation_id) is now optional. Verified live against us-chicago-1 (cohere.command-a-03-2025 and command-latest), single and multi-round tool loops. Adds unit regressions on the transformed request shape and on citation parsing, plus gated integration tests for the continuation. * feat: integrate Repelloai Argus guardrail (#30673) * feat(guardrails): add RepelloAI Argus guardrail integration (#1) * feat(guardrails): add RepelloAI Argus guardrail integration Add a new guardrail hook backed by RepelloAI Argus, with dashboard-managed asset policies enforced via an asset_id and X-API-Key auth. * fix(guardrails): harden RepelloAI Argus guardrail - scan streaming responses on output (was bypassing the guardrail) - log blocked verdicts as guardrail_intervened instead of success - treat auth/config errors (401/403/404/422) as misconfiguration that always blocks, not a fail-open-able unreachable error - default unreachable_fallback to fail_closed and read it directly; block on unknown/malformed verdicts so an API change can't silently disable enforcement - type unreachable_fallback as a Literal, drop the duplicate config model, expose unreachable_fallback in the config schema, and stop leaking the raw provider response / exception strings to the client * fix(guardrails): address RepelloAI Argus review feedback - support ARGUS_API_KEY (with REPELLOAI_API_KEY fallback) - make asset_id required in the config model - normalize unreachable_fallback so only fail_open opens; block on 400 misconfig - correct the shared unreachable_fallback field description * docs(guardrails): add RepelloAI Argus docs page and dashboard listing - add docs page covering config, env vars, modes, verdicts, failure semantics - list RepelloAI Argus in the Guardrail Garden with provider/logo mappings - add a regression test for the provider logo and display-name resolution * fix(guardrails): keep RepelloAI asset_id optional in config model A required asset_id leaked onto the shared LitellmParams (which inherits RepelloAIGuardrailConfigModel), breaking validation for every other guardrail. Keep it optional like sibling models; the guardrail __init__ still raises when asset_id is missing, which is the real enforcement. * Add comment for last user turn scanning * feat(guardrails): harden repelloai scanning * feat(guardrails): expand repelloai scanning to include tool definitions Add extraction of tool definitions and tool call arguments to the RepelloAI guardrail scanning. Improves detection coverage by including function schemas and parameters in the prompt sent to the guardrail service. Also captures detailed error responses in logs and adds guardrail header to streaming responses. * refactor(guardrails): fix and harden repelloai schema text extraction - Fix duplicate text in _iter_schema_text: previously all dict values were re-queued onto the stack even after scalar/list keys were already extracted explicitly, causing names/descriptions to appear twice in the scanned prompt - Extract schema key frozensets to module-level constants so they are not reconstructed on every call - Change _iter_schema_text from @classmethod to @staticmethod (cls unused) - Narrow _call_analyze stage param from str to Literal["prompt", "response"] - Add HttpxResponse type annotation to _raise_for_config_error - Add LLMResponseTypes annotation to async_post_call_success_hook response param * fix(guardrails): resolve pyright type errors in repelloai guardrail - Narrow async_handler.post return from Response|None to Response with explicit None guard before calling raise_for_status/json - Fix list comprehension returning str|None by switching to explicit loop with isinstance guard so pyright tracks the narrowing - Cast model_dump() result to Dict since hasattr does not narrow object type in pyright * fix(guardrails/repello): include Responses API instructions field in prompt scan The /v1/responses top-level `instructions` field was not included in _extract_prompt_text, allowing a caller to bypass guardrail policy checks by putting blocked content in `instructions` while keeping `input` benign. * feat: add api_key to config model and read prompt from data dict * fix(guardrails/repello): plug input_text and tool-call response bypass gaps Responses API input content parts with type 'input_text' were silently dropped by build_inspection_messages (which only handles type='text'), allowing callers to send blocked content via that path without triggering the pre-call scan. Fix: add _extract_input_text_parts to RepelloAIGuardrail and call it when walking the Responses API input messages. Post-call scanning skipped responses whose choices contained only tool_calls or function_call (message.content=None), letting models put blocked output in function arguments undetected. Fix: _extract_chat_completion_text now calls _extract_tool_call_args_from_message on each choice message. Also replace typing.Dict/List with builtin dict/list to clear TID251 strict ruff violations introduced by this file. * fix(guardrails/repello): scan Responses API function_call output arguments Output items with type 'function_call' in a /v1/responses response were skipped by _extract_responses_api_text; only 'message' items were walked. A model could return blocked content in function_call.arguments undetected. Now extract arguments from function_call output items before scanning. * refactor(guardrails/repello): clean up typing and remove lint-any workarounds - Replace Optional[X]/Union[X,Y] with X|None/X|Y union syntax throughout - Use dict[str, object] instead of bare dict in all signatures - Remove **kwargs from __init__; declare guardrail_name, event_hook, default_on explicitly - Replace getattr(litellm_params, ...) with direct attribute access now that LitellmParams inherits RepelloAIGuardrailConfigModel - Add _event_hook_from_mode() to convert str|list[str]|Mode to typed GuardrailEventHooks - Use TypeAdapter.validate_json() instead of response.json() + manual dict construction - Add _is_object_dict/_is_object_list TypeGuard helpers to narrow object types without Any - Remove cast() workarounds and typed intermediate variables that existed only for the now-removed lint-any CI check - Drop _AddLiteLLMCallback Protocol; budget has sufficient slack for the one reportUnknownMemberType - Fix GuardrailConfigModel missing type arg: GuardrailConfigModel[BaseModel] * fix(guardrails/repello): suppress LIT007 on TypeGuard helpers and add streaming scan-skip warning - Add guard-ok suppressions to _is_object_dict and _is_object_list to satisfy the LIT007 hard-zero budget gate - Emit verbose_proxy_logger.warning when the streaming hook finds no inspectable text after assembly, matching observability of pre/post hooks * refactor: modifications for lint check * feat: add Pinstripes as an OpenAI-compatible provider (#30567) * feat: add Pinstripes as an OpenAI-compatible provider Pinstripes (https://pinstripes.io) is an OpenAI-compatible inference provider serving open-source models (GLM-4.5-Air, Qwen3, DeepSeek, etc.) with per-token pricing and no subscriptions. Changes: - `litellm/llms/openai_like/providers.json`: register pinstripes with base_url, api_key_env, and max_completion_tokens→max_tokens mapping - `litellm/types/utils.py`: add `PINSTRIPES = "pinstripes"` to LlmProviders - `litellm/constants.py`: add to openai_compatible_providers and openai_compatible_endpoints lists - `litellm/litellm_core_utils/get_llm_provider_logic.py`: auto-detect provider when api_base is "https://pinstripes.io/v1" - `provider_endpoints_support.json`: document supported endpoints - `tests/`: 7 unit tests covering provider registration, resolution, URL auto-detection, api_base override, and Router config Usage: import litellm response = litellm.completion( model="pinstripes/ps/glm-4.5-air", messages=[{"role": "user", "content": "Hello"}], api_key=os.environ["PINSTRIPES_API_KEY"], ) Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> * fix(pinstripes): resolve Greptile P1 review comments - Add api_base_env: PINSTRIPES_API_BASE to providers.json so env var override works - Set responses: false in provider_endpoints_support.json — not actually wired up - Remove docs/my-website/docs/providers/pinstripes.md — belongs in litellm-docs repo Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> * fix(pinstripes): add api_base_env and correct responses capability - Add api_base_env: PINSTRIPES_API_BASE to providers.json - Set responses: false in provider_endpoints_support.json Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> * fix(pinstripes): wire up Responses API — add supported_endpoints Adds supported_endpoints: ["/v1/chat/completions", "/v1/responses"] so JSONProviderRegistry.supports_responses_api returns true correctly, matching what provider_endpoints_support.json advertises. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> * feat(pinstripes): enable embeddings endpoint Pinstripes serves nomic-embed-text-v1.5 and bge-m3 via /v1/embeddings. Add /v1/embeddings to supported_endpoints and set embeddings: true. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> * fix(pinstripes): use 4-space indentation in model_prices_and_context_window.json Matches the file's existing convention. Flagged by Greptile review. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> * fix(pinstripes): set a2a: false — A2A protocol not implemented All comparable JSON-configured providers (tensormesh, parasail, empiriolabs, libertai, neosantara) have a2a: false. Pinstripes does not implement the Google A2A protocol, so this should be false to match. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> --------- Co-authored-by: inference_provider <max@redactedlab.com> Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com> * fix(rag): attach existing OpenAI file ids (#30628) * fix(rag): attach existing OpenAI file ids * chore: use modern typing in rag ingest fix * chore: retrigger ci * fix(anthropic-messages): apply cache_control_injection_points on /v1/messages path (#30341) cache_control_injection_points was only consumed by the chat/completions prompt-management hook; on the native Anthropic /v1/messages path it was forwarded unused, so deployment-level cache injection was silently dropped (cache_creation_input_tokens stayed 0 for Anthropic-native clients). Add AnthropicCacheControlHook.apply_to_anthropic_messages_request to inject cache_control at block level for system / tools / message locations (the only forms /v1/messages accepts), wire it into the native anthropic_messages handler, and pop the param so it does not leak upstream as an unknown field. A {location: message, role: system} config is redirected to the top-level system prompt so the same YAML works on both endpoints. Injection respects Anthropic's 4-block cache_control limit shared across system, tools, and messages: client-supplied markers count toward the cap and are never overwritten, a slot is reserved per Bedrock tool_config point, and injection stops once the budget is exhausted. Locations this path cannot represent (tool_config) are forwarded downstream instead of being silently consumed, mirroring get_chat_completion_prompt's remaining_points pass-through. Built on litellm_internal_staging. Refs BerriAI/litellm#30293 * fix(proxy): release budget reservation when a request is cancelled mid-flight (#30522) * fix(proxy): release budget reservation on cancel when no chunk was delivered The pre-call budget reservation increments the cross-pod spend counter by a request's worst-case cost, then reconciles it on success (cost callback) or error (failure hook). A client disconnect or timeout cancels the request and surfaces as CancelledError / GeneratorExit, which neither path catches, so the reservation leaks. Under a retry storm the leaked holds accumulate, pin the counter above real spend, and return spurious 429 "Budget has been exceeded" to keys whose spend is far below budget; the counter only recovers when its TTL lapses, so the failure is intermittent and self-healing. Release the reservation in async_streaming_data_generator (which the Anthropic and Google SSE generators delegate to) on the (CancelledError, GeneratorExit) path, alongside the existing max_parallel_requests release. release_budget_ reservation_on_cancel runs under asyncio.shield so it completes despite the in-progress cancellation, is guarded by the reservation's finalized flag, and swallows a failing release so it cannot replace the in-flight cancellation. The refund is gated on whether a chunk reached the client. The flag is set immediately before the yield, after the slow-path hook await: an async generator suspends at the yield, so a GeneratorExit on disconnect after a delivered chunk sees it True (keep the hold), while a cancellation during the slow-path await leaves it False (refund, nothing sent). A non-streaming cancellation delivers nothing and a completed non-streaming response is reconciled by the success callback, so neither needs a release here. Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> * fix(proxy): reconcile a cancelled reservation to input cost, not zero A streaming request cancelled before the first chunk previously reconciled its reservation to zero and finalized it. But by the time the generator is consuming the response the provider call was already dispatched, so the input tokens were billed even though no chunk reached the client, and the success/failure cost callbacks are skipped on cancellation. Refunding to zero let a caller send an expensive request and abort pre-token to dodge the input charge. Compute the request's input-token cost at reservation time and reconcile the cancelled reservation to it instead of zero. The worst-case output portion of the reservation is still released (so a legitimate mid-flight cancellation no longer pins the counter and 429s the key), while the input the provider already processed is charged. --------- Co-authored-by: Bytechoreographer <Bytechoreographer@users.noreply.github.com> Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> * fix(caching): encode object name in GCS cache GET path (#30378) GCS cache reads always missed when gcs_path was set. The GET methods interpolated the object name directly into the URL path, while the GCS JSON API requires it to be URL-encoded (a "/" must be sent as %2F). With gcs_path configured the object name is "<prefix>/<sha256>", so the raw slash produced a malformed object path and GCS returned 404. httpx does not raise on 4xx, so the status_code == 200 check fell through and get/async_get returned None, silently missing on every read. Without gcs_path the key has no slash, which is why this went unnoticed. Wrap the object name with urllib.parse.quote(..., safe="") in get_cache and async_get_cache. Apply the same encoding to the name= query parameter in set_cache and async_set_cache so the key written matches the key read back. Adds regression tests asserting the GET path and SET query are encoded (%2F) when gcs_path is set, for both sync and async paths; these fail on the unpatched code. Fixes #30377 * chore: add soniox stt-async-v5 model (#30672) * fix(proxy): include model group aliases in v1 model info (#30626) * Include model group aliases in v1 model info * Fix model info alias implementation * removed extra blank line * chore: rerun CI * fix(lint): remove redundant noqa directive in proxy_cli.py * fix: address greptile review - restore bedrock_mantle auth symbols, guard OCI empty message list, validate DIRECT_URL scheme * Revert "fix: address greptile review - restore bedrock_mantle auth symbols, guard OCI empty message list, validate DIRECT_URL scheme" This reverts commit52c7a07777. * Revert "fix(anthropic-messages): apply cache_control_injection_points on /v1/messages path (#30341)" This reverts commitc9e8a177bd. * Revert "fix(proxy): stop IAM-refresh engine restart from cascading reconnects (#29176) (#30183)" This reverts commit85828da695. * fix(proxy): stop IAM-refresh engine restart from cascading reconnects (#29176) (#30183) An RDS IAM token refresh recreates the Prisma client, which SIGKILLs the running query-engine and spawns a new one. That planned kill was indistinguishable from a crash, and three reconnect paths used two uncoordinated locks, so a single refresh triggered a cascade of engine kill/respawn cycles: 1. `_safe_refresh_token` (holds `_reconnection_lock`) -> recreate -> kill old engine, spawn new one. 2. The engine-death watcher sees that kill, assumes a crash, and calls `attempt_db_reconnect(force=True)` (a different lock, `_db_reconnect_lock`) -> recreate again -> kills the fresh engine. 3. In-flight queries failing during the swap are classified as transport errors and trigger their own `attempt_db_reconnect` -> recreate again. Fix coordinates planned restarts across the wrapper and the watcher: - PrismaWrapper records the old engine PID in `_expected_engine_deaths` before killing it; all four watcher death-detectors (waitpid thread, pidfd, already-dead probe, os.kill poll) consume that PID and skip the reconnect instead of treating it as a crash. - `recreate_prisma_client` now serializes through `_reconnection_lock` and bumps a monotonic `_engine_generation`. Callers pass `expected_generation` as an optimistic-lock token, so racing/cascading recreates collapse into a single restart (losers no-op). This closes the two-lock gap. - The direct reconnect path probes the writer with SELECT 1 before recreating; a healthy connection (e.g. engine already replaced by a refresh) skips the recreate entirely. - `_safe_refresh_token` coalesces: it skips when the current token still has more than the refresh buffer of runway, so stacked triggers (proactive loop + __getattr__ fallback) don't each restart the engine. An `on_engine_replaced` hook re-arms the watcher on the new PID. RoutingPrismaWrapper forwards `expected_generation` and skips recreating the reader when the writer recreate was skipped. * fix(lint): modernize type annotations in IAM-refresh prisma client files (UP006/UP045) * Revert "feat(proxy): show session-aggregate cost and duration in request logs (#25708) (#30507)" This reverts commitf530b2237c. * Revert "fix(dashscope): treat an explicit 0.0 tier cost as a real price, not missing (#30653)" This reverts commit4f58bd0df5. * Revert "fix(oci): make Cohere {{trace}} judges work (tool param types + agentic tool-calling continuation) (#30646)" This reverts commit50f34e0b15. * Revert "fix(proxy): fail fast on non-PostgreSQL DATABASE_URL instead of hanging on startup (#30366)" This reverts commit0544eed6ea. * fix(bedrock_mantle): restore BedrockMantleAuthMixin and constants removed by routing rewrite * fix(key management): restore exact /key/list user_id & key_alias matching by default (#30593) Before substring search was added (commit33bd570d5e), /key/list matched user_id and key_alias exactly. That change made admin-authenticated calls substring-match by default, breaking the prior contract: a caller passing an exact user_id as an access filter (e.g. an integration scoping to one user with an admin key) then received other users' keys -- user_id="alice" also returned "alice2", "alice-test", etc. This is a cross-user key disclosure. Make substring matching opt-in via a new admin-only substring_matching=true query param; default to exact, restoring the prior behavior. The dashboard search box (keyListCall) passes the flag so partial search still works. Non-admins remain exact and scoped to their own keys. Updates the proxy-behavior key_alias test to opt in and adds an exact-by-default guard; adds list_keys unit coverage for the opt-in gate. --------- Co-authored-by: perseus <51974392+tcconnally@users.noreply.github.com> Co-authored-by: Hannah Smith <64043506+hannahmadison@users.noreply.github.com> Co-authored-by: Charlie Patterson <Pattersoncharlesl@gmail.com> Co-authored-by: Matthew Lapointe <mlapointe@alpha-sense.com> Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com> Co-authored-by: KRISH SONI <67964054+krishvsoni@users.noreply.github.com> Co-authored-by: Yash Raj Pandey <55940078+devYRPauli@users.noreply.github.com> Co-authored-by: Nitish Agarwal <1592163+nitishagar@users.noreply.github.com> Co-authored-by: hcl <chenglunhu@gmail.com> Co-authored-by: tushar8408 <32977767+tushar8408@users.noreply.github.com> Co-authored-by: AD Mohanraj <admohanraj@gmail.com> Co-authored-by: Fede Kamelhar <federico.kamelhar@oracle.com> Co-authored-by: Lavish Bansal <lavish.bansal619@gmail.com> Co-authored-by: max-amos <gruffulom@gmail.com> Co-authored-by: inference_provider <max@redactedlab.com> Co-authored-by: NK <93352237+Nithish-Yenaganti@users.noreply.github.com> Co-authored-by: 安妮的心动录 <74543653+anneheartrecord@users.noreply.github.com> Co-authored-by: Rick <26716961+Bytechoreographer@users.noreply.github.com> Co-authored-by: Bytechoreographer <Bytechoreographer@users.noreply.github.com> Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> Co-authored-by: Burak Ömür <burak.omur.1998@gmail.com> Co-authored-by: Dan Lemon <daniel.lemon@amazee.io> Co-authored-by: Vanika Dangi <166420943+vanika02@users.noreply.github.com> Co-authored-by: Jay Gowdy <130084966+jgowdy-godaddy@users.noreply.github.com>
1702 lines
68 KiB
Python
1702 lines
68 KiB
Python
# What is this?
|
|
## This hook is used to check for LiteLLM managed files in the request body, and replace them with model-specific file id
|
|
|
|
import base64
|
|
import json
|
|
from types import MappingProxyType
|
|
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast
|
|
|
|
from fastapi import HTTPException
|
|
|
|
import litellm
|
|
from litellm import Router, verbose_logger
|
|
from litellm._uuid import uuid
|
|
from litellm.caching.caching import DualCache
|
|
from litellm.integrations.custom_logger import CustomLogger
|
|
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
|
|
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
|
|
from litellm.llms.base_llm.managed_resources.isolation import (
|
|
build_list_page,
|
|
build_owner_filter,
|
|
can_access_resource,
|
|
)
|
|
from litellm.proxy._types import (
|
|
CallTypes,
|
|
LiteLLM_ManagedFileTable,
|
|
LiteLLM_ManagedObjectTable,
|
|
UserAPIKeyAuth,
|
|
)
|
|
from litellm.proxy.openai_files_endpoints.common_utils import (
|
|
_is_base64_encoded_unified_file_id,
|
|
get_batch_id_from_unified_batch_id,
|
|
get_content_type_from_file_object,
|
|
get_model_id_from_unified_batch_id,
|
|
get_models_from_unified_file_id,
|
|
normalize_mime_type_for_provider,
|
|
)
|
|
from litellm.types.llms.openai import ( # pyright: ignore[reportAttributeAccessIssue]
|
|
AllMessageValues,
|
|
AsyncCursorPage,
|
|
ChatCompletionFileObject,
|
|
CreateFileRequest,
|
|
FileObject,
|
|
OpenAIFileObject,
|
|
OpenAIFilesPurpose,
|
|
ResponsesAPIResponse,
|
|
)
|
|
from litellm.types.utils import (
|
|
CallTypesLiteral,
|
|
LiteLLMBatch,
|
|
LiteLLMFineTuningJob,
|
|
LLMResponseTypes,
|
|
SpecialEnums,
|
|
)
|
|
|
|
if TYPE_CHECKING:
|
|
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
|
|
|
|
|
if TYPE_CHECKING:
|
|
from opentelemetry.trace import Span as _Span
|
|
|
|
from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
|
|
from litellm.proxy.utils import PrismaClient as _PrismaClient
|
|
|
|
Span = Union[_Span, Any]
|
|
InternalUsageCache = _InternalUsageCache
|
|
PrismaClient = _PrismaClient
|
|
else:
|
|
Span = Any
|
|
InternalUsageCache = Any
|
|
PrismaClient = Any
|
|
|
|
|
|
class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|
# Class variables or attributes
|
|
def __init__(
|
|
self, internal_usage_cache: InternalUsageCache, prisma_client: PrismaClient
|
|
):
|
|
self.internal_usage_cache = internal_usage_cache
|
|
self.prisma_client = prisma_client
|
|
|
|
@staticmethod
|
|
def _get_prometheus_logger():
|
|
"""Find PrometheusLogger from litellm.callbacks, if registered."""
|
|
from litellm.integrations.prometheus import PrometheusLogger
|
|
|
|
return PrometheusLogger.get_instance()
|
|
|
|
async def store_unified_file_id(
|
|
self,
|
|
file_id: str,
|
|
file_object: Optional[OpenAIFileObject],
|
|
litellm_parent_otel_span: Optional[Span],
|
|
model_mappings: Dict[str, str],
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
) -> None:
|
|
verbose_logger.info(
|
|
f"Storing LiteLLM Managed File object with id={file_id} in cache"
|
|
)
|
|
if file_object is not None:
|
|
litellm_managed_file_object = LiteLLM_ManagedFileTable(
|
|
unified_file_id=file_id,
|
|
file_object=file_object,
|
|
model_mappings=model_mappings,
|
|
flat_model_file_ids=list(model_mappings.values()),
|
|
created_by=user_api_key_dict.user_id,
|
|
team_id=user_api_key_dict.team_id,
|
|
updated_by=user_api_key_dict.user_id,
|
|
)
|
|
await self.internal_usage_cache.async_set_cache(
|
|
key=file_id,
|
|
value=litellm_managed_file_object.model_dump(),
|
|
litellm_parent_otel_span=litellm_parent_otel_span,
|
|
)
|
|
|
|
## STORE MODEL MAPPINGS IN DB
|
|
|
|
db_data = {
|
|
"unified_file_id": file_id,
|
|
"model_mappings": json.dumps(model_mappings),
|
|
"flat_model_file_ids": list(model_mappings.values()),
|
|
"created_by": user_api_key_dict.user_id,
|
|
"team_id": user_api_key_dict.team_id,
|
|
"updated_by": user_api_key_dict.user_id,
|
|
}
|
|
|
|
if file_object is not None:
|
|
db_data["file_object"] = file_object.model_dump_json()
|
|
# Extract storage metadata from hidden params if present
|
|
hidden_params = getattr(file_object, "_hidden_params", {}) or {}
|
|
if "storage_backend" in hidden_params:
|
|
db_data["storage_backend"] = hidden_params["storage_backend"]
|
|
if "storage_url" in hidden_params:
|
|
db_data["storage_url"] = hidden_params["storage_url"]
|
|
|
|
verbose_logger.debug(
|
|
f"Storage metadata: storage_backend={db_data.get('storage_backend')}, "
|
|
f"storage_url={db_data.get('storage_url')}"
|
|
)
|
|
|
|
result = await self.prisma_client.db.litellm_managedfiletable.create(
|
|
data=db_data
|
|
)
|
|
verbose_logger.debug(
|
|
f"LiteLLM Managed File object with id={file_id} stored in db: {result}"
|
|
)
|
|
|
|
async def store_unified_object_id(
|
|
self,
|
|
unified_object_id: str,
|
|
file_object: Union[LiteLLMBatch, LiteLLMFineTuningJob, "ResponsesAPIResponse"],
|
|
litellm_parent_otel_span: Optional[Span],
|
|
model_object_id: str,
|
|
file_purpose: Literal["batch", "fine-tune", "response"],
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
) -> None:
|
|
verbose_logger.info(
|
|
f"Storing LiteLLM Managed {file_purpose} object with id={unified_object_id} in cache"
|
|
)
|
|
litellm_managed_object = LiteLLM_ManagedObjectTable(
|
|
unified_object_id=unified_object_id,
|
|
model_object_id=model_object_id,
|
|
file_purpose=file_purpose,
|
|
file_object=file_object,
|
|
)
|
|
await self.internal_usage_cache.async_set_cache(
|
|
key=unified_object_id,
|
|
value=litellm_managed_object.model_dump(),
|
|
litellm_parent_otel_span=litellm_parent_otel_span,
|
|
)
|
|
|
|
await self.prisma_client.db.litellm_managedobjecttable.upsert(
|
|
where={"unified_object_id": unified_object_id},
|
|
data={
|
|
"create": {
|
|
"unified_object_id": unified_object_id,
|
|
"file_object": file_object.model_dump_json(),
|
|
"model_object_id": model_object_id,
|
|
"file_purpose": file_purpose,
|
|
"created_by": user_api_key_dict.user_id,
|
|
"team_id": user_api_key_dict.team_id,
|
|
"updated_by": user_api_key_dict.user_id,
|
|
"status": file_object.status,
|
|
},
|
|
"update": {
|
|
"file_object": file_object.model_dump_json(),
|
|
"status": file_object.status,
|
|
"updated_by": user_api_key_dict.user_id,
|
|
}, # FIX: Update status and file_object on every operation to keep state in sync
|
|
},
|
|
)
|
|
|
|
async def get_unified_file_id(
|
|
self, file_id: str, litellm_parent_otel_span: Optional[Span] = None
|
|
) -> Optional[LiteLLM_ManagedFileTable]:
|
|
## CHECK CACHE
|
|
result = cast(
|
|
Optional[dict],
|
|
await self.internal_usage_cache.async_get_cache(
|
|
key=file_id,
|
|
litellm_parent_otel_span=litellm_parent_otel_span,
|
|
),
|
|
)
|
|
|
|
if result:
|
|
return LiteLLM_ManagedFileTable(**result)
|
|
|
|
## CHECK DB
|
|
db_object = await self.prisma_client.db.litellm_managedfiletable.find_first(
|
|
where={"unified_file_id": file_id}
|
|
)
|
|
|
|
if db_object:
|
|
return LiteLLM_ManagedFileTable(**db_object.model_dump())
|
|
return None
|
|
|
|
async def delete_unified_file_id(
|
|
self, file_id: str, litellm_parent_otel_span: Optional[Span] = None
|
|
) -> OpenAIFileObject:
|
|
## get old value
|
|
initial_value = await self.prisma_client.db.litellm_managedfiletable.find_first(
|
|
where={"unified_file_id": file_id}
|
|
)
|
|
if initial_value is None:
|
|
raise Exception(f"LiteLLM Managed File object with id={file_id} not found")
|
|
## delete old value
|
|
await self.internal_usage_cache.async_set_cache(
|
|
key=file_id,
|
|
value=None,
|
|
litellm_parent_otel_span=litellm_parent_otel_span,
|
|
)
|
|
await self.prisma_client.db.litellm_managedfiletable.delete(
|
|
where={"unified_file_id": file_id}
|
|
)
|
|
return initial_value.file_object
|
|
|
|
async def can_user_call_unified_file_id(
|
|
self, unified_file_id: str, user_api_key_dict: UserAPIKeyAuth
|
|
) -> bool:
|
|
managed_file = await self.prisma_client.db.litellm_managedfiletable.find_first(
|
|
where={"unified_file_id": unified_file_id}
|
|
)
|
|
|
|
if managed_file:
|
|
return can_access_resource(
|
|
user_api_key_dict=user_api_key_dict,
|
|
created_by=managed_file.created_by,
|
|
resource_team_id=managed_file.team_id,
|
|
)
|
|
raise HTTPException(
|
|
status_code=404,
|
|
detail=f"File not found: {unified_file_id}",
|
|
)
|
|
|
|
async def can_user_call_unified_object_id(
|
|
self, unified_object_id: str, user_api_key_dict: UserAPIKeyAuth
|
|
) -> bool:
|
|
managed_object = (
|
|
await self.prisma_client.db.litellm_managedobjecttable.find_first(
|
|
where={"unified_object_id": unified_object_id}
|
|
)
|
|
)
|
|
|
|
if managed_object:
|
|
return can_access_resource(
|
|
user_api_key_dict=user_api_key_dict,
|
|
created_by=managed_object.created_by,
|
|
resource_team_id=managed_object.team_id,
|
|
)
|
|
raise HTTPException(
|
|
status_code=404,
|
|
detail=f"Object not found: {unified_object_id}",
|
|
)
|
|
|
|
async def list_user_batches(
|
|
self,
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
limit: Optional[int] = None,
|
|
after: Optional[str] = None,
|
|
provider: Optional[str] = None,
|
|
target_model_names: Optional[str] = None,
|
|
llm_router: Optional[Router] = None,
|
|
) -> Dict[str, Any]:
|
|
# Provider filtering is not supported for managed batches
|
|
# This is because the encoded object ids stored in the managed objects table do not contain the provider information
|
|
# To support provider filtering, we would need to store the provider information in the encoded object ids
|
|
if provider:
|
|
raise Exception(
|
|
"Filtering by 'provider' is not supported when using managed batches."
|
|
)
|
|
|
|
# Model name filtering is not supported for managed batches
|
|
# This is because the encoded object ids stored in the managed objects table do not contain the model name
|
|
# A hash of the model name + litellm_params for the model name is encoded as the model id. This is not sufficient to reliably map the target model names to the model ids.
|
|
if target_model_names:
|
|
raise Exception(
|
|
"Filtering by 'target_model_names' is not supported when using managed batches."
|
|
)
|
|
|
|
owner_filter = build_owner_filter(user_api_key_dict)
|
|
if owner_filter is None:
|
|
return build_list_page([])
|
|
|
|
where_clause: Dict[str, Any] = {"file_purpose": "batch", **owner_filter}
|
|
|
|
if after:
|
|
where_clause["id"] = {"gt": after}
|
|
|
|
fetch_limit = limit or 20
|
|
if target_model_names:
|
|
# Oversample so post-fetch model-name filtering still has enough rows.
|
|
fetch_limit = max(fetch_limit * 3, 100)
|
|
|
|
batches = await self.prisma_client.db.litellm_managedobjecttable.find_many(
|
|
where=where_clause,
|
|
take=fetch_limit,
|
|
order={"created_at": "desc"},
|
|
)
|
|
|
|
batch_objects: List[LiteLLMBatch] = []
|
|
for batch in batches:
|
|
try:
|
|
# Stop once we have enough after filtering
|
|
if len(batch_objects) >= (limit or 20):
|
|
break
|
|
|
|
batch_data = (
|
|
json.loads(batch.file_object)
|
|
if isinstance(batch.file_object, str)
|
|
else batch.file_object
|
|
)
|
|
batch_obj = LiteLLMBatch(**batch_data)
|
|
batch_obj.id = batch.unified_object_id
|
|
batch_objects.append(batch_obj)
|
|
|
|
except Exception as e:
|
|
verbose_logger.warning(
|
|
f"Failed to parse batch object {batch.unified_object_id}: {e}"
|
|
)
|
|
continue
|
|
|
|
return build_list_page(
|
|
batch_objects, has_more=len(batch_objects) == (limit or 20)
|
|
)
|
|
|
|
async def get_user_created_file_ids(
|
|
self, user_api_key_dict: UserAPIKeyAuth, model_object_ids: List[str]
|
|
) -> List[OpenAIFileObject]:
|
|
"""
|
|
Get all file ids the caller is allowed to see for a list of model
|
|
object ids. Service-account keys (no user_id) are scoped to their
|
|
team via ``team_id``; admins see all matches.
|
|
|
|
Returns:
|
|
- List of OpenAIFileObject's
|
|
"""
|
|
owner_filter = build_owner_filter(user_api_key_dict)
|
|
if owner_filter is None:
|
|
return []
|
|
|
|
file_ids = await self.prisma_client.db.litellm_managedfiletable.find_many(
|
|
where={
|
|
**owner_filter,
|
|
"flat_model_file_ids": {"hasSome": model_object_ids},
|
|
}
|
|
)
|
|
return [OpenAIFileObject(**file_object.file_object) for file_object in file_ids]
|
|
|
|
async def check_managed_file_id_access(
|
|
self, data: Dict, user_api_key_dict: UserAPIKeyAuth
|
|
) -> bool:
|
|
retrieve_file_id = cast(Optional[str], data.get("file_id"))
|
|
potential_file_id = (
|
|
_is_base64_encoded_unified_file_id(retrieve_file_id)
|
|
if retrieve_file_id
|
|
else False
|
|
)
|
|
if potential_file_id and retrieve_file_id:
|
|
if await self.can_user_call_unified_file_id(
|
|
retrieve_file_id, user_api_key_dict
|
|
):
|
|
return True
|
|
else:
|
|
raise HTTPException(
|
|
status_code=403,
|
|
detail=f"User {user_api_key_dict.user_id} does not have access to the file {retrieve_file_id}",
|
|
)
|
|
return False
|
|
|
|
async def check_file_ids_access(
|
|
self, file_ids: List[str], user_api_key_dict: UserAPIKeyAuth
|
|
) -> None:
|
|
"""
|
|
Check if the user has access to a list of file IDs.
|
|
Only checks managed (unified) file IDs.
|
|
|
|
Args:
|
|
file_ids: List of file IDs to check access for
|
|
user_api_key_dict: User API key authentication details
|
|
|
|
Raises:
|
|
HTTPException: If user doesn't have access to any of the files
|
|
"""
|
|
for file_id in file_ids:
|
|
is_unified_file_id = _is_base64_encoded_unified_file_id(file_id)
|
|
if is_unified_file_id:
|
|
if not await self.can_user_call_unified_file_id(
|
|
file_id, user_api_key_dict
|
|
):
|
|
raise HTTPException(
|
|
status_code=403,
|
|
detail=f"User {user_api_key_dict.user_id} does not have access to the file {file_id}",
|
|
)
|
|
|
|
async def async_pre_call_hook(
|
|
self,
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
cache: DualCache,
|
|
data: Dict,
|
|
call_type: CallTypesLiteral,
|
|
) -> Union[Exception, str, Dict, None]:
|
|
"""
|
|
- Detect litellm_proxy/ file_id
|
|
- add dictionary of mappings of litellm_proxy/ file_id -> provider_file_id => {litellm_proxy/file_id: {"model_id": id, "file_id": provider_file_id}}
|
|
"""
|
|
### HANDLE FILE ACCESS ### - ensure user has access to the file
|
|
if (
|
|
call_type == CallTypes.afile_content.value
|
|
or call_type == CallTypes.afile_delete.value
|
|
or call_type == CallTypes.afile_retrieve.value
|
|
or call_type == CallTypes.afile_content.value
|
|
):
|
|
await self.check_managed_file_id_access(data, user_api_key_dict)
|
|
|
|
### HANDLE TRANSFORMATIONS ###
|
|
# Check both completion and acompletion call types
|
|
is_completion_call = (
|
|
call_type == CallTypes.completion.value
|
|
or call_type == CallTypes.acompletion.value
|
|
)
|
|
|
|
if is_completion_call:
|
|
messages = data.get("messages")
|
|
model = data.get("model", "")
|
|
if messages:
|
|
file_ids = self.get_file_ids_from_messages(messages)
|
|
if file_ids:
|
|
# Check user has access to all managed files
|
|
await self.check_file_ids_access(file_ids, user_api_key_dict)
|
|
|
|
# Check if any files are stored in storage backends and need base64 conversion
|
|
# This is needed for Vertex AI/Gemini which requires base64 content
|
|
is_vertex_ai = model and (
|
|
"vertex_ai" in model or "gemini" in model.lower()
|
|
)
|
|
if is_vertex_ai:
|
|
await self._convert_storage_files_to_base64(
|
|
messages=messages,
|
|
file_ids=file_ids,
|
|
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
|
|
)
|
|
|
|
model_file_id_mapping = await self.get_model_file_id_mapping(
|
|
file_ids, user_api_key_dict.parent_otel_span
|
|
)
|
|
data["model_file_id_mapping"] = model_file_id_mapping
|
|
elif (
|
|
call_type == CallTypes.aresponses.value
|
|
or call_type == CallTypes.responses.value
|
|
):
|
|
# Handle managed files in responses API input and tools
|
|
file_ids = []
|
|
|
|
# Extract file IDs from input parameter
|
|
input_data = data.get("input")
|
|
if input_data:
|
|
file_ids.extend(self.get_file_ids_from_responses_input(input_data))
|
|
|
|
# Extract file IDs from tools parameter (e.g., code_interpreter container)
|
|
tools = data.get("tools")
|
|
if tools:
|
|
file_ids.extend(self.get_file_ids_from_responses_tools(tools))
|
|
|
|
if file_ids:
|
|
# Check user has access to all managed files
|
|
await self.check_file_ids_access(file_ids, user_api_key_dict)
|
|
|
|
model_file_id_mapping = await self.get_model_file_id_mapping(
|
|
file_ids, user_api_key_dict.parent_otel_span
|
|
)
|
|
data["model_file_id_mapping"] = model_file_id_mapping
|
|
|
|
# Check access for file_search vector_store_ids
|
|
if tools:
|
|
unified_vs_ids = self.get_vector_store_ids_from_file_search_tools(tools)
|
|
if unified_vs_ids:
|
|
await self.check_vector_store_ids_access(
|
|
unified_vs_ids, user_api_key_dict
|
|
)
|
|
elif call_type == CallTypes.afile_content.value:
|
|
retrieve_file_id = cast(Optional[str], data.get("file_id"))
|
|
potential_file_id = (
|
|
_is_base64_encoded_unified_file_id(retrieve_file_id)
|
|
if retrieve_file_id
|
|
else False
|
|
)
|
|
if potential_file_id and "llm_output_file_id," in potential_file_id:
|
|
model_id = self.get_model_id_from_unified_file_id(potential_file_id)
|
|
if model_id:
|
|
data["model"] = model_id
|
|
data["file_id"] = self.get_output_file_id_from_unified_file_id(
|
|
potential_file_id
|
|
)
|
|
elif call_type == CallTypes.acreate_batch.value:
|
|
input_file_id = cast(Optional[str], data.get("input_file_id"))
|
|
if input_file_id:
|
|
model_file_id_mapping = await self.get_model_file_id_mapping(
|
|
[input_file_id], user_api_key_dict.parent_otel_span
|
|
)
|
|
|
|
data["model_file_id_mapping"] = model_file_id_mapping
|
|
elif (
|
|
call_type == CallTypes.aretrieve_batch.value
|
|
or call_type == CallTypes.acancel_batch.value
|
|
or call_type == CallTypes.acancel_fine_tuning_job.value
|
|
or call_type == CallTypes.aretrieve_fine_tuning_job.value
|
|
):
|
|
accessor_key: Optional[str] = None
|
|
retrieve_object_id: Optional[str] = None
|
|
if (
|
|
call_type == CallTypes.aretrieve_batch.value
|
|
or call_type == CallTypes.acancel_batch.value
|
|
):
|
|
accessor_key = "batch_id"
|
|
elif (
|
|
call_type == CallTypes.acancel_fine_tuning_job.value
|
|
or call_type == CallTypes.aretrieve_fine_tuning_job.value
|
|
):
|
|
accessor_key = "fine_tuning_job_id"
|
|
|
|
if accessor_key:
|
|
retrieve_object_id = cast(Optional[str], data.get(accessor_key))
|
|
|
|
potential_llm_object_id = (
|
|
_is_base64_encoded_unified_file_id(retrieve_object_id)
|
|
if retrieve_object_id
|
|
else False
|
|
)
|
|
if potential_llm_object_id and retrieve_object_id:
|
|
## VALIDATE USER HAS ACCESS TO THE OBJECT ##
|
|
if not await self.can_user_call_unified_object_id(
|
|
retrieve_object_id, user_api_key_dict
|
|
):
|
|
raise HTTPException(
|
|
status_code=403,
|
|
detail=f"User {user_api_key_dict.user_id} does not have access to the object {retrieve_object_id}",
|
|
)
|
|
|
|
## for managed batch id - get the model id
|
|
potential_model_id = get_model_id_from_unified_batch_id(
|
|
potential_llm_object_id
|
|
)
|
|
if potential_model_id is None:
|
|
raise Exception(
|
|
f"LiteLLM Managed {accessor_key} with id={retrieve_object_id} is invalid - does not contain encoded model_id."
|
|
)
|
|
data["model"] = potential_model_id
|
|
data[accessor_key] = get_batch_id_from_unified_batch_id(
|
|
potential_llm_object_id
|
|
)
|
|
elif call_type == CallTypes.acreate_fine_tuning_job.value:
|
|
input_file_id = cast(Optional[str], data.get("training_file"))
|
|
if input_file_id:
|
|
model_file_id_mapping = await self.get_model_file_id_mapping(
|
|
[input_file_id], user_api_key_dict.parent_otel_span
|
|
)
|
|
|
|
return data
|
|
|
|
async def async_filter_deployments(
|
|
self,
|
|
model: str,
|
|
healthy_deployments: List,
|
|
messages: Optional[List[AllMessageValues]],
|
|
request_kwargs: Optional[Dict] = None,
|
|
parent_otel_span: Optional[Span] = None,
|
|
) -> List[Dict]:
|
|
if request_kwargs is None:
|
|
return healthy_deployments
|
|
|
|
input_file_id = cast(Optional[str], request_kwargs.get("input_file_id"))
|
|
model_file_id_mapping = cast(
|
|
Optional[Dict[str, Dict[str, str]]],
|
|
request_kwargs.get("model_file_id_mapping"),
|
|
)
|
|
allowed_model_ids = []
|
|
if input_file_id and model_file_id_mapping:
|
|
model_id_dict = model_file_id_mapping.get(input_file_id, {})
|
|
allowed_model_ids = list(model_id_dict.keys())
|
|
|
|
if len(allowed_model_ids) == 0:
|
|
return healthy_deployments
|
|
|
|
return [
|
|
deployment
|
|
for deployment in healthy_deployments
|
|
if deployment.get("model_info", {}).get("id") in allowed_model_ids
|
|
]
|
|
|
|
async def async_pre_call_deployment_hook(
|
|
self, kwargs: Dict[str, Any], call_type: Optional[CallTypes]
|
|
) -> Optional[dict]:
|
|
"""
|
|
Allow modifying the request just before it's sent to the deployment.
|
|
"""
|
|
accessor_key: Optional[str] = None
|
|
if call_type and call_type == CallTypes.acreate_batch:
|
|
accessor_key = "input_file_id"
|
|
elif call_type and call_type == CallTypes.acreate_fine_tuning_job:
|
|
accessor_key = "training_file"
|
|
else:
|
|
return kwargs
|
|
|
|
if accessor_key:
|
|
input_file_id = cast(Optional[str], kwargs.get(accessor_key))
|
|
model_file_id_mapping = cast(
|
|
Optional[Dict[str, Dict[str, str]]], kwargs.get("model_file_id_mapping")
|
|
)
|
|
# model_info may be at top-level or nested under litellm_metadata
|
|
# (batch/file operations use litellm_metadata)
|
|
model_id = cast(Optional[str], kwargs.get("model_info", {}).get("id", None))
|
|
if model_id is None:
|
|
model_id = cast(
|
|
Optional[str],
|
|
kwargs.get("litellm_metadata", {})
|
|
.get("model_info", {})
|
|
.get("id", None),
|
|
)
|
|
mapped_file_id: Optional[str] = None
|
|
if input_file_id and model_file_id_mapping and model_id:
|
|
mapped_file_id = model_file_id_mapping.get(input_file_id, {}).get(
|
|
model_id, None
|
|
)
|
|
if mapped_file_id:
|
|
kwargs[accessor_key] = mapped_file_id
|
|
|
|
return kwargs
|
|
|
|
def get_file_ids_from_messages(self, messages: List[AllMessageValues]) -> List[str]:
|
|
"""
|
|
Gets file ids from messages
|
|
"""
|
|
file_ids = []
|
|
for message in messages:
|
|
if message.get("role") == "user":
|
|
content = message.get("content")
|
|
if content:
|
|
if isinstance(content, str):
|
|
continue
|
|
for c in content:
|
|
if c.get("type") == "file":
|
|
file_object = cast(ChatCompletionFileObject, c)
|
|
file_object_file_field = file_object["file"]
|
|
file_id = file_object_file_field.get("file_id")
|
|
if file_id:
|
|
file_ids.append(file_id)
|
|
return file_ids
|
|
|
|
def get_file_ids_from_responses_input(
|
|
self, input: Union[str, List[Dict[str, Any]]]
|
|
) -> List[str]:
|
|
"""
|
|
Gets file ids from responses API input.
|
|
|
|
The input can be:
|
|
- A string (no files)
|
|
- A list of input items, where each item can have:
|
|
- type: "input_file" with file_id
|
|
- content: a list that can contain items with type: "input_file" and file_id
|
|
"""
|
|
file_ids: List[str] = []
|
|
|
|
if isinstance(input, str):
|
|
return file_ids
|
|
|
|
if not isinstance(input, list):
|
|
return file_ids
|
|
|
|
for item in input:
|
|
if not isinstance(item, dict):
|
|
continue
|
|
|
|
# Check for direct input_file type
|
|
if item.get("type") == "input_file":
|
|
file_id = item.get("file_id")
|
|
if file_id:
|
|
file_ids.append(file_id)
|
|
|
|
# Check for input_file in content array
|
|
content = item.get("content")
|
|
if isinstance(content, list):
|
|
for content_item in content:
|
|
if (
|
|
isinstance(content_item, dict)
|
|
and content_item.get("type") == "input_file"
|
|
):
|
|
file_id = content_item.get("file_id")
|
|
if file_id:
|
|
file_ids.append(file_id)
|
|
|
|
return file_ids
|
|
|
|
def get_file_ids_from_responses_tools(
|
|
self, tools: List[Dict[str, Any]]
|
|
) -> List[str]:
|
|
"""
|
|
Gets file ids from responses API tools parameter.
|
|
|
|
The tools can contain code_interpreter with container.file_ids:
|
|
[
|
|
{
|
|
"type": "code_interpreter",
|
|
"container": {"type": "auto", "file_ids": ["file-123", "file-456"]}
|
|
}
|
|
]
|
|
"""
|
|
file_ids: List[str] = []
|
|
|
|
if not isinstance(tools, list):
|
|
return file_ids
|
|
|
|
for tool in tools:
|
|
if not isinstance(tool, dict):
|
|
continue
|
|
|
|
# Check for code_interpreter with container file_ids
|
|
if tool.get("type") == "code_interpreter":
|
|
container = tool.get("container")
|
|
if isinstance(container, dict):
|
|
container_file_ids = container.get("file_ids")
|
|
if isinstance(container_file_ids, list):
|
|
for file_id in container_file_ids:
|
|
if isinstance(file_id, str):
|
|
file_ids.append(file_id)
|
|
|
|
return file_ids
|
|
|
|
def get_vector_store_ids_from_file_search_tools(
|
|
self, tools: List[Dict[str, Any]]
|
|
) -> List[str]:
|
|
"""
|
|
Extract unified vector_store_ids from file_search tools.
|
|
|
|
Only returns IDs that are LiteLLM-managed (base64 unified IDs).
|
|
Native provider IDs are skipped — they have no LiteLLM access record.
|
|
"""
|
|
from litellm.llms.base_llm.managed_resources.utils import (
|
|
is_base64_encoded_unified_id,
|
|
)
|
|
|
|
vs_ids: List[str] = []
|
|
if not isinstance(tools, list):
|
|
return vs_ids
|
|
|
|
for tool in tools:
|
|
if not isinstance(tool, dict) or tool.get("type") != "file_search":
|
|
continue
|
|
vector_store_ids = tool.get("vector_store_ids")
|
|
if not isinstance(vector_store_ids, list):
|
|
continue
|
|
for vs_id in vector_store_ids:
|
|
if isinstance(vs_id, str) and is_base64_encoded_unified_id(vs_id):
|
|
vs_ids.append(vs_id)
|
|
|
|
return vs_ids
|
|
|
|
async def check_vector_store_ids_access(
|
|
self,
|
|
vector_store_ids: List[str],
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
) -> None:
|
|
"""
|
|
Verify the caller's team can access each LiteLLM-managed vector store.
|
|
|
|
Batch-fetches vector stores from DB and checks team_id.
|
|
Raises HTTPException(403) on the first access violation.
|
|
Non-managed (native) IDs should already be filtered out before calling this.
|
|
"""
|
|
from litellm.llms.base_llm.managed_resources.utils import (
|
|
extract_unified_uuid_from_unified_id,
|
|
)
|
|
from litellm.proxy.auth.auth_checks import (
|
|
get_managed_vector_store_rows_by_uuids,
|
|
)
|
|
from litellm.proxy.proxy_server import (
|
|
prisma_client,
|
|
proxy_logging_obj,
|
|
user_api_key_cache,
|
|
)
|
|
|
|
if not vector_store_ids or prisma_client is None:
|
|
return
|
|
|
|
# Map each unified ID to its internal UUID for a single batch DB fetch
|
|
uuid_to_unified: Dict[str, str] = {}
|
|
for vs_id in vector_store_ids:
|
|
uuid = extract_unified_uuid_from_unified_id(vs_id)
|
|
if uuid:
|
|
uuid_to_unified[uuid] = vs_id
|
|
|
|
if not uuid_to_unified:
|
|
return
|
|
|
|
rows = await get_managed_vector_store_rows_by_uuids(
|
|
uuids=list(uuid_to_unified.keys()),
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
|
|
found_uuids = {row.vector_store_id for row in rows}
|
|
|
|
for uuid, original_id in uuid_to_unified.items():
|
|
if uuid not in found_uuids:
|
|
raise HTTPException(
|
|
status_code=403,
|
|
detail=f"Vector store '{original_id}' not found or access denied.",
|
|
)
|
|
|
|
caller_team_id = user_api_key_dict.team_id
|
|
for row in rows:
|
|
vs_team_id = getattr(row, "team_id", None)
|
|
if vs_team_id is not None and vs_team_id != caller_team_id:
|
|
raise HTTPException(
|
|
status_code=403,
|
|
detail=(
|
|
f"Team '{caller_team_id}' does not have access to vector "
|
|
f"store '{row.vector_store_id}'. The store belongs to team "
|
|
f"'{vs_team_id}'."
|
|
),
|
|
)
|
|
|
|
async def get_model_file_id_mapping(
|
|
self, file_ids: List[str], litellm_parent_otel_span: Span
|
|
) -> dict:
|
|
"""
|
|
Get model-specific file IDs for a list of proxy file IDs.
|
|
Returns a dictionary mapping litellm_proxy/ file_id -> model_id -> model_file_id
|
|
|
|
1. Get all the litellm_proxy/ file_ids from the messages
|
|
2. For each file_id, search for cache keys matching the pattern file_id:*
|
|
3. Return a dictionary of mappings of litellm_proxy/ file_id -> model_id -> model_file_id
|
|
|
|
Example:
|
|
{
|
|
"litellm_proxy/file_id": {
|
|
"model_id": "model_file_id"
|
|
}
|
|
}
|
|
"""
|
|
|
|
file_id_mapping: Dict[str, Dict[str, str]] = {}
|
|
litellm_managed_file_ids = []
|
|
|
|
for file_id in file_ids:
|
|
## CHECK IF FILE ID IS MANAGED BY LITELM
|
|
is_base64_unified_file_id = _is_base64_encoded_unified_file_id(file_id)
|
|
if is_base64_unified_file_id:
|
|
litellm_managed_file_ids.append(file_id)
|
|
|
|
if litellm_managed_file_ids:
|
|
# Get all cache keys matching the pattern file_id:*
|
|
for file_id in litellm_managed_file_ids:
|
|
# Search for any cache key starting with this file_id
|
|
unified_file_object = await self.get_unified_file_id(
|
|
file_id, litellm_parent_otel_span
|
|
)
|
|
|
|
if unified_file_object:
|
|
file_id_mapping[file_id] = unified_file_object.model_mappings
|
|
|
|
return file_id_mapping
|
|
|
|
async def create_file_for_each_model(
|
|
self,
|
|
llm_router: Optional[Router],
|
|
_create_file_request: CreateFileRequest,
|
|
target_model_names_list: List[str],
|
|
litellm_parent_otel_span: Span,
|
|
) -> List[OpenAIFileObject]:
|
|
if llm_router is None:
|
|
raise Exception("LLM Router not initialized. Ensure models added to proxy.")
|
|
responses = []
|
|
for model in target_model_names_list:
|
|
individual_response = await llm_router.acreate_file(
|
|
model=model, **_create_file_request
|
|
)
|
|
responses.append(individual_response)
|
|
|
|
return responses
|
|
|
|
async def acreate_file(
|
|
self,
|
|
create_file_request: CreateFileRequest,
|
|
llm_router: Router,
|
|
target_model_names_list: List[str],
|
|
litellm_parent_otel_span: Span,
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
) -> OpenAIFileObject:
|
|
responses = await self.create_file_for_each_model(
|
|
llm_router=llm_router,
|
|
_create_file_request=create_file_request,
|
|
target_model_names_list=target_model_names_list,
|
|
litellm_parent_otel_span=litellm_parent_otel_span,
|
|
)
|
|
response = await _PROXY_LiteLLMManagedFiles.return_unified_file_id(
|
|
file_objects=responses,
|
|
create_file_request=create_file_request,
|
|
internal_usage_cache=self.internal_usage_cache,
|
|
litellm_parent_otel_span=litellm_parent_otel_span,
|
|
target_model_names_list=target_model_names_list,
|
|
)
|
|
|
|
## STORE MODEL MAPPINGS IN DB
|
|
model_mappings: Dict[str, str] = {}
|
|
|
|
for file_object in responses:
|
|
model_file_id_mapping = file_object._hidden_params.get(
|
|
"model_file_id_mapping"
|
|
)
|
|
if model_file_id_mapping and isinstance(model_file_id_mapping, dict):
|
|
model_mappings.update(model_file_id_mapping)
|
|
|
|
await self.store_unified_file_id(
|
|
file_id=response.id,
|
|
file_object=response,
|
|
litellm_parent_otel_span=litellm_parent_otel_span,
|
|
model_mappings=model_mappings,
|
|
user_api_key_dict=user_api_key_dict,
|
|
)
|
|
|
|
# Emit Prometheus metrics for managed file creation
|
|
prom_logger = self._get_prometheus_logger()
|
|
if prom_logger:
|
|
first_model = (
|
|
target_model_names_list[0] if target_model_names_list else None
|
|
)
|
|
first_provider = ""
|
|
if responses:
|
|
first_provider = (
|
|
getattr(responses[0], "_hidden_params", {}).get(
|
|
"custom_llm_provider"
|
|
)
|
|
or ""
|
|
)
|
|
prom_logger.record_managed_file_created(
|
|
model=first_model or "",
|
|
api_provider=first_provider,
|
|
user=user_api_key_dict.user_id or "",
|
|
user_email=getattr(user_api_key_dict, "user_email", None) or "",
|
|
api_key_alias=user_api_key_dict.key_alias or "",
|
|
)
|
|
if response.bytes and response.bytes > 0:
|
|
prom_logger.record_managed_file_size(
|
|
size_bytes=response.bytes,
|
|
purpose=response.purpose or "batch",
|
|
file_type="input",
|
|
model=first_model,
|
|
api_provider=first_provider,
|
|
user=user_api_key_dict.user_id,
|
|
)
|
|
|
|
return response
|
|
|
|
@staticmethod
|
|
async def return_unified_file_id(
|
|
file_objects: List[OpenAIFileObject],
|
|
create_file_request: CreateFileRequest,
|
|
internal_usage_cache: InternalUsageCache,
|
|
litellm_parent_otel_span: Span,
|
|
target_model_names_list: List[str],
|
|
) -> OpenAIFileObject:
|
|
## GET THE FILE TYPE FROM THE CREATE FILE REQUEST
|
|
file_data = extract_file_data(create_file_request["file"])
|
|
|
|
file_type = file_data["content_type"]
|
|
|
|
output_file_id = file_objects[0].id
|
|
model_id = file_objects[0]._hidden_params.get("model_id")
|
|
|
|
unified_file_id = SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format(
|
|
file_type,
|
|
str(uuid.uuid4()),
|
|
",".join(target_model_names_list),
|
|
output_file_id,
|
|
model_id,
|
|
)
|
|
|
|
# Convert to URL-safe base64 and strip padding
|
|
base64_unified_file_id = (
|
|
base64.urlsafe_b64encode(unified_file_id.encode()).decode().rstrip("=")
|
|
)
|
|
|
|
## CREATE RESPONSE OBJECT
|
|
|
|
response = OpenAIFileObject(
|
|
id=base64_unified_file_id,
|
|
object="file",
|
|
purpose=create_file_request["purpose"],
|
|
created_at=file_objects[0].created_at,
|
|
bytes=file_objects[0].bytes,
|
|
filename=file_objects[0].filename,
|
|
status="uploaded",
|
|
expires_at=file_objects[0].expires_at,
|
|
)
|
|
|
|
return response
|
|
|
|
def get_unified_generic_response_id(
|
|
self, model_id: str, generic_response_id: str
|
|
) -> str:
|
|
unified_generic_response_id = (
|
|
SpecialEnums.LITELLM_MANAGED_GENERIC_RESPONSE_COMPLETE_STR.value.format(
|
|
model_id, generic_response_id
|
|
)
|
|
)
|
|
return (
|
|
base64.urlsafe_b64encode(unified_generic_response_id.encode())
|
|
.decode()
|
|
.rstrip("=")
|
|
)
|
|
|
|
def get_unified_batch_id(self, batch_id: str, model_id: str) -> str:
|
|
unified_batch_id = SpecialEnums.LITELLM_MANAGED_BATCH_COMPLETE_STR.value.format(
|
|
model_id, batch_id
|
|
)
|
|
return base64.urlsafe_b64encode(unified_batch_id.encode()).decode().rstrip("=")
|
|
|
|
def get_unified_output_file_id(
|
|
self, output_file_id: str, model_id: str, model_name: Optional[str]
|
|
) -> str:
|
|
unified_output_file_id = (
|
|
SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format(
|
|
"application/json",
|
|
str(uuid.uuid4()),
|
|
model_name or "",
|
|
output_file_id,
|
|
model_id,
|
|
)
|
|
)
|
|
return (
|
|
base64.urlsafe_b64encode(unified_output_file_id.encode())
|
|
.decode()
|
|
.rstrip("=")
|
|
)
|
|
|
|
def get_model_id_from_unified_file_id(self, file_id: str) -> str:
|
|
return file_id.split("llm_output_file_model_id,")[1].split(";")[0]
|
|
|
|
def get_output_file_id_from_unified_file_id(self, file_id: str) -> str:
|
|
marker = "llm_output_file_id,"
|
|
if marker not in file_id:
|
|
raise ValueError(
|
|
f"Unified id does not contain {marker!r}: {file_id[:80]!r}"
|
|
)
|
|
return file_id.split(marker, 1)[1].split(";")[0]
|
|
|
|
async def async_post_call_success_hook(
|
|
self, data: Dict, user_api_key_dict: UserAPIKeyAuth, response: LLMResponseTypes
|
|
) -> Any:
|
|
if isinstance(response, LiteLLMBatch):
|
|
## Check if unified_file_id is in the response
|
|
unified_file_id = response._hidden_params.get(
|
|
"unified_file_id"
|
|
) # managed file id
|
|
unified_batch_id = response._hidden_params.get(
|
|
"unified_batch_id"
|
|
) # managed batch id
|
|
model_id = cast(Optional[str], response._hidden_params.get("model_id"))
|
|
model_name = cast(Optional[str], response._hidden_params.get("model_name"))
|
|
resolved_model_name = model_name
|
|
|
|
# Some providers (e.g. Vertex batch retrieve) do not set model_name on
|
|
# the response. In that case, recover target_model_names from the input
|
|
# managed file metadata so unified output IDs preserve routing metadata.
|
|
if not resolved_model_name and isinstance(unified_file_id, str):
|
|
decoded_unified_file_id = (
|
|
_is_base64_encoded_unified_file_id(unified_file_id)
|
|
or unified_file_id
|
|
)
|
|
target_model_names = get_models_from_unified_file_id(
|
|
decoded_unified_file_id
|
|
)
|
|
if target_model_names:
|
|
resolved_model_name = ",".join(target_model_names)
|
|
original_response_id = response.id
|
|
|
|
if (unified_batch_id or unified_file_id) and model_id:
|
|
response.id = self.get_unified_batch_id(
|
|
batch_id=response.id, model_id=model_id
|
|
)
|
|
|
|
# Handle both output_file_id and error_file_id
|
|
for file_attr in ["output_file_id", "error_file_id"]:
|
|
file_id_value = getattr(response, file_attr, None)
|
|
if file_id_value and model_id:
|
|
decoded_output_file_id = _is_base64_encoded_unified_file_id(
|
|
file_id_value
|
|
)
|
|
if (
|
|
decoded_output_file_id
|
|
and "llm_output_file_id," in decoded_output_file_id
|
|
):
|
|
provider_file_id = (
|
|
self.get_output_file_id_from_unified_file_id(
|
|
decoded_output_file_id
|
|
)
|
|
)
|
|
unified_file_id = file_id_value
|
|
elif decoded_output_file_id:
|
|
verbose_logger.warning(
|
|
f"Skipping {file_attr}={file_id_value!r}: "
|
|
"unified id is not a managed file output id"
|
|
)
|
|
continue
|
|
else:
|
|
provider_file_id = file_id_value
|
|
unified_file_id = self.get_unified_output_file_id(
|
|
output_file_id=provider_file_id,
|
|
model_id=model_id,
|
|
model_name=resolved_model_name,
|
|
)
|
|
setattr(response, file_attr, unified_file_id)
|
|
|
|
# Use llm_router credentials when available. Without credentials,
|
|
# Azure and other auth-required providers return 500/401.
|
|
file_object = None
|
|
try:
|
|
# Import module and use getattr for better testability with mocks
|
|
import litellm.proxy.proxy_server as proxy_server_module
|
|
|
|
_llm_router = getattr(
|
|
proxy_server_module, "llm_router", None
|
|
)
|
|
if _llm_router is not None and model_id:
|
|
_creds = (
|
|
_llm_router.get_deployment_credentials_with_provider(
|
|
model_id
|
|
)
|
|
or {}
|
|
)
|
|
file_object = await litellm.afile_retrieve(
|
|
file_id=provider_file_id,
|
|
**_creds,
|
|
)
|
|
else:
|
|
file_object = await litellm.afile_retrieve(
|
|
custom_llm_provider=model_name.split("/")[0] if model_name and "/" in model_name else "openai", # type: ignore[arg-type]
|
|
file_id=provider_file_id,
|
|
)
|
|
verbose_logger.debug(
|
|
f"Successfully retrieved file object for {file_attr}={provider_file_id}"
|
|
)
|
|
except Exception as e:
|
|
verbose_logger.warning(
|
|
f"Failed to retrieve file object for {file_attr}={provider_file_id}: {str(e)}. Storing with None and will fetch on-demand."
|
|
)
|
|
|
|
await self.store_unified_file_id(
|
|
file_id=unified_file_id,
|
|
file_object=file_object,
|
|
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
|
|
model_mappings={model_id: provider_file_id},
|
|
user_api_key_dict=user_api_key_dict,
|
|
)
|
|
await self.store_unified_object_id(
|
|
unified_object_id=response.id,
|
|
file_object=response,
|
|
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
|
|
model_object_id=original_response_id,
|
|
file_purpose="batch",
|
|
user_api_key_dict=user_api_key_dict,
|
|
)
|
|
|
|
# Only record batch creation metric on actual create (not retrieve/cancel).
|
|
# unified_file_id in _hidden_params is only set by the create_batch endpoint.
|
|
original_unified_file_id = response._hidden_params.get("unified_file_id")
|
|
if original_unified_file_id:
|
|
prom_logger = self._get_prometheus_logger()
|
|
if prom_logger:
|
|
batch_provider = ""
|
|
if model_name:
|
|
try:
|
|
from litellm.litellm_core_utils.get_llm_provider_logic import (
|
|
get_llm_provider,
|
|
)
|
|
|
|
_, batch_provider, _, _ = get_llm_provider(model=model_name)
|
|
except Exception:
|
|
if "/" in model_name:
|
|
batch_provider = model_name.split("/")[0]
|
|
prom_logger.record_managed_batch_created(
|
|
model=model_name or "",
|
|
api_provider=batch_provider,
|
|
user=user_api_key_dict.user_id or "",
|
|
user_email=getattr(user_api_key_dict, "user_email", None) or "",
|
|
api_key_alias=user_api_key_dict.key_alias or "",
|
|
)
|
|
|
|
elif isinstance(response, LiteLLMFineTuningJob):
|
|
## Check if unified_file_id is in the response
|
|
unified_file_id = response._hidden_params.get(
|
|
"unified_file_id"
|
|
) # managed file id
|
|
unified_finetuning_job_id = response._hidden_params.get(
|
|
"unified_finetuning_job_id"
|
|
) # managed finetuning job id
|
|
model_id = cast(Optional[str], response._hidden_params.get("model_id"))
|
|
model_name = cast(Optional[str], response._hidden_params.get("model_name"))
|
|
original_response_id = response.id
|
|
if (unified_file_id or unified_finetuning_job_id) and model_id:
|
|
response.id = self.get_unified_generic_response_id(
|
|
model_id=model_id, generic_response_id=response.id
|
|
)
|
|
await self.store_unified_object_id(
|
|
unified_object_id=response.id,
|
|
file_object=response,
|
|
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
|
|
model_object_id=original_response_id,
|
|
file_purpose="fine-tune",
|
|
user_api_key_dict=user_api_key_dict,
|
|
)
|
|
elif isinstance(response, AsyncCursorPage):
|
|
"""
|
|
For listing files, filter for the ones created by the user
|
|
"""
|
|
## check if file object
|
|
if hasattr(response, "data") and isinstance(response.data, list):
|
|
if all(
|
|
isinstance(file_object, FileObject) for file_object in response.data
|
|
):
|
|
## Get all file id's
|
|
## Check which file id's were created by the user
|
|
## Filter the response to only include the files created by the user
|
|
## Return the filtered response
|
|
file_ids = [
|
|
file_object.id
|
|
for file_object in cast(List[FileObject], response.data) # type: ignore
|
|
]
|
|
user_created_file_ids = await self.get_user_created_file_ids(
|
|
user_api_key_dict, file_ids
|
|
)
|
|
## Filter the response to only include the files created by the user
|
|
response.data = user_created_file_ids # type: ignore
|
|
return response
|
|
return response
|
|
return response
|
|
|
|
async def afile_retrieve(
|
|
self, file_id: str, litellm_parent_otel_span: Optional[Span], llm_router=None
|
|
) -> OpenAIFileObject:
|
|
stored_file_object = await self.get_unified_file_id(
|
|
file_id, litellm_parent_otel_span
|
|
)
|
|
|
|
# Case 1 : This is not a managed file
|
|
if not stored_file_object:
|
|
raise Exception(f"LiteLLM Managed File object with id={file_id} not found")
|
|
|
|
# Case 2: Managed file and the file object exists in the database
|
|
# The stored file_object has the raw provider ID. Replace with the unified ID
|
|
# so callers see a consistent ID (matching Case 3 which does response.id = file_id).
|
|
if stored_file_object and stored_file_object.file_object:
|
|
# Use model_copy to ensure the ID update persists (Pydantic v2 compatibility)
|
|
response = stored_file_object.file_object.model_copy(update={"id": file_id})
|
|
return response
|
|
|
|
# Case 3: Managed file exists in the database but not the file object (for. e.g the batch task might not have run)
|
|
# So we fetch the file object from the provider. We deliberately do not store the result to avoid interfering with batch cost tracking code.
|
|
if not llm_router:
|
|
raise Exception(
|
|
f"LiteLLM Managed File object with id={file_id} has no file_object "
|
|
f"and llm_router is required to fetch from provider"
|
|
)
|
|
|
|
try:
|
|
model_id, model_file_id = next(
|
|
iter(stored_file_object.model_mappings.items())
|
|
)
|
|
credentials = (
|
|
llm_router.get_deployment_credentials_with_provider(model_id) or {}
|
|
)
|
|
response = await litellm.afile_retrieve(
|
|
file_id=model_file_id, **credentials
|
|
)
|
|
response.id = file_id # Replace with unified ID
|
|
return response
|
|
except Exception as e:
|
|
raise Exception(
|
|
f"Failed to retrieve file {file_id} from provider: {str(e)}"
|
|
) from e
|
|
|
|
async def afile_list(
|
|
self,
|
|
purpose: Optional[OpenAIFilesPurpose],
|
|
litellm_parent_otel_span: Optional[Span],
|
|
**data: Dict,
|
|
) -> List[OpenAIFileObject]:
|
|
"""Handled in files_endpoints.py"""
|
|
return []
|
|
|
|
def _is_batch_polling_enabled(self) -> bool:
|
|
"""
|
|
Check if batch cost tracking is actually enabled and running.
|
|
Returns:
|
|
bool: True if batch cost tracking is active, False otherwise
|
|
"""
|
|
try:
|
|
# Import here to avoid circular dependencies
|
|
import litellm.proxy.proxy_server as proxy_server_module
|
|
|
|
# Check if the scheduler has the batch cost checking job registered
|
|
scheduler = getattr(proxy_server_module, "scheduler", None)
|
|
if scheduler is None:
|
|
return False
|
|
|
|
# Check if the check_batch_cost_job exists in the scheduler
|
|
try:
|
|
job = scheduler.get_job("check_batch_cost_job")
|
|
if job is not None:
|
|
return True
|
|
except Exception:
|
|
# Job not found or scheduler doesn't support get_job
|
|
pass
|
|
|
|
return False
|
|
except Exception as e:
|
|
verbose_logger.warning(
|
|
f"Error checking batch polling configuration: {e}. Assuming disabled."
|
|
)
|
|
return False
|
|
|
|
async def _get_batches_referencing_file(self, file_id: str) -> List[Dict[str, Any]]:
|
|
"""
|
|
Find batches that reference this file and still need cost tracking.
|
|
Find batches that are in non-terminal state and have not yet been processed by CheckBatchCost.
|
|
Args:
|
|
file_id: The unified file ID to check
|
|
|
|
Returns:
|
|
List of batch objects referencing this file in non-terminal state
|
|
(max 10 for error message display)
|
|
"""
|
|
# Prepare list of file IDs to check (both unified and provider IDs)
|
|
file_ids_to_check = [file_id]
|
|
|
|
# Get model-specific file IDs for this unified file ID if it's a managed file
|
|
try:
|
|
model_file_id_mapping = await self.get_model_file_id_mapping(
|
|
[file_id], litellm_parent_otel_span=None
|
|
)
|
|
|
|
if model_file_id_mapping and file_id in model_file_id_mapping:
|
|
# Add all provider file IDs for this unified file
|
|
provider_file_ids = list(model_file_id_mapping[file_id].values())
|
|
file_ids_to_check.extend(provider_file_ids)
|
|
except Exception as e:
|
|
verbose_logger.debug(
|
|
f"Could not get model file ID mapping for {file_id}: {e}. "
|
|
f"Will only check unified file ID."
|
|
)
|
|
MAX_MATCHES_TO_RETURN = 10
|
|
|
|
batches = await self.prisma_client.db.litellm_managedobjecttable.find_many(
|
|
where={
|
|
"file_purpose": "batch",
|
|
"batch_processed": False,
|
|
"status": {"not_in": ["failed", "expired", "cancelled"]},
|
|
},
|
|
take=MAX_MATCHES_TO_RETURN,
|
|
order={"created_at": "desc"},
|
|
)
|
|
|
|
referencing_batches = []
|
|
for batch in batches:
|
|
try:
|
|
# Parse the batch file_object to check for file references
|
|
batch_data = (
|
|
json.loads(batch.file_object)
|
|
if isinstance(batch.file_object, str)
|
|
else batch.file_object
|
|
)
|
|
|
|
# Extract file IDs from batch
|
|
# Batches typically reference the unified file ID in input_file_id
|
|
# Output and error files are generated by the provider
|
|
input_file_id = batch_data.get("input_file_id")
|
|
output_file_id = batch_data.get("output_file_id")
|
|
error_file_id = batch_data.get("error_file_id")
|
|
|
|
referenced_file_ids = [
|
|
fid for fid in [input_file_id, output_file_id, error_file_id] if fid
|
|
]
|
|
|
|
# Check if any referenced file ID matches the file we're trying to delete
|
|
if any(ref_id in file_ids_to_check for ref_id in referenced_file_ids):
|
|
referencing_batches.append(
|
|
{
|
|
"batch_id": batch.unified_object_id,
|
|
"status": batch.status,
|
|
"created_at": batch.created_at,
|
|
}
|
|
)
|
|
except Exception as e:
|
|
verbose_logger.warning(
|
|
f"Error parsing batch object {batch.unified_object_id}: {e}"
|
|
)
|
|
continue
|
|
|
|
return referencing_batches
|
|
|
|
async def _check_file_deletion_allowed(self, file_id: str) -> None:
|
|
"""
|
|
Check if file deletion should be blocked due to batch references.
|
|
|
|
Blocks deletion if:
|
|
1. File is referenced by any batch in non-terminal state, AND
|
|
2. Batch polling is configured (user wants cost tracking)
|
|
|
|
Args:
|
|
file_id: The unified file ID to check
|
|
|
|
Raises:
|
|
HTTPException: If file deletion should be blocked
|
|
"""
|
|
# Check if batch polling is enabled
|
|
if not self._is_batch_polling_enabled():
|
|
# Batch polling not configured, allow deletion
|
|
return
|
|
|
|
# Check if file is referenced by any non-terminal batches
|
|
referencing_batches = await self._get_batches_referencing_file(file_id)
|
|
|
|
if referencing_batches:
|
|
# File is referenced by non-terminal batches and polling is enabled
|
|
MAX_BATCHES_IN_ERROR = (
|
|
5 # Limit batches shown in error message for readability
|
|
)
|
|
|
|
# Show up to MAX_BATCHES_IN_ERROR in the error message
|
|
batches_to_show = referencing_batches[:MAX_BATCHES_IN_ERROR]
|
|
batch_statuses = [
|
|
f"{b['batch_id']}: {b['status']}" for b in batches_to_show
|
|
]
|
|
|
|
# Determine the count message
|
|
count_message = f"{len(referencing_batches)}"
|
|
if (
|
|
len(referencing_batches) >= 10
|
|
): # MAX_MATCHES_TO_RETURN from _get_batches_referencing_file
|
|
count_message = "10+"
|
|
|
|
error_message = (
|
|
f"Cannot delete file {file_id}. "
|
|
f"The file is referenced by {count_message} batch(es) in non-terminal state"
|
|
)
|
|
|
|
# Add specific batch details if not too many
|
|
if len(referencing_batches) <= MAX_BATCHES_IN_ERROR:
|
|
error_message += f": {', '.join(batch_statuses)}. "
|
|
else:
|
|
error_message += f" (showing {MAX_BATCHES_IN_ERROR} most recent): {', '.join(batch_statuses)}. "
|
|
|
|
error_message += (
|
|
"To delete this file before complete cost tracking, please delete or cancel the referencing batch(es) first. "
|
|
"Alternatively, wait for all batches to complete and for cost to be computed (batch_processed=true)."
|
|
)
|
|
|
|
# Record blocked deletion metric
|
|
prom_logger = self._get_prometheus_logger()
|
|
if prom_logger:
|
|
prom_logger.record_managed_file_deleted(result="blocked")
|
|
|
|
raise HTTPException(
|
|
status_code=400,
|
|
detail=error_message,
|
|
)
|
|
|
|
async def afile_delete(
|
|
self,
|
|
file_id: str,
|
|
litellm_parent_otel_span: Optional[Span],
|
|
llm_router: Router,
|
|
**data: Dict,
|
|
) -> OpenAIFileObject:
|
|
|
|
# Check if file deletion should be blocked due to batch references
|
|
await self._check_file_deletion_allowed(file_id)
|
|
|
|
# file_id = convert_b64_uid_to_unified_uid(file_id)
|
|
model_file_id_mapping = await self.get_model_file_id_mapping(
|
|
[file_id], litellm_parent_otel_span
|
|
)
|
|
|
|
delete_response = None
|
|
specific_model_file_id_mapping = model_file_id_mapping.get(file_id)
|
|
if specific_model_file_id_mapping:
|
|
# Remove conflicting keys from data to avoid duplicate keyword arguments
|
|
filtered_data = {
|
|
k: v for k, v in data.items() if k not in ("model", "file_id")
|
|
}
|
|
for model_id, model_file_id in specific_model_file_id_mapping.items():
|
|
delete_response = await llm_router.afile_delete(model=model_id, file_id=model_file_id, **filtered_data) # type: ignore
|
|
|
|
stored_file_object = await self.delete_unified_file_id(
|
|
file_id, litellm_parent_otel_span
|
|
)
|
|
|
|
# Record successful deletion metric only on actual success
|
|
if stored_file_object or delete_response:
|
|
prom_logger = self._get_prometheus_logger()
|
|
if prom_logger:
|
|
prom_logger.record_managed_file_deleted(result="success")
|
|
|
|
if stored_file_object:
|
|
return stored_file_object
|
|
elif delete_response:
|
|
delete_response.id = file_id
|
|
return delete_response
|
|
else:
|
|
raise Exception(f"LiteLLM Managed File object with id={file_id} not found")
|
|
|
|
async def afile_content(
|
|
self,
|
|
file_id: str,
|
|
litellm_parent_otel_span: Optional[Span],
|
|
llm_router: Router,
|
|
**data: Dict,
|
|
) -> "HttpxBinaryResponseContent":
|
|
"""
|
|
Get the content of a file from first model that has it
|
|
"""
|
|
model_file_id_mapping = data.pop("model_file_id_mapping", None)
|
|
model_file_id_mapping = (
|
|
model_file_id_mapping
|
|
or await self.get_model_file_id_mapping([file_id], litellm_parent_otel_span)
|
|
)
|
|
|
|
specific_model_file_id_mapping = model_file_id_mapping.get(file_id)
|
|
|
|
if specific_model_file_id_mapping:
|
|
exception_dict = {}
|
|
for model_id, provider_file_id in specific_model_file_id_mapping.items():
|
|
try:
|
|
# Cloud-storage providers (e.g. Bedrock S3) validate file ids
|
|
# against the deployment's configured bucket, which they only
|
|
# trust from this immutable server-side snapshot, never from
|
|
# request params.
|
|
credentials = llm_router.get_deployment_credentials_with_provider(
|
|
model_id=model_id
|
|
)
|
|
if credentials is not None:
|
|
data["_litellm_internal_model_credentials"] = cast(
|
|
Dict, MappingProxyType(dict(credentials))
|
|
)
|
|
else:
|
|
data.pop("_litellm_internal_model_credentials", None)
|
|
return await llm_router.afile_content(model=model_id, file_id=provider_file_id, **data) # type: ignore
|
|
except Exception as e:
|
|
exception_dict[model_id] = str(e)
|
|
raise Exception(
|
|
f"LiteLLM Managed File object with id={file_id} not found. Checked model id's: {specific_model_file_id_mapping.keys()}. Errors: {exception_dict}"
|
|
)
|
|
else:
|
|
raise Exception(f"LiteLLM Managed File object with id={file_id} not found")
|
|
|
|
async def _convert_storage_files_to_base64(
|
|
self,
|
|
messages: List[AllMessageValues],
|
|
file_ids: List[str],
|
|
litellm_parent_otel_span: Optional[Span],
|
|
) -> None:
|
|
"""
|
|
Convert files stored in storage backends to base64 format for Vertex AI/Gemini.
|
|
|
|
This method checks if any managed files are stored in storage backends,
|
|
downloads them, and converts them to base64 format in the messages.
|
|
"""
|
|
# Check each file_id to see if it's stored in a storage backend
|
|
for file_id in file_ids:
|
|
# Check if this is a base64 encoded unified file ID
|
|
decoded_unified_file_id = _is_base64_encoded_unified_file_id(file_id)
|
|
|
|
if not decoded_unified_file_id:
|
|
continue
|
|
|
|
# Check database for storage backend info
|
|
# IMPORTANT: The database stores the base64 encoded unified_file_id (not the decoded version)
|
|
# So we query with the original file_id (which is base64 encoded)
|
|
db_file = await self.prisma_client.db.litellm_managedfiletable.find_first(
|
|
where={"unified_file_id": file_id}
|
|
)
|
|
|
|
if not db_file or not db_file.storage_backend or not db_file.storage_url:
|
|
continue
|
|
|
|
# File is stored in a storage backend, download and convert to base64
|
|
try:
|
|
from litellm.llms.base_llm.files.storage_backend_factory import (
|
|
get_storage_backend,
|
|
)
|
|
|
|
storage_backend_name = db_file.storage_backend
|
|
storage_url = db_file.storage_url
|
|
|
|
# Get storage backend (uses same env vars as callback)
|
|
try:
|
|
storage_backend = get_storage_backend(storage_backend_name)
|
|
except ValueError as e:
|
|
verbose_logger.warning(
|
|
f"Storage backend '{storage_backend_name}' error for file {file_id}: {str(e)}"
|
|
)
|
|
continue
|
|
|
|
file_content = await storage_backend.download_file(storage_url)
|
|
|
|
# Determine content type from file object
|
|
content_type = self._get_content_type_from_file_object(
|
|
db_file.file_object
|
|
)
|
|
|
|
# Convert to base64
|
|
base64_data = base64.b64encode(file_content).decode("utf-8")
|
|
base64_data_uri = f"data:{content_type};base64,{base64_data}"
|
|
|
|
# Update messages to use base64 instead of file_id
|
|
self._update_messages_with_base64_data(
|
|
messages, file_id, base64_data_uri, content_type
|
|
)
|
|
except Exception as e:
|
|
verbose_logger.exception(
|
|
f"Error converting file {file_id} from storage backend to base64: {str(e)}"
|
|
)
|
|
# Continue with other files even if one fails
|
|
continue
|
|
|
|
def _get_content_type_from_file_object(self, file_object: Optional[Any]) -> str:
|
|
"""
|
|
Determine content type from file object.
|
|
|
|
Uses the MIME type utility for consistent detection and normalization.
|
|
|
|
Args:
|
|
file_object: The file object from the database (can be dict, JSON string, or None)
|
|
|
|
Returns:
|
|
str: MIME type (defaults to "application/octet-stream" if cannot be determined)
|
|
"""
|
|
# Use utility function for detection
|
|
content_type = get_content_type_from_file_object(file_object)
|
|
|
|
# Normalize for Gemini/Vertex AI (requires image/jpeg, not image/jpg)
|
|
content_type = normalize_mime_type_for_provider(content_type, provider="gemini")
|
|
|
|
return content_type
|
|
|
|
def _update_messages_with_base64_data(
|
|
self,
|
|
messages: List[AllMessageValues],
|
|
file_id: str,
|
|
base64_data_uri: str,
|
|
content_type: str,
|
|
) -> None:
|
|
"""
|
|
Update messages to replace file_id with base64 data URI.
|
|
|
|
Args:
|
|
messages: List of messages to update
|
|
file_id: The file ID to replace
|
|
base64_data_uri: The base64 data URI to use as replacement
|
|
content_type: The MIME type of the file (e.g., "image/jpeg", "application/pdf")
|
|
"""
|
|
for message in messages:
|
|
if message.get("role") == "user":
|
|
content = message.get("content")
|
|
if content and isinstance(content, list):
|
|
for element in content:
|
|
if element.get("type") == "file":
|
|
file_element = cast(ChatCompletionFileObject, element)
|
|
file_element_file = file_element.get("file", {})
|
|
|
|
if file_element_file.get("file_id") == file_id:
|
|
# Replace file_id with base64 data
|
|
file_element_file["file_data"] = base64_data_uri
|
|
# Set format to help Gemini determine mime type
|
|
file_element_file["format"] = content_type
|
|
# Remove file_id to ensure only file_data is used
|
|
file_element_file.pop("file_id", None)
|
|
|
|
verbose_logger.debug(
|
|
f"Converted file {file_id} from storage backend to base64 with format {content_type}"
|
|
)
|