Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/auth/auth_utils.py: 59%
780 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1import importlib.util
2import os
3import re
4import sys
5from collections.abc import Collection, Iterator, Mapping
6from functools import lru_cache
7from logging import Logger
8from typing import Any, Final, Protocol
10from fastapi import HTTPException, Request, status
11from pydantic import PositiveInt, TypeAdapter, ValidationError
13import litellm
14from litellm import Router, constants, provider_list
15from litellm._logging import verbose_proxy_logger
16from litellm.constants import (
17 BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY,
18 EMPTY_MAPPING,
19 INVALID_VIRTUAL_KEY_ERROR_MARKER,
20 MINIMUM_CUSTOM_KEY_LENGTH,
21 STANDARD_CUSTOMER_ID_HEADERS,
22)
23from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
24from litellm.litellm_core_utils.url_utils import (
25 SSRFError,
26 is_url_destination_allowed_by_host,
27 provider_url_destination_candidates,
28 validate_url,
29)
30from litellm.llms.azure.passthrough.transformation import azure_router_model_in_endpoint
31from litellm.llms.nvidia_nim.passthrough.transformation import nvidia_nim_model_group_in_path
32from litellm.proxy._types import *
33from litellm.proxy.common_utils.http_parsing_utils import extract_nested_form_metadata
34from litellm.types.passthrough_endpoints.pass_through_endpoints import (
35 LITELLM_PASS_THROUGH_ENDPOINT_MARKER,
36)
37from litellm.types.router import CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS, Deployment
38from litellm.types.utils import CustomPricingLiteLLMParams
41def is_invalid_virtual_key_error(exception: BaseException | None) -> bool:
42 """True when an authentication error rejects a malformed virtual key.
44 Classifies only by the marker stamped where that 401 is raised. Message
45 content is never inspected: other 401s interpolate caller-supplied values
46 (vector store ids, organization ids) into their messages, so a phrase
47 match would let a request body demote an authorization failure to the
48 quiet log path.
49 """
50 if not isinstance(exception, (HTTPException, ProxyException)):
51 return False
53 code: Final[object] = getattr(exception, "code", None)
54 status_code: Final[object] = code if code is not None else getattr(exception, "status_code", None)
55 if str(status_code) != str(status.HTTP_401_UNAUTHORIZED):
56 return False
58 return getattr(exception, INVALID_VIRTUAL_KEY_ERROR_MARKER, False) is True
61def mark_invalid_virtual_key_error(exception: ProxyException, is_invalid_virtual_key: bool) -> ProxyException:
62 """Return an independently marked malformed-key exception after callback transformations."""
63 if not is_invalid_virtual_key or str(exception.code) != str(status.HTTP_401_UNAUTHORIZED):
64 return exception
65 marked_exception: Final = ProxyException(
66 message=exception.message,
67 type=exception.type,
68 param=exception.param,
69 code=exception.code,
70 headers=exception.headers.copy(),
71 openai_code=None if exception.openai_code is None else str(exception.openai_code),
72 provider_specific_fields=exception.provider_specific_fields,
73 )
74 setattr(marked_exception, INVALID_VIRTUAL_KEY_ERROR_MARKER, True)
75 return marked_exception
78def _get_request_ip_address(request: Request, use_x_forwarded_for: bool | None = False) -> str | None:
79 client_ip = None
80 if use_x_forwarded_for is True and "x-forwarded-for" in request.headers: 80 ↛ 81line 80 didn't jump to line 81 because the condition on line 80 was never true
81 client_ip = request.headers["x-forwarded-for"]
82 elif request.client is not None: 82 ↛ 85line 82 didn't jump to line 85 because the condition on line 82 was always true
83 client_ip = request.client.host
84 else:
85 client_ip = ""
87 return client_ip
90def _check_valid_ip(
91 allowed_ips: list[str] | None,
92 request: Request,
93 use_x_forwarded_for: bool | None = False,
94) -> tuple[bool, str | None]:
95 """
96 Returns if ip is allowed or not
97 """
98 if allowed_ips is None: # if not set, assume true 98 ↛ 102line 98 didn't jump to line 102 because the condition on line 98 was always true
99 return True, None
101 # if general_settings.get("use_x_forwarded_for") is True then use x-forwarded-for
102 client_ip: Final = _get_request_ip_address(request=request, use_x_forwarded_for=use_x_forwarded_for)
104 # Check if IP address is allowed
105 if client_ip not in allowed_ips:
106 return False, client_ip
108 return True, client_ip
111def check_complete_credentials(request_body: dict) -> bool:
112 """
113 if 'api_base' in request body. Check if complete credentials given. Prevent malicious attacks.
115 Supplying an ``api_key`` is necessary but not sufficient: even with
116 credentials supplied, an ``api_base`` / ``base_url`` that resolves to a
117 private/internal/cloud-metadata address would still allow the proxy to
118 be used as an SSRF pivot. Validate any URL fields here so the gate
119 can't be bypassed with ``api_key=anything`` plus a malicious target.
120 """
121 given_model: str | None = None
123 given_model = request_body.get("model")
124 if given_model is None:
125 return False
127 if (
128 "sagemaker" in given_model
129 or "bedrock" in given_model
130 or "vertex_ai" in given_model
131 or "vertex_ai_beta" in given_model
132 ):
133 # complex credentials - easier to make a malicious request
134 return False
136 api_key_value: Final = request_body.get("api_key")
137 if not (api_key_value and isinstance(api_key_value, str) and api_key_value.strip()):
138 return False
140 # ``validate_url`` itself doesn't consult the toggle; ``safe_get`` /
141 # ``async_safe_get`` do. Mirror that here so admins who explicitly
142 # disabled URL validation (e.g. for an internal Ollama endpoint they
143 # accept the SSRF risk for) aren't blocked at the proxy boundary.
144 if getattr(litellm, "user_url_validation", False):
145 for url_field in ("api_base", "base_url"):
146 url_value = request_body.get(url_field)
147 if not url_value or not isinstance(url_value, str):
148 continue
149 try:
150 validate_url(url_value)
151 except SSRFError as e:
152 raise ValueError(
153 f"Rejected request: client-side {url_field}={url_value!r} is rejected by the SSRF guard ({e})."
154 )
156 return True
159def check_regex_or_str_match(request_body_value: Any, regex_str: str) -> bool:
160 """
161 Check if request_body_value matches the regex_str or is equal to param
162 """
163 if re.match(regex_str, request_body_value) or regex_str == request_body_value:
164 return True
165 return False
168def _is_param_allowed(
169 param: str,
170 request_body_value: object,
171 configurable_clientside_auth_params: CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS,
172) -> bool:
173 """
174 Check if param is a str or dict and if request_body_value is in the list of allowed values
175 """
176 if configurable_clientside_auth_params is None:
177 return False
179 for item in configurable_clientside_auth_params:
180 if isinstance(item, str) and param == item:
181 return True
182 elif isinstance(item, dict):
183 if param == "api_base" and check_regex_or_str_match(
184 request_body_value=request_body_value,
185 regex_str=item["api_base"],
186 ): # assume param is a regex
187 return True
189 return False
192def _allow_model_level_clientside_configurable_parameters(
193 model: str, param: str, request_body_value: object, llm_router: Router | None
194) -> bool:
195 """
196 Check if model is allowed to use configurable client-side params
197 - get matching model
198 - check if 'clientside_configurable_parameters' is set for model
199 -
200 """
201 if llm_router is None: 201 ↛ 202line 201 didn't jump to line 202 because the condition on line 201 was never true
202 return False
203 # check if model is set
204 model_info = llm_router.get_model_group_info(model_group=model)
205 if model_info is None: 205 ↛ 210line 205 didn't jump to line 210 because the condition on line 205 was always true
206 # check if wildcard model is set
207 if model.split("/", 1)[0] in provider_list: 207 ↛ 208line 207 didn't jump to line 208 because the condition on line 207 was never true
208 model_info = llm_router.get_model_group_info(model_group=model.split("/", 1)[0])
210 if model_info is None: 210 ↛ 213line 210 didn't jump to line 213 because the condition on line 210 was always true
211 return False
213 if model_info is None or model_info.configurable_clientside_auth_params is None:
214 return False
216 return _is_param_allowed(
217 param=param,
218 request_body_value=request_body_value,
219 configurable_clientside_auth_params=model_info.configurable_clientside_auth_params,
220 )
223# Config dicts whose entries are spread as ``**dict`` into outbound LLM
224# API calls. ``litellm_embedding_config`` is consumed by the Milvus
225# vector store transformer. ``extra_body`` is the OpenAI-SDK passthrough
226# container: provider modules pull provider-auth fields out of it
227# (e.g. Azure's ``extra_body.azure_ad_token``, Bedrock's
228# ``extra_body.aws_web_identity_token``) without re-validating, so the
229# banned-key check has to descend into it the same way it descends into
230# ``litellm_embedding_config``.
231_NESTED_CONFIG_KEYS: Final[tuple[str, ...]] = ("litellm_embedding_config", "extra_body")
233# Metadata containers that carry per-request configuration consumed by the
234# observability callbacks. The same banned-param list applies — a value
235# under ``metadata.langfuse_host`` redirects the same Langfuse client and
236# leaks the same credentials as the root-level ``langfuse_host``, but the
237# original check only walked the request-body root, so the metadata path
238# was an unintentional bypass.
239_NESTED_METADATA_KEYS: Final[tuple[str, ...]] = ("metadata", "litellm_metadata")
241# Banned request-body params. The same list applies to every entry in
242# ``_NESTED_CONFIG_KEYS`` (dicts spread as ``**kwargs`` into outbound
243# calls) and ``_NESTED_METADATA_KEYS`` (dicts read directly by integration
244# callbacks), so a single banned name is enforced wherever the field can
245# reach the call path from.
246# Per-request observability params that are SAFE to accept from clients.
247# These describe the request being logged (prompt version, sampling rate)
248# without choosing the destination or the credentials, so they don't
249# contribute to the data-exfil primitive that the rest of
250# ``_supported_callback_params`` does.
251_SAFE_CLIENT_CALLBACK_PARAMS: Final[frozenset[str]] = frozenset(
252 {
253 "langfuse_prompt_version",
254 "langsmith_sampling_rate",
255 }
256)
258# Observability fields that integrations read from the request body or
259# metadata but that are not (yet) listed in ``_supported_callback_params``.
260# Listed here so the proxy bans them today; the long-term cleanup is to
261# fold these into the canonical allowlist so they share one source of
262# truth with the rest.
263_EXTRA_BANNED_OBSERVABILITY_PARAMS: Final[frozenset[str]] = frozenset(
264 {
265 "posthog_api_url",
266 # ``phoenix_project_name`` / ``phoenix_project_name_override`` are NOT
267 # banned: on the proxy the Phoenix integrations only read them from
268 # ``user_api_key_auth_metadata`` (key/team config), so the bare request
269 # fields are inert and rejecting them just breaks SDK-style callers.
270 # Server-reserved: written exclusively by add_user_api_key_auth_to_request_metadata
271 # from the authenticated key's database record. A caller-supplied value
272 # would survive the server merge and let an authenticated user redirect
273 # their Arize/Phoenix telemetry into arbitrary projects.
274 "user_api_key_auth_metadata",
275 "wandb_api_key",
276 "weave_project_id",
277 }
278)
281def _build_banned_observability_params() -> frozenset[str]:
282 """Derive the observability ban list from the canonical allowlist.
284 ``_supported_callback_params`` and ``_request_blocked_callback_params`` in
285 ``litellm/litellm_core_utils/initialize_dynamic_callback_params.py`` is
286 the single place that enumerates every observability field integrations
287 resolve from kwargs/metadata, plus fields that integration code explicitly
288 blocks from request-supplied callback params. Subtract the small set of
289 informational fields (``_SAFE_CLIENT_CALLBACK_PARAMS``) and union with the
290 extras the canonical allowlist hasn't caught up to yet. New integrations
291 added to the canonical allowlist are banned by default, which is the safe
292 failure mode.
293 """
294 from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
295 _request_blocked_callback_params,
296 _supported_callback_params,
297 )
299 return (
300 (frozenset(_supported_callback_params) - _SAFE_CLIENT_CALLBACK_PARAMS)
301 | frozenset(_request_blocked_callback_params)
302 | _EXTRA_BANNED_OBSERVABILITY_PARAMS
303 )
306_BANNED_REQUEST_BODY_PARAMS: Final[tuple[str, ...]] = (
307 "api_base",
308 "base_url",
309 "user_config",
310 "aws_sts_endpoint",
311 "aws_web_identity_token",
312 "aws_role_name",
313 # Remaining AWS identity selectors. ``get_credentials`` prefers a named
314 # profile over the deployment's static keys, so a caller-supplied
315 # ``aws_profile_name`` signs Bedrock and S3 requests as any profile
316 # present on the proxy host; the two AssumeRole knobs are banned with it
317 # so the whole identity-selection family lives behind the same opt-in.
318 "aws_profile_name",
319 "aws_session_name",
320 "aws_external_id",
321 "aws_session_tags",
322 "vertex_credentials",
323 # Azure managed-identity / federated-auth token. The Azure provider
324 # transformer reads ``azure_ad_token`` (top-level or via
325 # ``extra_body``) and resolves it through ``get_secret`` before
326 # passing it as the bearer token to the Azure endpoint, so a
327 # caller-supplied value is the same exfil shape as
328 # ``aws_web_identity_token`` on the Bedrock path.
329 "azure_ad_token",
330 # Endpoint-targeting fields that retarget the outbound request or
331 # an observability callback. An attacker-controlled value either
332 # exfiltrates the request payload (incl. messages + admin-set
333 # tokens) to the attacker's host, or coerces the proxy into
334 # authenticating against the attacker's host with admin secrets.
335 "aws_bedrock_runtime_endpoint",
336 # Bedrock project/workspace association. Deployments pin this to
337 # enforce a data-retention policy, so a caller-supplied value would
338 # re-route the request's retention and accounting to any project
339 # reachable with the deployment's shared AWS credentials.
340 "aws_bedrock_project_id",
341 "workspace_id",
342 "aws_workspace_id",
343 "anthropic_workspace_id",
344 "anthropic-workspace-id",
345 "bedrock_tags",
346 # Provider-specific endpoint overrides that flow into the outbound
347 # request via ``optional_params``. Same threat as ``api_base``:
348 # ``s3_endpoint_url`` redirects Bedrock file uploads to attacker
349 # S3; ``sagemaker_base_url`` redirects all SageMaker traffic;
350 # ``deployment_url`` redirects SAP deployments.
351 "s3_endpoint_url",
352 "sagemaker_base_url",
353 "deployment_url",
354 # NVIDIA Riva fields consumed by the audio-transcription handler
355 # via ``optional_params``. Banned for the same reason as the
356 # provider-specific entries above: a caller-supplied value retargets
357 # the request away from the admin's pinned configuration.
358 "nvcf_function_id",
359 "use_ssl",
360 # Per-deployment opt-in that hands the whole call to the Rust core. It is a
361 # deployment decision, not a request one: the Rust path uses its own client
362 # rather than the one the deployment configured, and reports no post_call,
363 # so a caller-supplied value picks a transport and a callback surface the
364 # admin did not choose.
365 "rust",
366 # SDK-only field; also rejected outright in is_request_body_safe.
367 "model_list",
368 "vertex_ai_credentials",
369 # Observability credentials, hosts, and project identifiers: derived
370 # from the canonical ``_supported_callback_params`` allowlist so new
371 # integrations are covered automatically. Sorted for stable iteration
372 # order and reviewable diffs.
373 *sorted(_build_banned_observability_params()),
374 *sorted(CustomPricingLiteLLMParams.model_fields.keys()),
375)
378def _check_banned_params(
379 body: dict,
380 general_settings: dict,
381 llm_router: Router | None,
382 model: str,
383) -> None:
384 """Raise ``ValueError`` if ``body`` carries a banned param without admin opt-in.
386 Shared between the root-level check and the nested-config check so a
387 new banned param only needs to be added in one place.
388 """
389 for param in _BANNED_REQUEST_BODY_PARAMS:
390 if param not in body:
391 continue
392 if general_settings.get("allow_client_side_credentials") is True: 392 ↛ 395line 392 didn't jump to line 395 because the condition on line 392 was never true
393 # Proxy-wide opt-in: every banned param is permitted, exit
394 # entirely so the rest of the loop doesn't waste work.
395 return
396 if ( 396 ↛ 410line 396 didn't jump to line 410 because the condition on line 396 was never true
397 _allow_model_level_clientside_configurable_parameters(
398 model=model,
399 param=param,
400 request_body_value=body[param],
401 llm_router=llm_router,
402 )
403 is True
404 ):
405 # Per-param opt-in: only THIS param is permitted by the
406 # deployment's ``configurable_clientside_auth_params``. Skip
407 # to the next banned param so a body that pairs an allowed
408 # ``api_base`` with an unallowed ``langfuse_host`` is still
409 # rejected for the second field.
410 continue
411 raise ValueError(
412 f"Rejected Request: {param} is not allowed in request body. "
413 "Clientside passthrough requires explicit admin opt-in via "
414 "either `general_settings.allow_client_side_credentials = true` "
415 "(proxy-wide) or `configurable_clientside_auth_params` on the "
416 "deployment in your proxy config.yaml. "
417 "Relevant Issue: https://huntr.com/bounties/4001e1a2-7b7a-4776-a3ae-e6692ec3d997",
418 )
421_FALLBACK_FIELDS: Final[tuple[str, ...]] = (
422 "fallbacks",
423 "context_window_fallbacks",
424 "content_policy_fallbacks",
425)
428def _iter_fallback_field_values(request_body: Mapping[str, object]) -> Iterator[object]:
429 override: Final = request_body.get("router_settings_override")
430 for source in (request_body, override):
431 if isinstance(source, Mapping):
432 for field in _FALLBACK_FIELDS:
433 yield source.get(field)
436def _iter_fallback_targets(value: object, depth: int) -> Iterator[str | Mapping[str, object]]:
437 if depth > 2 * litellm.ROUTER_MAX_FALLBACKS: 437 ↛ 438line 437 didn't jump to line 438 because the condition on line 437 was never true
438 raise ValueError("Rejected Request: fallback nesting exceeds the allowed validation depth.")
439 if not isinstance(value, list):
440 return
441 for item in value:
442 if isinstance(item, str):
443 yield item
444 elif isinstance(item, Mapping):
445 values = tuple(item.values())
446 if not (values and all(isinstance(v, list) for v in values)): 446 ↛ 448line 446 didn't jump to line 448 because the condition on line 446 was always true
447 yield item
448 if isinstance(item.get("model"), str): 448 ↛ 449line 448 didn't jump to line 449 because the condition on line 448 was never true
449 for field in _FALLBACK_FIELDS:
450 yield from _iter_fallback_targets(item.get(field), depth + 1)
451 else:
452 for target_list in values: 452 ↛ 453line 452 didn't jump to line 453 because the loop on line 452 never started
453 yield from _iter_fallback_targets(target_list, depth + 1)
456def iter_request_fallback_targets(request_body: Mapping[str, object]) -> Iterator[str | Mapping[str, object]]:
457 for value in _iter_fallback_field_values(request_body):
458 yield from _iter_fallback_targets(value, 0)
461def _reject_url_valued_fallback_target(value: str) -> None:
462 allowed_hosts: Final = getattr(litellm, "provider_url_destination_allowed_hosts", []) or []
463 for candidate in provider_url_destination_candidates(value):
464 if not candidate.lower().startswith(("http://", "https://")): 464 ↛ 466line 464 didn't jump to line 466 because the condition on line 464 was always true
465 continue
466 if is_url_destination_allowed_by_host(candidate, allowed_hosts):
467 continue
468 raise ValueError(
469 f"Rejected Request: URL-valued fallback destination '{value}' is not allowed. "
470 "Configure custom endpoints with api_base instead, or add the destination host to "
471 "`provider_url_destination_allowed_hosts` in litellm_settings."
472 )
475def is_request_body_safe(request_body: dict, general_settings: dict, llm_router: Router | None, model: str) -> bool:
476 """
477 Check if the request body is safe.
479 A malicious user can set the api_base to their own domain and invoke POST /chat/completions to intercept and steal the OpenAI API key.
480 Relevant issue: https://huntr.com/bounties/4001e1a2-7b7a-4776-a3ae-e6692ec3d997
482 The blocklist is enforced unconditionally. Legitimate clientside
483 credential / endpoint passthrough goes through one of the two
484 explicit admin opt-ins (``general_settings.allow_client_side_credentials``
485 proxy-wide or ``configurable_clientside_auth_params`` per deployment).
486 Historically there was a third, *implicit*, *caller-controlled* path:
487 ``check_complete_credentials`` returned True when the caller supplied
488 any non-empty ``api_key``, which made the entire blocklist a no-op.
489 That bypass turned every missing entry on the blocklist into an
490 exploitable SSRF / credential-exfil hole — see GHSA-jh89-88fc-qrfp,
491 GHSA-3frq-6r6h-7j64, and the chain of veria-admin findings (Dv_m860l,
492 b_yRJeQ5, stN90yjP, LBlyOAc8, U2TD78kg). Removed: the blocklist now
493 has a single, predictable failure mode for missing entries (a 400),
494 not a credential leak.
496 Iterative single-level descent into ``_NESTED_CONFIG_KEYS`` (rather
497 than recursion) covers nested-config attacks like Milvus's
498 ``litellm_embedding_config.api_base`` (VERIA-6) without exposing a
499 recursion-depth DoS surface.
500 """
501 if "model_list" in request_body: 501 ↛ 502line 501 didn't jump to line 502 because the condition on line 501 was never true
502 raise ValueError("Rejected Request: model_list is not allowed in the request body.")
503 _check_banned_params(request_body, general_settings, llm_router, model)
504 for nested_key in _NESTED_CONFIG_KEYS:
505 nested = _coerce_metadata_to_dict(request_body.get(nested_key))
506 if nested is not None: 506 ↛ 507line 506 didn't jump to line 507 because the condition on line 506 was never true
507 _check_banned_params(nested, general_settings, llm_router, model)
508 for metadata_key in _NESTED_METADATA_KEYS:
509 metadata = _coerce_metadata_to_dict(request_body.get(metadata_key))
510 if metadata is not None:
511 _check_banned_params(metadata, general_settings, llm_router, model)
512 if any(isinstance(key, str) and key.startswith(f"{metadata_key}[") for key in request_body): 512 ↛ 513line 512 didn't jump to line 513 because the condition on line 512 was never true
513 _check_banned_params(
514 extract_nested_form_metadata(form_data=request_body, prefix=f"{metadata_key}["),
515 general_settings,
516 llm_router,
517 model,
518 )
519 for target in iter_request_fallback_targets(request_body):
520 if isinstance(target, dict):
521 _check_banned_params(target, general_settings, llm_router, model)
522 target_model = target.get("model")
523 if isinstance(target_model, str): 523 ↛ 524line 523 didn't jump to line 524 because the condition on line 523 was never true
524 _reject_url_valued_fallback_target(target_model)
525 elif isinstance(target, str): 525 ↛ 519line 525 didn't jump to line 519 because the condition on line 525 was always true
526 _reject_url_valued_fallback_target(target)
527 litellm_params: Final = _coerce_metadata_to_dict(request_body.get("litellm_params"))
528 if litellm_params is not None:
529 litellm_params_metadata: Final = _coerce_metadata_to_dict(litellm_params.get("metadata"))
530 if litellm_params_metadata is not None: 530 ↛ 531line 530 didn't jump to line 531 because the condition on line 530 was never true
531 _check_banned_params(
532 litellm_params_metadata,
533 general_settings,
534 llm_router,
535 model,
536 )
537 return True
540def _coerce_metadata_to_dict(value: object) -> dict[str, object] | None:
541 """Return ``value`` as a dict, parsing it from JSON if delivered as a string.
543 Multipart/form-data and ``extra_body`` callers send ``litellm_metadata``
544 as a JSON-encoded string; the proxy parses it into a dict later in
545 ``add_litellm_data_to_request``, but the auth-time bouncer runs first
546 and would otherwise miss the banned-param check on a still-stringified
547 metadata blob.
548 """
549 if isinstance(value, dict):
550 return value
551 if isinstance(value, str):
552 from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
554 parsed: Final = safe_json_loads(value)
555 if isinstance(parsed, dict):
556 return parsed
557 return None
560async def pre_db_read_auth_checks(
561 request: Request,
562 request_data: dict,
563 route: str,
564):
565 """
566 1. Checks if request size is under max_request_size_mb (if set)
567 2. Check if request body is safe (example user has not set api_base in request body)
568 3. Check if IP address is allowed (if set)
569 4. Check if request route is an allowed route on the proxy (if set)
571 Returns:
572 - True
574 Raises:
575 - HTTPException if request fails initial auth checks
576 """
577 from litellm.proxy.proxy_server import general_settings, llm_router, premium_user
579 # Check 1. request size
580 await check_if_request_size_is_safe(request=request)
582 # Check 2. Request body is safe
583 is_request_body_safe(
584 request_body=request_data,
585 general_settings=general_settings,
586 llm_router=llm_router,
587 model=request_data.get("model", ""), # [TODO] use model passed in url as well (azure openai routes)
588 )
590 # Check 3. Check if IP address is allowed
591 is_valid_ip, passed_in_ip = _check_valid_ip(
592 allowed_ips=general_settings.get("allowed_ips", None),
593 use_x_forwarded_for=general_settings.get("use_x_forwarded_for", False),
594 request=request,
595 )
597 if not is_valid_ip: 597 ↛ 598line 597 didn't jump to line 598 because the condition on line 597 was never true
598 raise HTTPException(
599 status_code=status.HTTP_403_FORBIDDEN,
600 detail=f"Access forbidden: IP address {passed_in_ip} not allowed.",
601 )
603 # Check 4. Check if request route is an allowed route on the proxy
604 if "allowed_routes" in general_settings: 604 ↛ 605line 604 didn't jump to line 605 because the condition on line 604 was never true
605 _allowed_routes: Final = general_settings["allowed_routes"]
606 if premium_user is not True:
607 verbose_proxy_logger.error(
608 "Trying to set allowed_routes. This is an Enterprise feature. %s",
609 CommonProxyErrors.not_premium_user.value,
610 )
611 if route not in _allowed_routes:
612 verbose_proxy_logger.error("Route %s not in allowed_routes=%s", route, _allowed_routes)
613 raise HTTPException(
614 status_code=status.HTTP_403_FORBIDDEN,
615 detail=f"Access forbidden: Route {route} not allowed",
616 )
619def route_in_additonal_public_routes(current_route: str):
620 """
621 Helper to check if the user defined public_routes on config.yaml
623 Parameters:
624 - current_route: str - the route the user is trying to call
626 Returns:
627 - bool - True if the route is defined in public_routes
628 - bool - False if the route is not defined in public_routes
630 Supports wildcard patterns (e.g., "/api/*" matches "/api/users", "/api/users/123")
632 In order to use this the litellm config.yaml should have the following in general_settings:
634 ```yaml
635 general_settings:
636 master_key: sk-1234
637 public_routes: ["LiteLLMRoutes.public_routes", "/spend/calculate", "/api/*"]
638 ```
639 """
640 from litellm.proxy.auth.route_checks import RouteChecks
641 from litellm.proxy.proxy_server import general_settings, premium_user
643 try:
644 if premium_user is not True: 644 ↛ 646line 644 didn't jump to line 646 because the condition on line 644 was always true
645 return False
646 if general_settings is None:
647 return False
649 routes_defined: Final = general_settings.get("public_routes", [])
651 # Check exact match first
652 if current_route in routes_defined:
653 return True
655 # Check wildcard patterns
656 for route_pattern in routes_defined:
657 if RouteChecks.route_matches_wildcard_pattern(route=current_route, pattern=route_pattern):
658 return True
660 return False
661 except Exception as e:
662 verbose_proxy_logger.error("route_in_additonal_public_routes: %s", e)
663 return False
666def get_request_route(request: Request) -> str:
667 """
668 Resolve the request route from the ASGI scope, with ``root_path`` stripped.
670 Prefer this over ``request.url.path`` for any auth, ACL, routing, or
671 audit-log decision: Starlette reconstructs ``url.path`` by interpolating
672 the Host header into a URL string and re-parsing with ``urlsplit``, so a
673 malformed Host (e.g. ``localhost/?x=1``) collapses ``url.path`` to ``"/"``
674 while FastAPI continues to dispatch on ``scope["path"]``. ``scope["path"]``
675 is uvicorn's parse of the HTTP request line and matches the actual
676 handler, so it's the authoritative route.
678 Also normalizes sub-path deployments by stripping ``scope["root_path"]``
679 e.g. ``/genai/chat/completions`` -> ``/chat/completions``.
680 """
681 try:
682 scope: Final = request.scope
683 if not isinstance(scope, dict): 683 ↛ 684line 683 didn't jump to line 684 because the condition on line 683 was never true
684 return str(request.url.path)
685 raw_path: Final[str] = str(scope.get("path", request.url.path))
686 root_path: Final[str] = str(scope.get("app_root_path", scope.get("root_path", ""))).rstrip("/")
687 if not isinstance(raw_path, str): 687 ↛ 688line 687 didn't jump to line 688 because the condition on line 687 was never true
688 return str(request.url.path)
689 # Strip root_path only when it matches whole path segments — guarding
690 # against sibling paths like "/apifoo" being truncated under
691 # root_path="/api". Trailing slashes on root_path are stripped above,
692 # so bare "/" or "/prefix/" still leave the leading "/" intact.
693 if root_path and (raw_path == root_path or raw_path.startswith(root_path + "/")): 693 ↛ 694line 693 didn't jump to line 694 because the condition on line 693 was never true
694 stripped: Final = raw_path[len(root_path) :]
695 return stripped or "/"
696 return raw_path
697 except Exception as e:
698 verbose_proxy_logger.debug(
699 "error on get_request_route: %s, defaulting to request.url.path=%s", e, request.url.path
700 )
701 return str(request.url.path)
704def get_request_route_template(request: Request) -> str | None:
705 """
706 Return the low-cardinality route template, e.g.
707 ``/v1/threads/{thread_id}/runs`` (vs. the literal path from
708 ``get_request_route``). FastAPI sets ``scope["route"]`` before endpoint
709 dependencies run. Returns None if unavailable (unmatched path, Mount).
710 """
711 try:
712 scope: Final = request.scope
713 if not isinstance(scope, dict): 713 ↛ 714line 713 didn't jump to line 714 because the condition on line 713 was never true
714 return None
715 route: Final = scope.get("route")
716 template: Final = getattr(route, "path", None)
717 return template if isinstance(template, str) and template else None
718 except Exception as e:
719 verbose_proxy_logger.debug("error on get_request_route_template: %s", e)
720 return None
723@lru_cache(maxsize=256)
724def normalize_request_route(route: str) -> str:
725 """
726 Normalize request routes by replacing dynamic path parameters with placeholders.
728 This prevents high cardinality in Prometheus metrics by collapsing routes like:
729 - /v1/responses/1234567890 -> /v1/responses/{response_id}
730 - /v1/threads/thread_123 -> /v1/threads/{thread_id}
732 Args:
733 route: The request route path
735 Returns:
736 Normalized route with dynamic parameters replaced by placeholders
738 Examples:
739 >>> normalize_request_route("/v1/responses/abc123")
740 '/v1/responses/{response_id}'
741 >>> normalize_request_route("/v1/responses/abc123/cancel")
742 '/v1/responses/{response_id}/cancel'
743 >>> normalize_request_route("/chat/completions")
744 '/chat/completions'
745 """
746 # Define patterns for routes with dynamic IDs
747 # Format: (regex_pattern, replacement_template)
748 patterns: Final = [
749 # Responses API - must come before generic patterns
750 (r"^(/(?:openai/)?v1/responses)/([^/]+)(/input_items)$", r"\1/{response_id}\3"),
751 (r"^(/(?:openai/)?v1/responses)/([^/]+)(/cancel)$", r"\1/{response_id}\3"),
752 (r"^(/(?:openai/)?v1/responses)/([^/]+)$", r"\1/{response_id}"),
753 (r"^(/responses)/([^/]+)(/input_items)$", r"\1/{response_id}\3"),
754 (r"^(/responses)/([^/]+)(/cancel)$", r"\1/{response_id}\3"),
755 (r"^(/responses)/([^/]+)$", r"\1/{response_id}"),
756 # Threads API
757 (
758 r"^(/(?:openai/)?v1/threads)/([^/]+)(/runs)/([^/]+)(/steps)/([^/]+)$",
759 r"\1/{thread_id}\3/{run_id}\5/{step_id}",
760 ),
761 (
762 r"^(/(?:openai/)?v1/threads)/([^/]+)(/runs)/([^/]+)(/steps)$",
763 r"\1/{thread_id}\3/{run_id}\5",
764 ),
765 (
766 r"^(/(?:openai/)?v1/threads)/([^/]+)(/runs)/([^/]+)(/cancel)$",
767 r"\1/{thread_id}\3/{run_id}\5",
768 ),
769 (
770 r"^(/(?:openai/)?v1/threads)/([^/]+)(/runs)/([^/]+)(/submit_tool_outputs)$",
771 r"\1/{thread_id}\3/{run_id}\5",
772 ),
773 (
774 r"^(/(?:openai/)?v1/threads)/([^/]+)(/runs)/([^/]+)$",
775 r"\1/{thread_id}\3/{run_id}",
776 ),
777 (r"^(/(?:openai/)?v1/threads)/([^/]+)(/runs)$", r"\1/{thread_id}\3"),
778 (
779 r"^(/(?:openai/)?v1/threads)/([^/]+)(/messages)/([^/]+)$",
780 r"\1/{thread_id}\3/{message_id}",
781 ),
782 (r"^(/(?:openai/)?v1/threads)/([^/]+)(/messages)$", r"\1/{thread_id}\3"),
783 (r"^(/(?:openai/)?v1/threads)/([^/]+)$", r"\1/{thread_id}"),
784 # Vector Stores API
785 (
786 r"^(/(?:openai/)?v1/vector_stores)/([^/]+)(/files)/([^/]+)$",
787 r"\1/{vector_store_id}\3/{file_id}",
788 ),
789 (
790 r"^(/(?:openai/)?v1/vector_stores)/([^/]+)(/files)$",
791 r"\1/{vector_store_id}\3",
792 ),
793 (
794 r"^(/(?:openai/)?v1/vector_stores)/([^/]+)(/file_batches)/([^/]+)$",
795 r"\1/{vector_store_id}\3/{batch_id}",
796 ),
797 (
798 r"^(/(?:openai/)?v1/vector_stores)/([^/]+)(/file_batches)$",
799 r"\1/{vector_store_id}\3",
800 ),
801 (r"^(/(?:openai/)?v1/vector_stores)/([^/]+)$", r"\1/{vector_store_id}"),
802 # Assistants API
803 (r"^(/(?:openai/)?v1/assistants)/([^/]+)$", r"\1/{assistant_id}"),
804 # Files API
805 (r"^(/(?:openai/)?v1/files)/([^/]+)(/content)$", r"\1/{file_id}\3"),
806 (r"^(/(?:openai/)?v1/files)/([^/]+)$", r"\1/{file_id}"),
807 # Batches API
808 (r"^(/(?:openai/)?v1/batches)/([^/]+)(/cancel)$", r"\1/{batch_id}\3"),
809 (r"^(/(?:openai/)?v1/batches)/([^/]+)$", r"\1/{batch_id}"),
810 # Fine-tuning API
811 (
812 r"^(/(?:openai/)?v1/fine_tuning/jobs)/([^/]+)(/events)$",
813 r"\1/{fine_tuning_job_id}\3",
814 ),
815 (
816 r"^(/(?:openai/)?v1/fine_tuning/jobs)/([^/]+)(/cancel)$",
817 r"\1/{fine_tuning_job_id}\3",
818 ),
819 (
820 r"^(/(?:openai/)?v1/fine_tuning/jobs)/([^/]+)(/checkpoints)$",
821 r"\1/{fine_tuning_job_id}\3",
822 ),
823 (r"^(/(?:openai/)?v1/fine_tuning/jobs)/([^/]+)$", r"\1/{fine_tuning_job_id}"),
824 # Models API
825 (r"^(/(?:openai/)?v1/models)/([^/]+)$", r"\1/{model}"),
826 ]
828 # Apply patterns in order
829 for pattern, replacement in patterns:
830 normalized = re.sub(pattern, replacement, route)
831 if normalized != route:
832 return normalized
834 # Return original route if no pattern matched
835 return route
838async def check_if_request_size_is_safe(request: Request) -> bool:
839 """
840 Enterprise Only:
841 - Checks if the request size is within the limit
843 Args:
844 request (Request): The incoming request.
846 Returns:
847 bool: True if the request size is within the limit
849 Raises:
850 ProxyException: If the request size is too large
852 """
853 from litellm.proxy.proxy_server import general_settings, premium_user
855 max_request_size_mb: Final = general_settings.get("max_request_size_mb", None)
857 if max_request_size_mb is not None: 857 ↛ 859line 857 didn't jump to line 859 because the condition on line 857 was never true
858 # Check if premium user
859 if premium_user is not True:
860 verbose_proxy_logger.warning(
861 "using max_request_size_mb - not checking - this is an enterprise only feature. %s",
862 CommonProxyErrors.not_premium_user.value,
863 )
864 return True
866 # Get the request body
867 content_length: Final = request.headers.get("content-length")
869 if content_length:
870 header_size: Final = int(content_length)
871 header_size_mb: Final = bytes_to_mb(bytes_value=header_size)
872 verbose_proxy_logger.debug("content_length request size in MB=%s", header_size_mb)
874 if header_size_mb > max_request_size_mb:
875 raise ProxyException(
876 message=f"Request size is too large. Request size is {header_size_mb} MB. Max size is {max_request_size_mb} MB",
877 type=ProxyErrorTypes.bad_request_error.value,
878 code=400,
879 param="content-length",
880 )
881 else:
882 # If Content-Length is not available, read the body
883 body: Final = await request.body()
884 body_size: Final = len(body)
885 request_size_mb: Final = bytes_to_mb(bytes_value=body_size)
887 verbose_proxy_logger.debug("request body request size in MB=%s", request_size_mb)
888 if request_size_mb > max_request_size_mb:
889 raise ProxyException(
890 message=f"Request size is too large. Request size is {request_size_mb} MB. Max size is {max_request_size_mb} MB",
891 type=ProxyErrorTypes.bad_request_error.value,
892 code=400,
893 param="content-length",
894 )
896 return True
899async def check_response_size_is_safe(response: object) -> bool:
900 """
901 Enterprise Only:
902 - Checks if the response size is within the limit
904 Args:
905 response (Any): The response to check.
907 Returns:
908 bool: True if the response size is within the limit
910 Raises:
911 ProxyException: If the response size is too large
913 """
915 from litellm.proxy.proxy_server import general_settings, premium_user
917 max_response_size_mb: Final = general_settings.get("max_response_size_mb", None)
918 if max_response_size_mb is not None: 918 ↛ 920line 918 didn't jump to line 920 because the condition on line 918 was never true
919 # Check if premium user
920 if premium_user is not True:
921 verbose_proxy_logger.warning(
922 "using max_response_size_mb - not checking - this is an enterprise only feature. %s",
923 CommonProxyErrors.not_premium_user.value,
924 )
925 return True
927 response_size_mb: Final = bytes_to_mb(bytes_value=sys.getsizeof(response))
928 verbose_proxy_logger.debug("response size in MB=%s", response_size_mb)
929 if response_size_mb > max_response_size_mb:
930 raise ProxyException(
931 message=f"Response size is too large. Response size is {response_size_mb} MB. Max size is {max_response_size_mb} MB",
932 type=ProxyErrorTypes.bad_request_error.value,
933 code=400,
934 param="content-length",
935 )
937 return True
940def bytes_to_mb(bytes_value: int):
941 """
942 Helper to convert bytes to MB
943 """
944 return bytes_value / (1024 * 1024)
947# helpers used by parallel request limiter to handle model rpm/tpm limits for a given api key
948def _get_deployment_default_limit(model_name: str, field: str) -> int | None:
949 """
950 Return the minimum value of `field` across all deployments for model_name,
951 or None if no deployment has the field set.
953 When multiple deployments share the same model name, taking the minimum is
954 the safest choice for load-balanced setups: it ensures no deployment is
955 over-consumed regardless of which one actually serves a given request.
956 """
957 from litellm.proxy.proxy_server import llm_router
959 if llm_router is None: 959 ↛ 960line 959 didn't jump to line 960 because the condition on line 959 was never true
960 return None
961 deployments: Final = llm_router.get_model_list(model_name=model_name)
962 if not deployments: 962 ↛ 964line 962 didn't jump to line 964 because the condition on line 962 was always true
963 return None
964 limits: Final = []
965 for deployment in deployments:
966 raw = deployment.get("litellm_params", {}).get(field)
967 if raw is not None:
968 try:
969 if isinstance(raw, (int, float, str, bytes, bytearray)):
970 limits.append(int(raw))
971 except (ValueError, TypeError):
972 pass
973 return min(limits) if limits else None
976def _get_deployment_default_rpm_limit(model_name: str) -> int | None:
977 return _get_deployment_default_limit(model_name, "default_api_key_rpm_limit")
980def _get_deployment_default_tpm_limit(model_name: str) -> int | None:
981 return _get_deployment_default_limit(model_name, "default_api_key_tpm_limit")
984def get_key_own_model_rate_limit(
985 user_api_key_dict: UserAPIKeyAuth,
986 rate_limit_key: Literal["model_rpm_limit", "model_tpm_limit"],
987) -> dict[str, int] | None:
988 if user_api_key_dict.metadata: 988 ↛ 989line 988 didn't jump to line 989 because the condition on line 988 was never true
989 result: Final = user_api_key_dict.metadata.get(rate_limit_key)
990 if result:
991 return result
993 if not user_api_key_dict.model_max_budget: 993 ↛ 995line 993 didn't jump to line 995 because the condition on line 993 was always true
994 return None
995 budget_key: Final = "rpm_limit" if rate_limit_key == "model_rpm_limit" else "tpm_limit"
996 model_limit: Final = {
997 model: budget[budget_key]
998 for model, budget in user_api_key_dict.model_max_budget.items()
999 if isinstance(budget, dict) and budget.get(budget_key) is not None
1000 }
1001 return model_limit or None
1004def get_key_model_rpm_limit(
1005 user_api_key_dict: UserAPIKeyAuth,
1006 model_name: str | None = None,
1007) -> dict[str, int] | None:
1008 """
1009 Get the model rpm limit for a given api key.
1011 Priority order (returns first found):
1012 1. Key metadata (model_rpm_limit)
1013 2. Key model_max_budget (rpm_limit per model)
1014 3. Team metadata (model_rpm_limit)
1015 4. Deployment default_api_key_rpm_limit (when model_name is provided)
1016 """
1017 key_own_limit: Final = get_key_own_model_rate_limit(user_api_key_dict, "model_rpm_limit")
1018 if key_own_limit is not None: 1018 ↛ 1019line 1018 didn't jump to line 1019 because the condition on line 1018 was never true
1019 return key_own_limit
1021 # 3. Fallback to team metadata
1022 if user_api_key_dict.team_metadata: 1022 ↛ 1023line 1022 didn't jump to line 1023 because the condition on line 1022 was never true
1023 team_limit: Final = user_api_key_dict.team_metadata.get("model_rpm_limit")
1024 if team_limit is not None:
1025 return team_limit
1027 # 4. Fallback to deployment default_api_key_rpm_limit
1028 if model_name is not None: 1028 ↛ 1033line 1028 didn't jump to line 1033 because the condition on line 1028 was always true
1029 default_limit: Final = _get_deployment_default_rpm_limit(model_name)
1030 if default_limit is not None: 1030 ↛ 1031line 1030 didn't jump to line 1031 because the condition on line 1030 was never true
1031 return {model_name: default_limit}
1033 return None
1036def get_key_model_tpm_limit(
1037 user_api_key_dict: UserAPIKeyAuth,
1038 model_name: str | None = None,
1039) -> dict[str, int] | None:
1040 """
1041 Get the model tpm limit for a given api key.
1043 Priority order (returns first found):
1044 1. Key metadata (model_tpm_limit)
1045 2. Key model_max_budget (tpm_limit per model)
1046 3. Team metadata (model_tpm_limit)
1047 4. Deployment default_api_key_tpm_limit (when model_name is provided)
1048 """
1049 key_own_limit: Final = get_key_own_model_rate_limit(user_api_key_dict, "model_tpm_limit")
1050 if key_own_limit is not None: 1050 ↛ 1051line 1050 didn't jump to line 1051 because the condition on line 1050 was never true
1051 return key_own_limit
1053 # 3. Fallback to team metadata
1054 if user_api_key_dict.team_metadata: 1054 ↛ 1055line 1054 didn't jump to line 1055 because the condition on line 1054 was never true
1055 team_limit: Final = user_api_key_dict.team_metadata.get("model_tpm_limit")
1056 if team_limit is not None:
1057 return team_limit
1059 # 4. Fallback to deployment default_api_key_tpm_limit
1060 if model_name is not None: 1060 ↛ 1065line 1060 didn't jump to line 1065 because the condition on line 1060 was always true
1061 default_limit: Final = _get_deployment_default_tpm_limit(model_name)
1062 if default_limit is not None: 1062 ↛ 1063line 1062 didn't jump to line 1063 because the condition on line 1062 was never true
1063 return {model_name: default_limit}
1065 return None
1068ESTIMATED_OUTPUT_TOKENS_FIELD: Final = "default_estimated_output_tokens"
1069ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD: Final = "default_estimated_output_tokens_per_model"
1070ESTIMATED_OUTPUT_TOKENS_METADATA_FIELDS: Final = frozenset(
1071 {ESTIMATED_OUTPUT_TOKENS_FIELD, ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD}
1072)
1074_ESTIMATED_OUTPUT_TOKENS_ADAPTER: Final = TypeAdapter(PositiveInt)
1075_ESTIMATED_OUTPUT_TOKENS_PER_MODEL_ADAPTER: Final = TypeAdapter(Mapping[str, PositiveInt])
1078def _validated_output_token_estimate(raw: object) -> int | None:
1079 """Coerce one declared estimate to a positive int, or ignore it."""
1080 if raw is None:
1081 return None
1082 try:
1083 return _ESTIMATED_OUTPUT_TOKENS_ADAPTER.validate_python(raw)
1084 except ValidationError as validation_error:
1085 verbose_proxy_logger.warning(
1086 "Ignoring malformed %s in metadata: %s",
1087 ESTIMATED_OUTPUT_TOKENS_FIELD,
1088 validation_error,
1089 )
1090 return None
1093def _validated_output_token_estimates_per_model(raw: object) -> Mapping[str, int] | None:
1094 """Coerce a declared per-model estimate map, or ignore it."""
1095 if raw is None:
1096 return None
1097 try:
1098 return _ESTIMATED_OUTPUT_TOKENS_PER_MODEL_ADAPTER.validate_python(raw)
1099 except ValidationError as validation_error:
1100 verbose_proxy_logger.warning(
1101 "Ignoring malformed %s in metadata: %s",
1102 ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD,
1103 validation_error,
1104 )
1105 return None
1108def _estimated_output_tokens_from_metadata(
1109 metadata: Mapping[str, object] | None,
1110 model_name: str | None,
1111) -> int | None:
1112 """Resolve the per-model, then global, estimate out of one metadata blob.
1114 The two fields are validated independently so a malformed per-model map
1115 cannot discard a valid global estimate, or the other way round.
1116 """
1117 if not metadata or ESTIMATED_OUTPUT_TOKENS_METADATA_FIELDS.isdisjoint(metadata):
1118 return None
1120 if model_name is not None:
1121 per_model: Final = _validated_output_token_estimates_per_model(
1122 metadata.get(ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD)
1123 )
1124 per_model_estimate: Final = per_model.get(model_name) if per_model is not None else None
1125 if per_model_estimate is not None:
1126 return per_model_estimate
1128 return _validated_output_token_estimate(metadata.get(ESTIMATED_OUTPUT_TOKENS_FIELD))
1131def get_estimated_output_tokens(
1132 user_api_key_dict: UserAPIKeyAuth,
1133 model_name: str | None = None,
1134) -> int | None:
1135 """Resolve the operator-declared output-token estimate for TPM reservation.
1137 Priority order (returns first found):
1138 1. Key metadata ``default_estimated_output_tokens_per_model[model_name]``
1139 2. Key metadata ``default_estimated_output_tokens``
1140 3. Team metadata ``default_estimated_output_tokens_per_model[model_name]``
1141 4. Team metadata ``default_estimated_output_tokens``
1143 Returns ``None`` when nothing is configured, which leaves the static
1144 heuristic floor in place.
1145 """
1146 key_estimate: Final = _estimated_output_tokens_from_metadata(user_api_key_dict.metadata, model_name)
1147 if key_estimate is not None:
1148 return key_estimate
1149 return _estimated_output_tokens_from_metadata(user_api_key_dict.team_metadata, model_name)
1152class OutputTokenEstimateRequest(Protocol):
1153 """The shape of any management request that can carry an output-token estimate.
1155 Read-only members: the gate inspects a request, it never writes one back.
1156 """
1158 @property
1159 def metadata(self) -> Mapping[str, object] | None: ... 1159 ↛ exitline 1159 didn't return from function 'metadata' because
1161 @property
1162 def default_estimated_output_tokens(self) -> int | None: ... 1162 ↛ exitline 1162 didn't return from function 'default_estimated_output_tokens' because
1164 @property
1165 def default_estimated_output_tokens_per_model(self) -> Mapping[str, int] | None: ... 1165 ↛ exitline 1165 didn't return from function 'default_estimated_output_tokens_per_model' because
1167 @property
1168 def model_fields_set(self) -> Collection[str]: ... 1168 ↛ exitline 1168 didn't return from function 'model_fields_set' because
1171def _requested_output_token_estimates(
1172 data: OutputTokenEstimateRequest,
1173 existing_metadata: Mapping[str, object],
1174) -> tuple[object, object]:
1175 """The output-token estimates this request would leave stored on the entity.
1177 Mirrors how the management endpoints merge metadata: a supplied ``metadata``
1178 replaces the stored blob wholesale, an omitted one preserves it, and the
1179 dedicated top-level fields overlay whatever survives. Both sources are read
1180 because the same declaration reaches the same stored field either way.
1181 """
1182 base: Final[Mapping[str, object]] = (
1183 (data.metadata or {}) if "metadata" in data.model_fields_set else existing_metadata
1184 )
1185 return (
1186 data.default_estimated_output_tokens
1187 if data.default_estimated_output_tokens is not None
1188 else base.get(ESTIMATED_OUTPUT_TOKENS_FIELD),
1189 data.default_estimated_output_tokens_per_model
1190 if data.default_estimated_output_tokens_per_model is not None
1191 else base.get(ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD),
1192 )
1195def enforce_output_token_estimates_are_admin_only(
1196 data: OutputTokenEstimateRequest,
1197 existing_metadata: Mapping[str, object] | None,
1198 user_api_key_dict: UserAPIKeyAuth,
1199 entity: Literal["key", "team"],
1200) -> None:
1201 """Only a proxy admin may change what a key or team declares its models emit.
1203 That declaration is what the TPM limiter reserves for a request omitting
1204 ``max_tokens``, so lowering or clearing it under-reserves against every
1205 window the request is charged against, including the team and organization
1206 ones the writer may not own. A key's metadata is writable by its holder and
1207 a team's by its team admin, so neither is a trustworthy source for a value
1208 that weakens a limit set above them. Gated on the resulting value rather
1209 than on presence, so a form resending the stored declaration stays a no-op.
1210 """
1211 if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value: 1211 ↛ 1213line 1211 didn't jump to line 1213 because the condition on line 1211 was always true
1212 return
1213 stored: Final[Mapping[str, object]] = existing_metadata or {}
1214 if _requested_output_token_estimates(data, stored) == (
1215 stored.get(ESTIMATED_OUTPUT_TOKENS_FIELD),
1216 stored.get(ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD),
1217 ):
1218 return
1219 raise HTTPException(
1220 status_code=403,
1221 detail={
1222 "error": f"Only proxy admins can set {ESTIMATED_OUTPUT_TOKENS_FIELD} or "
1223 f"{ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD} on a {entity}. They decide how many output tokens "
1224 "the rate limiter reserves for a request that omits max_tokens."
1225 },
1226 )
1229class BatchEnqueuedTokenLimitRequest(Protocol):
1230 """The shape of any management request that can carry a batch enqueued-token limit."""
1232 @property
1233 def metadata(self) -> Mapping[str, object] | None: ... 1233 ↛ exitline 1233 didn't return from function 'metadata' because
1235 @property
1236 def model_fields_set(self) -> Collection[str]: ... 1236 ↛ exitline 1236 didn't return from function 'model_fields_set' because
1239def enforce_batch_enqueued_token_limit_is_admin_only(
1240 data: BatchEnqueuedTokenLimitRequest,
1241 existing_metadata: Mapping[str, object] | None,
1242 user_api_key_dict: UserAPIKeyAuth,
1243 entity: Literal["key", "team"],
1244) -> None:
1245 """Only a proxy admin may change a key or team's batch enqueued-token limit.
1247 When set, ``batch_enqueued_token_limit`` replaces the standard RPM/TPM checks
1248 for batch submissions, so a holder-writable copy would let a caller lift their
1249 own batch quota. Gated on the resulting value rather than on presence, so a
1250 form resending the stored value stays a no-op.
1251 """
1252 if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value: 1252 ↛ 1254line 1252 didn't jump to line 1254 because the condition on line 1252 was always true
1253 return
1254 stored: Final[Mapping[str, object]] = existing_metadata or EMPTY_MAPPING
1255 requested: Final[Mapping[str, object]] = (
1256 (data.metadata or EMPTY_MAPPING) if "metadata" in data.model_fields_set else stored
1257 )
1258 if requested.get(BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY) == stored.get(BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY):
1259 return
1260 raise HTTPException(
1261 status_code=403,
1262 detail={ # mutable-ok: HTTPException.detail has no immutable form
1263 "error": f"Only proxy admins can set {BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY} on a {entity}. "
1264 "It replaces the standard rate limit checks for batch submissions."
1265 },
1266 )
1269def get_model_rate_limit_from_metadata(
1270 user_api_key_dict: UserAPIKeyAuth,
1271 metadata_accessor_key: Literal["team_metadata", "organization_metadata", "project_metadata"],
1272 rate_limit_key: Literal["model_rpm_limit", "model_tpm_limit", "model_itpm_limit", "model_otpm_limit"],
1273) -> dict[str, int] | None:
1274 if getattr(user_api_key_dict, metadata_accessor_key): 1274 ↛ 1275line 1274 didn't jump to line 1275 because the condition on line 1274 was never true
1275 return getattr(user_api_key_dict, metadata_accessor_key).get(rate_limit_key)
1276 return None
1279def get_team_model_rpm_limit(
1280 user_api_key_dict: UserAPIKeyAuth,
1281) -> dict[str, int] | None:
1282 if user_api_key_dict.team_metadata:
1283 return user_api_key_dict.team_metadata.get("model_rpm_limit")
1284 return None
1287def get_team_model_tpm_limit(
1288 user_api_key_dict: UserAPIKeyAuth,
1289) -> dict[str, int] | None:
1290 if user_api_key_dict.team_metadata:
1291 return user_api_key_dict.team_metadata.get("model_tpm_limit")
1292 return None
1295def get_key_mcp_rpm_limit(
1296 user_api_key_dict: UserAPIKeyAuth,
1297) -> dict[str, int] | None:
1298 """
1299 Get the per-MCP-server rpm limit for a given api key.
1301 Priority order (returns first found):
1302 1. Key metadata (mcp_rpm_limit)
1303 2. Team metadata (mcp_rpm_limit)
1305 The returned dict is keyed by MCP server name (alias if set, else the
1306 configured server name).
1307 """
1308 if user_api_key_dict.metadata:
1309 result: Final = user_api_key_dict.metadata.get("mcp_rpm_limit")
1310 if result is not None:
1311 return result
1313 if user_api_key_dict.team_metadata:
1314 team_limit: Final = user_api_key_dict.team_metadata.get("mcp_rpm_limit")
1315 if team_limit is not None:
1316 return team_limit
1318 return None
1321def get_team_mcp_rpm_limit(
1322 user_api_key_dict: UserAPIKeyAuth,
1323) -> dict[str, int] | None:
1324 if user_api_key_dict.team_metadata:
1325 return user_api_key_dict.team_metadata.get("mcp_rpm_limit")
1326 return None
1329def get_key_tag_rpm_limit(
1330 user_api_key_dict: UserAPIKeyAuth,
1331) -> dict[str, int] | None:
1332 """
1333 Get the per-request-tag rpm limit configured on a given api key.
1335 The returned dict is keyed by request tag, so each tag/group tracked on
1336 the key gets its own independent RPM counter.
1337 """
1338 if user_api_key_dict.metadata: 1338 ↛ 1339line 1338 didn't jump to line 1339 because the condition on line 1338 was never true
1339 return user_api_key_dict.metadata.get("tag_rpm_limit")
1340 return None
1343def get_project_model_rpm_limit(
1344 user_api_key_dict: UserAPIKeyAuth,
1345) -> dict[str, int] | None:
1346 if user_api_key_dict.project_metadata:
1347 return user_api_key_dict.project_metadata.get("model_rpm_limit")
1348 return None
1351def get_project_model_tpm_limit(
1352 user_api_key_dict: UserAPIKeyAuth,
1353) -> dict[str, int] | None:
1354 if user_api_key_dict.project_metadata:
1355 return user_api_key_dict.project_metadata.get("model_tpm_limit")
1356 return None
1359def custom_auth_common_checks_warning(
1360 *,
1361 custom_auth_configured: bool,
1362 run_common_checks: bool,
1363) -> str | None:
1364 if not custom_auth_configured or run_common_checks: 1364 ↛ 1366line 1364 didn't jump to line 1366 because the condition on line 1364 was always true
1365 return None
1366 return (
1367 "custom_auth is configured but 'custom_auth_run_common_checks' is not set. "
1368 "Problem: budgets, model-access allowlists, and per-model rate limits configured "
1369 "on your DB team/project records will NOT be enforced for custom-auth requests "
1370 "(rate limits set directly on the returned UserAPIKeyAuth still apply). "
1371 "Fix: set 'general_settings.custom_auth_run_common_checks: true'. "
1372 "Docs: https://docs.litellm.ai/docs/proxy/custom_auth"
1373 )
1376_custom_auth_common_checks_warning_emitted = False
1379def warn_once_if_custom_auth_skips_common_checks(
1380 *,
1381 custom_auth_configured: bool,
1382 run_common_checks: bool,
1383 logger: Logger = verbose_proxy_logger,
1384) -> None:
1385 global _custom_auth_common_checks_warning_emitted
1386 if _custom_auth_common_checks_warning_emitted: 1386 ↛ 1387line 1386 didn't jump to line 1387 because the condition on line 1386 was never true
1387 return
1388 message: Final = custom_auth_common_checks_warning(
1389 custom_auth_configured=custom_auth_configured,
1390 run_common_checks=run_common_checks,
1391 )
1392 if message is None: 1392 ↛ 1394line 1392 didn't jump to line 1394 because the condition on line 1392 was always true
1393 return
1394 logger.warning(message)
1395 _custom_auth_common_checks_warning_emitted = True
1398def log_once_if_budget_reservation_disabled(
1399 *,
1400 disabled: bool,
1401 logger: Logger = verbose_proxy_logger,
1402) -> None:
1403 if constants.budget_reservation_disabled_info_emitted or not disabled: 1403 ↛ 1405line 1403 didn't jump to line 1405 because the condition on line 1403 was always true
1404 return
1405 logger.info(
1406 "disable_budget_reservation is enabled: skipping optimistic budget "
1407 "reservation. Budget enforcement is read-time only. Concurrent "
1408 "requests can each pass the spend check before their cost is recorded, "
1409 "so a configured budget may be briefly exceeded under high concurrency. "
1410 "Set disable_budget_reservation to False or remove it to restore "
1411 "hard per-request budget enforcement."
1412 )
1413 constants.budget_reservation_disabled_info_emitted = True
1416def is_pass_through_provider_route(route: str) -> bool:
1417 PROVIDER_SPECIFIC_PASS_THROUGH_ROUTES: Final = [
1418 "vertex-ai",
1419 ]
1421 # check if any of the prefixes are in the route
1422 for prefix in PROVIDER_SPECIFIC_PASS_THROUGH_ROUTES:
1423 if prefix in route:
1424 return True
1426 return False
1429def has_user_setup_sso() -> bool:
1430 """
1431 Check if the user has set up single sign-on (SSO).
1433 Covers OAuth providers (Microsoft, Google, generic) and SAML IdP metadata.
1434 Used by UI discovery (``sso_configured``) so the login button enables when
1435 any supported SSO path is configured — including SAML-only setups.
1436 """
1437 microsoft_client_id: Final = os.getenv("MICROSOFT_CLIENT_ID", None)
1438 google_client_id: Final = os.getenv("GOOGLE_CLIENT_ID", None)
1439 generic_client_id: Final = os.getenv("GENERIC_CLIENT_ID", None)
1440 saml_idp_metadata_url: Final = os.getenv("SAML_IDP_METADATA_URL", None)
1441 saml_idp_metadata_xml: Final = os.getenv("SAML_IDP_METADATA_XML", None)
1443 return (
1444 microsoft_client_id is not None
1445 or google_client_id is not None
1446 or generic_client_id is not None
1447 or bool(saml_idp_metadata_url)
1448 or bool(saml_idp_metadata_xml)
1449 )
1452def _is_google_ready() -> bool:
1453 return bool(os.getenv("GOOGLE_CLIENT_ID")) and bool(os.getenv("GOOGLE_CLIENT_SECRET"))
1456def _is_microsoft_ready() -> bool:
1457 return (
1458 bool(os.getenv("MICROSOFT_CLIENT_ID"))
1459 and bool(os.getenv("MICROSOFT_CLIENT_SECRET"))
1460 and bool(os.getenv("MICROSOFT_TENANT"))
1461 )
1464def _is_generic_oauth_ready() -> bool:
1465 return (
1466 bool(os.getenv("GENERIC_CLIENT_ID"))
1467 and bool(os.getenv("GENERIC_CLIENT_SECRET"))
1468 and bool(os.getenv("GENERIC_AUTHORIZATION_ENDPOINT"))
1469 and bool(os.getenv("GENERIC_TOKEN_ENDPOINT"))
1470 and bool(os.getenv("GENERIC_USERINFO_ENDPOINT"))
1471 )
1474def _is_saml_ready() -> bool:
1475 if not (os.getenv("SAML_IDP_METADATA_URL") or os.getenv("SAML_IDP_METADATA_XML")):
1476 return False
1477 # SAML's runtime (python3-saml) is an optional dependency; the SAML
1478 # handler itself fails closed on every request when it is missing
1479 # (SAMLAuthHandler raises before touching the IdP), so metadata alone
1480 # is not "ready" either. find_spec raises ModuleNotFoundError (rather
1481 # than returning None) when the top-level package is absent entirely,
1482 # so this must not be a bare boolean expression or every password
1483 # login would 500 on a deployment that configured SAML metadata
1484 # without installing the optional extra.
1485 try:
1486 return importlib.util.find_spec("onelogin.saml2.auth") is not None
1487 except ModuleNotFoundError:
1488 return False
1491def is_sso_provider_fully_configured() -> bool:
1492 """Whether ANY configured SSO provider has every companion setting it
1493 needs to actually authenticate a user, not merely a client id.
1495 A lone ``MICROSOFT_CLIENT_ID`` with no secret or tenant makes
1496 ``has_user_setup_sso()`` return True while every real sign-in attempt
1497 fails, so a gate that BLOCKS the password fallback (unlike the UI
1498 discovery use of ``has_user_setup_sso()``, where a dead login button is
1499 merely confusing) must check readiness here, or it can lock every admin
1500 out with no way to sign in at all. Checks every provider independently
1501 (mirroring ``/sso/readiness``'s per-provider requirements) rather than
1502 stopping at the first one with a client id set, so a stray leftover
1503 client id for an unused provider can never mask a different, fully
1504 configured provider that would otherwise satisfy this gate.
1505 """
1506 return _is_google_ready() or _is_microsoft_ready() or _is_generic_oauth_ready() or _is_saml_ready()
1509def get_customer_user_header_from_mapping(user_id_mapping) -> list | None:
1510 """Return the header_name mapped to CUSTOMER role, if any (dict-based)."""
1511 if not user_id_mapping:
1512 return None
1513 items: Final = user_id_mapping if isinstance(user_id_mapping, list) else [user_id_mapping]
1514 customer_headers_mappings: Final = []
1515 for item in items:
1516 if not isinstance(item, dict):
1517 continue
1518 role = item.get("litellm_user_role")
1519 header_name = item.get("header_name")
1520 if role is None or not header_name:
1521 continue
1522 if str(role).lower() == str(LitellmUserRoles.CUSTOMER).lower():
1523 customer_headers_mappings.append(header_name.lower())
1525 if customer_headers_mappings:
1526 return customer_headers_mappings
1528 return None
1531def _get_customer_id_from_standard_headers(
1532 request_headers: Mapping[str, object] | None,
1533) -> str | None:
1534 """
1535 Check standard customer ID headers for a customer/end-user ID.
1537 This enables tools like Claude Code to pass customer IDs via ANTHROPIC_CUSTOM_HEADERS.
1538 No configuration required - these headers are always checked.
1540 Args:
1541 request_headers: The request headers dict
1543 Returns:
1544 The customer ID if found in standard headers, None otherwise
1545 """
1546 if request_headers is None: 1546 ↛ 1547line 1546 didn't jump to line 1547 because the condition on line 1546 was never true
1547 return None
1549 for standard_header in STANDARD_CUSTOMER_ID_HEADERS:
1550 for header_name, header_value in request_headers.items():
1551 if header_name.lower() == standard_header.lower(): 1551 ↛ 1552line 1551 didn't jump to line 1552 because the condition on line 1551 was never true
1552 user_id_str = _coerce_user_id_to_str(header_value)
1553 if user_id_str:
1554 return user_id_str
1555 return None
1558def _coerce_user_id_to_str(value: object) -> str | None:
1559 """Return a usable end-user identifier string, or None if the value isn't one.
1561 Always drops non-string structured values (dict/list/tuple/set) because
1562 stringifying them produces garbage spend-log rows like
1563 ``"{'device_id': ...}"``. Strings that *decode* to a structured payload
1564 are only rejected when ``litellm.validate_end_user_id_in_db`` is enabled
1565 — operators who currently pass JSON-encoded identifiers keep their
1566 existing behavior until they opt in. See
1567 auth_utils.py:get_end_user_id_from_request_body for the extraction chain.
1568 """
1569 if value is None:
1570 return None
1571 if isinstance(value, bool):
1572 # bool is an int subclass; handle explicitly to avoid "True"/"False".
1573 return None
1574 if isinstance(value, (int, float)):
1575 return str(value)
1576 if isinstance(value, str):
1577 stripped: Final = value.strip()
1578 if not stripped:
1579 return None
1580 # Reject strings that decode to a structured payload (JSON object/array)
1581 # only when the operator has opted into end-user validation. Gating
1582 # behind the flag preserves backwards compatibility for deployments
1583 # that intentionally pass JSON-encoded user identifiers.
1584 if litellm.validate_end_user_id_in_db and stripped[:1] in ("{", "["): 1584 ↛ 1585line 1584 didn't jump to line 1585 because the condition on line 1584 was never true
1585 parsed: Final[object] = safe_json_loads(stripped)
1586 if isinstance(parsed, (dict, list)):
1587 return None
1588 return stripped
1589 # dict, list, tuple, set, arbitrary objects -> drop.
1590 return None
1593def get_end_user_id_from_request_body(
1594 request_body: Mapping[str, object], request_headers: Mapping[str, object] | None = None
1595) -> str | None:
1596 # Import general_settings here to avoid potential circular import issues at module level
1597 # and to ensure it's fetched at runtime.
1598 from litellm.proxy.proxy_server import general_settings
1600 # Check 1: Standard customer ID headers (always checked, no configuration required)
1601 customer_id: Final = _get_customer_id_from_standard_headers(request_headers=request_headers)
1602 if customer_id is not None: 1602 ↛ 1603line 1602 didn't jump to line 1603 because the condition on line 1602 was never true
1603 return customer_id
1605 # Check 2: Follow the user header mappings feature, if not found, then check for deprecated user_header_name (only if request_headers is provided)
1606 # User query: "system not respecting user_header_name property"
1607 # This implies the key in general_settings is 'user_header_name'.
1608 if request_headers is not None: 1608 ↛ 1639line 1608 didn't jump to line 1639 because the condition on line 1608 was always true
1609 custom_header_name_to_check: list | str | None = None
1611 # Prefer user mappings (new behavior)
1612 user_id_mapping: Final = general_settings.get("user_header_mappings", None)
1613 if user_id_mapping: 1613 ↛ 1614line 1613 didn't jump to line 1614 because the condition on line 1613 was never true
1614 custom_header_name_to_check = get_customer_user_header_from_mapping(user_id_mapping)
1616 # Fallback to deprecated user_header_name if mapping did not specify
1617 if not custom_header_name_to_check: 1617 ↛ 1624line 1617 didn't jump to line 1624 because the condition on line 1617 was always true
1618 user_id_header_config_key: Final = "user_header_name"
1619 value: Final = general_settings.get(user_id_header_config_key)
1620 if isinstance(value, str) and value.strip() != "": 1620 ↛ 1621line 1620 didn't jump to line 1621 because the condition on line 1620 was never true
1621 custom_header_name_to_check = value
1623 # If we have a header name to check, try to read it from request headers
1624 if isinstance(custom_header_name_to_check, list): 1624 ↛ 1625line 1624 didn't jump to line 1625 because the condition on line 1624 was never true
1625 headers_lower: Final = {k.lower(): v for k, v in request_headers.items()}
1626 for expected_header in custom_header_name_to_check:
1627 user_id_str = _coerce_user_id_to_str(headers_lower.get(expected_header))
1628 if user_id_str:
1629 return user_id_str
1631 elif isinstance(custom_header_name_to_check, str): 1631 ↛ 1632line 1631 didn't jump to line 1632 because the condition on line 1631 was never true
1632 for header_name, header_value in request_headers.items():
1633 if header_name.lower() == custom_header_name_to_check.lower():
1634 user_id_str = _coerce_user_id_to_str(header_value)
1635 if user_id_str:
1636 return user_id_str
1638 # Check 3: 'user' field in request_body (commonly OpenAI)
1639 if "user" in request_body:
1640 user_id_str = _coerce_user_id_to_str(request_body["user"])
1641 if user_id_str:
1642 return user_id_str
1644 def _as_dict(value: object) -> dict:
1645 # metadata / litellm_metadata can arrive as JSON strings from
1646 # multipart/form-data or extra_body; coerce so string-encoded
1647 # payloads can't evade end-user attribution.
1648 if isinstance(value, dict):
1649 return value
1650 if isinstance(value, str):
1651 parsed: Final = safe_json_loads(value)
1652 return parsed if isinstance(parsed, dict) else {}
1653 return {}
1655 # Check 4: 'litellm_metadata.user' in request_body (commonly Anthropic)
1656 litellm_metadata: Final = _as_dict(request_body.get("litellm_metadata"))
1657 user_id_str = _coerce_user_id_to_str(litellm_metadata.get("user"))
1658 if user_id_str: 1658 ↛ 1659line 1658 didn't jump to line 1659 because the condition on line 1658 was never true
1659 return user_id_str
1661 # Check 5: 'metadata.user_id' in request_body (another common pattern)
1662 metadata_dict: Final = _as_dict(request_body.get("metadata"))
1663 user_id_str = _coerce_user_id_to_str(metadata_dict.get("user_id"))
1664 if user_id_str: 1664 ↛ 1665line 1664 didn't jump to line 1665 because the condition on line 1664 was never true
1665 return user_id_str
1667 # Check 6: 'safety_identifier' in request body (OpenAI Responses API parameter)
1668 # SECURITY NOTE: safety_identifier can be set by any caller in the request body.
1669 # Only use this for end-user identification in trusted environments where you control
1670 # the calling application. For untrusted callers, prefer using headers or server-side
1671 # middleware to set the end_user_id to prevent impersonation.
1672 user_id_str = _coerce_user_id_to_str(request_body.get("safety_identifier"))
1673 if user_id_str:
1674 return user_id_str
1676 return None
1679MODEL_ROUTING_HEADER_NAME: Final = "x-litellm-model"
1680_MODEL_ROUTING_ROUTE_MARKERS: Final = (
1681 "/files",
1682 "/batches",
1683 "/vector_stores",
1684 "/skills",
1685 "/evals",
1686 "/fine_tuning",
1687 "/videos",
1688)
1689_MODEL_ROUTING_HEADER_OR_QUERY_ROUTE_MARKERS: Final = (
1690 "/files",
1691 "/batches",
1692 "/skills",
1693 "/evals",
1694)
1695_MODEL_ROUTING_QUERY_TARGET_MODEL_ROUTE_MARKERS: Final = (
1696 "/files",
1697 "/batches",
1698 "/fine_tuning",
1699)
1700_MODEL_ROUTING_BODY_TARGET_MODEL_ROUTE_MARKERS: Final = (
1701 "/files",
1702 "/batches",
1703 "/vector_stores",
1704)
1705_MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS: Final = ("/evals",)
1706# Realtime WebRTC routes carry the effective model inside the nested
1707# ``session.model`` field (see realtime_endpoints.endpoints), so the model the
1708# request will actually use is not present at the top level. Extract it here so
1709# can_key_call_model() validates the real target model.
1710_MODEL_ROUTING_SESSION_MODEL_ROUTE_MARKERS: Final = (
1711 "/realtime/client_secrets",
1712 "/realtime/calls",
1713)
1714_MODEL_ROUTING_ID_FIELDS: Final = (
1715 "file_id",
1716 "input_file_id",
1717 "output_file_id",
1718 "error_file_id",
1719 "batch_id",
1720 "fine_tuning_job_id",
1721 "training_file",
1722 "validation_file",
1723 "vector_store_id",
1724 "video_id",
1725 "character_id",
1726)
1729def _append_model_candidates(candidates: list[str], value: object) -> None:
1730 if value is None:
1731 return
1733 values: Final[tuple[object, ...]] = tuple(value) if isinstance(value, (list, tuple, set)) else (value,)
1734 for item in values:
1735 if item is None:
1736 continue
1737 if isinstance(item, str):
1738 model_names = [model.strip() for model in item.split(",")]
1739 else:
1740 model_names = [str(item).strip()]
1741 candidates.extend(model for model in model_names if model)
1744def _dedupe_model_candidates(candidates: Collection[str]) -> list[str]:
1745 deduped: Final[list[str]] = []
1746 for model in candidates:
1747 if model not in deduped:
1748 deduped.append(model)
1749 return deduped
1752def _get_case_insensitive_mapping_value(mapping: Mapping[str, object] | None, key: str) -> object:
1753 if not mapping:
1754 return None
1755 if key in mapping:
1756 return mapping[key]
1757 key_lower: Final = key.lower()
1758 for mapping_key, value in mapping.items():
1759 if str(mapping_key).lower() == key_lower: 1759 ↛ 1760line 1759 didn't jump to line 1760 because the condition on line 1759 was never true
1760 return value
1761 return None
1764def _route_matches_any_marker(route: str, markers: tuple[str, ...]) -> bool:
1765 normalized_route: Final = route.lower()
1766 return any(marker in normalized_route for marker in markers)
1769def _route_uses_model_routing_sources(route: str) -> bool:
1770 return _route_matches_any_marker(route=route, markers=_MODEL_ROUTING_ROUTE_MARKERS)
1773def _extract_models_from_managed_resource_id(
1774 resource_id: object,
1775 resource_id_field: str | None = None,
1776 llm_router: Router | None = None,
1777) -> list[str]:
1778 if not isinstance(resource_id, str) or not resource_id:
1779 return []
1781 candidates: Final[list[str]] = []
1783 try:
1784 from litellm.proxy.openai_files_endpoints.common_utils import (
1785 _is_base64_encoded_unified_file_id,
1786 decode_model_from_file_id,
1787 get_model_id_from_unified_batch_id,
1788 get_models_from_unified_file_id,
1789 )
1791 _append_model_candidates(candidates=candidates, value=decode_model_from_file_id(resource_id))
1792 unified_file_id: Final = _is_base64_encoded_unified_file_id(resource_id)
1793 if unified_file_id: 1793 ↛ 1794line 1793 didn't jump to line 1794 because the condition on line 1793 was never true
1794 _append_model_candidates(
1795 candidates=candidates,
1796 value=get_models_from_unified_file_id(unified_file_id),
1797 )
1798 _append_model_candidates(
1799 candidates=candidates,
1800 value=_resolve_model_id_with_router(get_model_id_from_unified_batch_id(unified_file_id), llm_router),
1801 )
1802 except Exception as e:
1803 verbose_proxy_logger.debug("Unable to extract model from managed file/batch ID: %s", str(e))
1805 try:
1806 from litellm.llms.base_llm.managed_resources.utils import parse_unified_id
1808 parsed_id: Final = parse_unified_id(resource_id)
1809 if parsed_id: 1809 ↛ 1810line 1809 didn't jump to line 1810 because the condition on line 1809 was never true
1810 _append_model_candidates(
1811 candidates=candidates,
1812 value=_resolve_model_id_with_router(parsed_id.get("model_id"), llm_router),
1813 )
1814 _append_model_candidates(candidates=candidates, value=parsed_id.get("target_model_names"))
1815 except Exception as e:
1816 verbose_proxy_logger.debug("Unable to extract model from unified managed resource ID: %s", str(e))
1818 if resource_id_field in ("video_id", "character_id"):
1819 try:
1820 from litellm.types.videos.utils import (
1821 decode_character_id_with_provider,
1822 decode_video_id_with_provider,
1823 )
1825 if resource_id_field == "video_id":
1826 model_id = decode_video_id_with_provider(resource_id).get("model_id")
1827 _append_model_candidates(
1828 candidates=candidates,
1829 value=_resolve_model_id_with_router(model_id, llm_router),
1830 )
1831 else:
1832 model_id = decode_character_id_with_provider(resource_id).get("model_id")
1833 _append_model_candidates(
1834 candidates=candidates,
1835 value=_resolve_model_id_with_router(model_id, llm_router),
1836 )
1837 except Exception as e:
1838 verbose_proxy_logger.debug("Unable to extract model from managed video/character ID: %s", str(e))
1840 return _dedupe_model_candidates(candidates)
1843def _resolve_model_id_with_router(model_id: str | None, llm_router: Router | None) -> str | None:
1844 if model_id is None or llm_router is None: 1844 ↛ 1846line 1844 didn't jump to line 1846 because the condition on line 1844 was always true
1845 return model_id
1846 try:
1847 return llm_router.resolve_model_name_from_model_id(model_id) or model_id
1848 except Exception as e:
1849 verbose_proxy_logger.debug("Unable to resolve model_id from managed resource ID: %s", str(e))
1850 return model_id
1853def get_cache_prediction_deployments(
1854 *, current_deployment_id: str, candidate_deployment_id: str, llm_router: Router, team_id: str | None
1855) -> tuple[Deployment, Deployment] | None:
1856 current: Final = llm_router.get_deployment(current_deployment_id)
1857 candidate: Final = llm_router.get_deployment(candidate_deployment_id)
1858 if current is None or candidate is None: 1858 ↛ 1860line 1858 didn't jump to line 1860 because the condition on line 1858 was always true
1859 return None
1860 if any(deployment.model_info.team_id not in (None, team_id) for deployment in (current, candidate)):
1861 return None
1862 return current, candidate
1865def _cache_prediction_model_candidates(
1866 request_data: Mapping[str, object], llm_router: Router | None, team_id: str | None
1867) -> tuple[str, ...]:
1868 current_id: Final = request_data.get("current_deployment_id")
1869 candidate_id: Final = request_data.get("candidate_deployment_id")
1870 if llm_router is None or not isinstance(current_id, str) or not isinstance(candidate_id, str):
1871 return ()
1872 deployments: Final = get_cache_prediction_deployments(
1873 current_deployment_id=current_id, candidate_deployment_id=candidate_id, llm_router=llm_router, team_id=team_id
1874 )
1875 return tuple(deployment.model_name for deployment in deployments) if deployments is not None else ()
1878def _extract_model_candidates_from_request(
1879 request_data: dict,
1880 route: str,
1881 request_headers: Mapping[str, object] | None = None,
1882 request_query_params: Mapping[str, object] | None = None,
1883 llm_router: Router | None = None,
1884 team_id: str | None = None,
1885) -> list[str]:
1886 if route == "/cost/predict-cache":
1887 prediction_models: Final = _cache_prediction_model_candidates(request_data, llm_router, team_id) # pyright: ignore[reportUnknownArgumentType] # the typed reader validates each deployment ID from this legacy payload
1888 return _dedupe_model_candidates(prediction_models)
1889 candidates: Final[list[str]] = []
1890 uses_model_routing_sources: Final = _route_uses_model_routing_sources(route=route)
1891 uses_header_or_query_model_sources: Final = _route_matches_any_marker(
1892 route=route, markers=_MODEL_ROUTING_HEADER_OR_QUERY_ROUTE_MARKERS
1893 )
1894 uses_query_target_model_sources: Final = _route_matches_any_marker(
1895 route=route, markers=_MODEL_ROUTING_QUERY_TARGET_MODEL_ROUTE_MARKERS
1896 )
1897 uses_body_target_model_sources: Final = _route_matches_any_marker(
1898 route=route, markers=_MODEL_ROUTING_BODY_TARGET_MODEL_ROUTE_MARKERS
1899 )
1900 uses_completion_model_sources: Final = _route_matches_any_marker(
1901 route=route, markers=_MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS
1902 )
1904 body_model: Final = request_data.get("model")
1905 _append_model_candidates(candidates, body_model)
1906 if uses_body_target_model_sources or not body_model:
1907 _append_model_candidates(candidates, request_data.get("target_model_names"))
1908 if _route_matches_any_marker(route=route, markers=_MODEL_ROUTING_SESSION_MODEL_ROUTE_MARKERS):
1909 session: Final = request_data.get("session")
1910 if isinstance(session, dict): 1910 ↛ 1911line 1910 didn't jump to line 1911 because the condition on line 1910 was never true
1911 _append_model_candidates(candidates, session.get("model"))
1912 if uses_completion_model_sources and isinstance(request_data.get("completion"), dict): 1912 ↛ 1913line 1912 didn't jump to line 1913 because the condition on line 1912 was never true
1913 _append_model_candidates(candidates, request_data["completion"].get("model"))
1915 if uses_model_routing_sources:
1916 if uses_header_or_query_model_sources:
1917 _append_model_candidates(
1918 candidates,
1919 _get_case_insensitive_mapping_value(request_query_params, "model"),
1920 )
1921 _append_model_candidates(
1922 candidates,
1923 _get_case_insensitive_mapping_value(request_headers, MODEL_ROUTING_HEADER_NAME),
1924 )
1925 if uses_query_target_model_sources:
1926 _append_model_candidates(
1927 candidates,
1928 _get_case_insensitive_mapping_value(request_query_params, "target_model_names"),
1929 )
1931 for field in _MODEL_ROUTING_ID_FIELDS:
1932 _append_model_candidates(
1933 candidates,
1934 _extract_models_from_managed_resource_id(
1935 request_data.get(field),
1936 resource_id_field=field,
1937 llm_router=llm_router,
1938 ),
1939 )
1941 return _dedupe_model_candidates(candidates)
1944def _format_model_candidates(
1945 candidates: list[str],
1946) -> str | list[str] | None:
1947 if not candidates:
1948 return None
1949 if len(candidates) == 1:
1950 return candidates[0]
1951 return candidates
1954def request_dispatched_to_pass_through_endpoint(request: Request | None) -> bool:
1955 """Whether FastAPI resolved this request to a user-defined pass-through handler.
1957 Reads the marker set by ``create_pass_through_route`` off the dispatched endpoint
1958 (``request.scope["endpoint"]``). Because routing has already run by the time auth
1959 dependencies execute, this reflects the handler that actually serves the request:
1960 a custom path colliding with a built-in route resolves to the built-in handler,
1961 which carries no marker, so model-access checks are never wrongly skipped.
1962 """
1963 if request is None:
1964 return False
1965 scope: Final = getattr(request, "scope", None)
1966 if not isinstance(scope, dict): 1966 ↛ 1967line 1966 didn't jump to line 1967 because the condition on line 1966 was never true
1967 return False
1968 endpoint: Final = scope.get("endpoint")
1969 # Identity check against True (not truthiness): the marker is set to the literal
1970 # True, and this keeps a spec'd Mock request (whose attribute access yields truthy
1971 # child mocks) from being misread as a pass-through dispatch.
1972 return getattr(endpoint, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, False) is True
1975def request_dispatched_to_provider_pass_through(request: Request) -> bool:
1976 """Built-in provider pass-through handlers (``/anthropic/{endpoint:path}``, ...) bind ``endpoint``."""
1977 return "endpoint" in request.path_params
1980def get_model_from_request(
1981 request_data: dict,
1982 route: str,
1983 request_headers: Mapping[str, object] | None = None,
1984 request_query_params: Mapping[str, object] | None = None,
1985 llm_router: Router | None = None,
1986 request: Request | None = None,
1987 team_id: str | None = None,
1988) -> str | list[str] | None:
1989 """Resolve the model(s) a request targets, for model-access and budget checks.
1991 Returns ``None`` when the request was dispatched to a user-defined pass-through
1992 endpoint: its body is forwarded verbatim to the configured upstream, so a
1993 ``model`` field there names an upstream model, not a LiteLLM-managed one, and
1994 enforcing key/team model allowlists against it would reject valid requests. The
1995 check reads the FastAPI-resolved endpoint (``request.scope["endpoint"]``), not the
1996 request path, so a custom path that collides with a built-in route never
1997 suppresses model-access checks: on a collision the built-in handler is dispatched
1998 and does not carry the marker. Built-in provider passthrough routes
1999 (``/vertex_ai``, ``/gemini``, ...) are separate handlers and keep model enforcement.
2000 """
2001 if request_dispatched_to_pass_through_endpoint(request):
2002 return None
2004 candidates: Final = _extract_model_candidates_from_request(
2005 request_data=request_data,
2006 route=route,
2007 request_headers=request_headers,
2008 request_query_params=request_query_params,
2009 llm_router=llm_router,
2010 team_id=team_id,
2011 )
2012 model = _format_model_candidates(candidates)
2014 # If no explicit model was found, try to extract from route
2015 if model is None:
2016 # Parse model from route that follows the pattern /openai/deployments/{model}/*
2017 match: Final = re.match(r"/openai/deployments/([^/]+)", route)
2018 if match:
2019 model = match.group(1)
2021 # If still not found, extract model from Google generateContent-style routes.
2022 # These routes put the model in the path and allow "/" inside the model id.
2023 # Examples:
2024 # - /v1beta/models/gemini-2.0-flash:generateContent
2025 # - /v1beta/models/bedrock/claude-sonnet-3.7:generateContent
2026 # - /models/custom/ns/model:streamGenerateContent
2027 if model is None and not route.lower().startswith("/vertex"):
2028 google_match = re.search(r"/(?:v1beta|beta)/models/([^:]+):", route)
2029 if google_match:
2030 model = google_match.group(1)
2032 if model is None and not route.lower().startswith("/vertex"):
2033 google_match = re.search(r"^/models/([^:]+):", route)
2034 if google_match:
2035 model = google_match.group(1)
2037 # If still not found, extract from Vertex AI passthrough route
2038 # Pattern: /vertex_ai/.../models/{model_id}:*
2039 # Example: /vertex_ai/v1/.../models/gemini-1.5-pro:generateContent
2040 if model is None and route.lower().startswith("/vertex"):
2041 vertex_match: Final = re.search(r"/models/([^:]+)", route)
2042 if vertex_match: 2042 ↛ 2043line 2042 didn't jump to line 2043 because the condition on line 2042 was never true
2043 model = vertex_match.group(1)
2045 if route.lower().startswith("/bedrock"):
2046 bedrock_model: Final = _model_from_bedrock_route(route)
2047 return model if bedrock_model is None else bedrock_model
2049 if route.lower().startswith(("/azure/", "/azure_ai/")):
2050 azure_model: Final = _router_model_from_azure_route(route, llm_router)
2051 return model if azure_model is None else azure_model
2053 if route.lower().startswith("/nvidia_nim/"):
2054 nvidia_nim_model: Final = (
2055 nvidia_nim_model_group_in_path(route, llm_router.get_model_list()) if llm_router else None
2056 )
2057 return model if nvidia_nim_model is None else nvidia_nim_model
2059 return model
2062def _router_model_from_azure_route(route: str, llm_router: Router | None) -> str | None:
2063 if llm_router is None: 2063 ↛ 2064line 2063 didn't jump to line 2064 because the condition on line 2063 was never true
2064 return None
2065 endpoint: Final = re.sub(r"^/azure(?:_ai)?/", "", route, flags=re.IGNORECASE)
2066 return azure_router_model_in_endpoint(endpoint, frozenset(llm_router.get_model_names()))
2069def _model_from_bedrock_route(route: str) -> str | None:
2070 from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
2071 _extract_model_from_bedrock_endpoint,
2072 is_bedrock_count_tokens_endpoint,
2073 )
2075 bedrock_endpoint: Final = re.sub(r"^/bedrock/", "", route, flags=re.IGNORECASE)
2076 if is_bedrock_count_tokens_endpoint(bedrock_endpoint): 2076 ↛ 2077line 2076 didn't jump to line 2077 because the condition on line 2076 was never true
2077 return None
2078 try:
2079 return _extract_model_from_bedrock_endpoint(bedrock_endpoint)
2080 except ValueError:
2081 return None
2084def abbreviate_api_key(api_key: str) -> str:
2085 if len(api_key) < MINIMUM_CUSTOM_KEY_LENGTH: 2085 ↛ 2086line 2085 didn't jump to line 2086 because the condition on line 2085 was never true
2086 return "sk-..."
2087 return f"sk-...{api_key[-4:]}"