mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-19 00:01:29 +00:00
3522 lines
155 KiB
Python
3522 lines
155 KiB
Python
"""
|
|
This file handles authentication for the LiteLLM Proxy.
|
|
|
|
it checks if the user passed a valid API Key to the LiteLLM Proxy
|
|
|
|
Returns a UserAPIKeyAuth object if the API key is valid
|
|
|
|
"""
|
|
|
|
import asyncio
|
|
import fnmatch
|
|
import re
|
|
import secrets
|
|
from collections.abc import Mapping
|
|
from datetime import datetime, timezone
|
|
from typing import Any, Final, NamedTuple, Protocol, Union, cast
|
|
|
|
import fastapi
|
|
import orjson
|
|
from fastapi import HTTPException, Request, WebSocket, status
|
|
from fastapi.security.api_key import APIKeyHeader
|
|
from starlette.exceptions import WebSocketException
|
|
|
|
import litellm
|
|
from litellm._logging import verbose_logger, verbose_proxy_logger
|
|
from litellm._service_logger import ServiceLogging
|
|
from litellm.caching.redis_cache import RedisCache
|
|
from litellm.constants import (
|
|
GLOBAL_PROXY_SPEND_CACHE_KEY,
|
|
INVALID_VIRTUAL_KEY_ERROR_MARKER,
|
|
INVALID_VIRTUAL_KEY_ERROR_MESSAGE,
|
|
LITELLM_PROXY_BUDGET_NAME,
|
|
LITELLM_PROXY_MASTER_KEY_ALIAS,
|
|
)
|
|
from litellm.integrations.otel.model.config import is_otel_v2_enabled
|
|
from litellm.integrations.otel.runtime import phase_span, seed_request_identity
|
|
from litellm.litellm_core_utils.dd_tracing import tracer
|
|
from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value
|
|
from litellm.proxy._types import *
|
|
from litellm.proxy.auth.auth_checks import (
|
|
ExperimentalUIJWTToken,
|
|
TeamNotFoundError,
|
|
_cache_key_object,
|
|
_can_object_call_model,
|
|
_check_end_user_budget,
|
|
_delete_cache_key_object,
|
|
_get_user_role,
|
|
_is_model_cost_zero,
|
|
_is_user_proxy_admin,
|
|
_virtual_key_max_budget_alert_check,
|
|
_virtual_key_max_budget_check,
|
|
_virtual_key_soft_budget_check,
|
|
can_key_call_model,
|
|
common_checks,
|
|
get_end_user_object,
|
|
get_jwt_key_mapping_object,
|
|
get_object_permission,
|
|
get_project_object,
|
|
get_team_membership,
|
|
get_team_object,
|
|
get_user_object,
|
|
is_valid_fallback_model,
|
|
jwt_key_mapping_cache_key,
|
|
resolve_and_validate_end_user_id,
|
|
)
|
|
from litellm.proxy.auth.auth_exception_handler import UserAPIKeyAuthExceptionHandler
|
|
from litellm.proxy.auth.auth_method import AuthMethod
|
|
from litellm.proxy.auth.auth_object_prefetch import AuthObjectRefs, prefetch_auth_objects
|
|
from litellm.proxy.auth.auth_utils import (
|
|
abbreviate_api_key,
|
|
get_end_user_id_from_request_body,
|
|
get_model_from_request,
|
|
get_request_route,
|
|
get_request_route_template,
|
|
is_invalid_virtual_key_error,
|
|
iter_request_fallback_targets,
|
|
normalize_request_route,
|
|
pre_db_read_auth_checks,
|
|
route_in_additonal_public_routes,
|
|
)
|
|
from litellm.proxy.auth.handle_jwt import JWTAuthManager, JWTHandler
|
|
from litellm.proxy.auth.network import TrustedProxyConfig, resolve_network_context
|
|
from litellm.proxy.auth.oauth2_check import Oauth2Handler
|
|
from litellm.proxy.auth.oauth2_proxy_hook import handle_oauth2_proxy_request
|
|
from litellm.proxy.auth.resolvers import CredentialRef, Principal
|
|
from litellm.proxy.auth.resolvers.grants import (
|
|
GrantResolver,
|
|
LookupDegraded,
|
|
ResolvedGrants,
|
|
UserLookup,
|
|
raise_public,
|
|
user_models,
|
|
)
|
|
from litellm.proxy.auth.resolvers.store import IdentityStore
|
|
from litellm.proxy.auth.route_checks import RouteChecks
|
|
from litellm.proxy.auth.team_grants import team_grants
|
|
from litellm.proxy.auth.trusted_proxy_utils import get_trusted_proxy_cidrs
|
|
from litellm.proxy.common_utils.cache_coordinator import EventDrivenCacheCoordinator
|
|
from litellm.proxy.common_utils.http_parsing_utils import (
|
|
_read_request_body,
|
|
_safe_get_request_headers,
|
|
_safe_get_request_query_params,
|
|
_safe_set_request_parsed_body,
|
|
populate_request_with_path_params,
|
|
read_raw_json_body,
|
|
)
|
|
from litellm.proxy.common_utils.model_listing_utils import claude_code_requested_group
|
|
from litellm.proxy.common_utils.realtime_utils import _realtime_request_body
|
|
from litellm.proxy.common_utils.user_api_key_cache import (
|
|
UserApiKeyCache,
|
|
team_membership_auth_cache_key,
|
|
)
|
|
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
|
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
|
|
from litellm.proxy.spend_tracking.carried_budget_state import carry_team_and_user_budget_state
|
|
from litellm.proxy.spend_tracking.spend_counter_batch import (
|
|
bind_admission_counter_keys,
|
|
release_spend_counter_batch,
|
|
spend_counter_batch_scope,
|
|
)
|
|
from litellm.proxy.utils import (
|
|
PrismaClient,
|
|
ProxyLogging,
|
|
normalize_route_for_root_path,
|
|
)
|
|
from litellm.repositories.table_repositories import TeamMembershipRepository
|
|
from litellm.secret_managers.main import get_secret_bool
|
|
from litellm.types.services import ServiceTypes
|
|
|
|
try:
|
|
from litellm_enterprise.proxy.auth.user_api_key_auth import (
|
|
enterprise_custom_auth as _enterprise_custom_auth,
|
|
)
|
|
|
|
enterprise_custom_auth: Callable | None = _enterprise_custom_auth
|
|
except ImportError as e:
|
|
verbose_proxy_logger.debug("Error in enterprise custom auth: %s", e)
|
|
enterprise_custom_auth = None
|
|
|
|
user_api_key_service_logger_obj: Final = ServiceLogging() # used for tracking latency on OTEL
|
|
|
|
|
|
def _normalize_public_auth_route(route: str) -> str:
|
|
if route != "/" and route.endswith("/"):
|
|
return route.rstrip("/")
|
|
return route
|
|
|
|
|
|
def _route_requires_auth_despite_public(route: str, general_settings: dict | None) -> bool:
|
|
normalized_route: Final = _normalize_public_auth_route(route)
|
|
if normalized_route == "/metrics":
|
|
return litellm.require_auth_for_metrics_endpoint is not False
|
|
|
|
return False
|
|
|
|
|
|
custom_litellm_key_header: Final = APIKeyHeader(
|
|
name=SpecialHeaders.custom_litellm_api_key.value,
|
|
auto_error=False,
|
|
description="Bearer token",
|
|
)
|
|
api_key_header: Final = APIKeyHeader(
|
|
name=SpecialHeaders.openai_authorization.value,
|
|
auto_error=False,
|
|
description="Bearer token",
|
|
)
|
|
azure_api_key_header: Final = APIKeyHeader(
|
|
name=SpecialHeaders.azure_authorization.value,
|
|
auto_error=False,
|
|
description="Some older versions of the openai Python package will send an API-Key header with just the API key ",
|
|
)
|
|
anthropic_api_key_header: Final = APIKeyHeader(
|
|
name=SpecialHeaders.anthropic_authorization.value,
|
|
auto_error=False,
|
|
description="If anthropic client used.",
|
|
)
|
|
google_ai_studio_api_key_header: Final = APIKeyHeader(
|
|
name=SpecialHeaders.google_ai_studio_authorization.value,
|
|
auto_error=False,
|
|
description="If google ai studio client used.",
|
|
)
|
|
azure_apim_header: Final = APIKeyHeader(
|
|
name=SpecialHeaders.azure_apim_authorization.value,
|
|
auto_error=False,
|
|
description="The default name of the subscription key header of Azure",
|
|
)
|
|
|
|
|
|
def _get_model_from_request_context(
|
|
request_data: dict,
|
|
route: str,
|
|
request: Request | None,
|
|
llm_router: Any | None = None,
|
|
team_id: str | None = None,
|
|
) -> str | list[str] | None:
|
|
return get_model_from_request(
|
|
request_data=request_data,
|
|
route=route,
|
|
request_headers=_safe_get_request_headers(request=request),
|
|
request_query_params=_safe_get_request_query_params(request=request),
|
|
llm_router=llm_router,
|
|
request=request,
|
|
team_id=team_id,
|
|
)
|
|
|
|
|
|
_CLAUDE_MODEL_ROUTES: Final = frozenset(
|
|
f"/{prefix}{endpoint}" for prefix in ("", "v1/") for endpoint in ("messages", "chat/completions", "responses")
|
|
)
|
|
_CLAUDE_MODEL_NORMALIZED: Final = "litellm.claude_model_normalized"
|
|
|
|
|
|
async def _normalize_claude_model(
|
|
request_data: dict, valid_token: UserAPIKeyAuth, request: Request | None, route: str
|
|
) -> None:
|
|
from litellm.proxy.proxy_server import llm_router, prisma_client, proxy_config, proxy_logging_obj
|
|
|
|
if route not in _CLAUDE_MODEL_ROUTES or llm_router is None:
|
|
return
|
|
if request is not None and request.scope.get(_CLAUDE_MODEL_NORMALIZED) is True:
|
|
return
|
|
requested: Final = _get_model_from_request_context(request_data, route, request, llm_router, valid_token.team_id)
|
|
if not isinstance(requested, str) or requested != request_data.get("model"):
|
|
return
|
|
if not requested.startswith("claude-router-") and not requested.lower().endswith("[1m]"):
|
|
return
|
|
settings: Final = await proxy_config.get_hierarchical_router_settings(
|
|
user_api_key_dict=valid_token, prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj
|
|
)
|
|
aliases: Final = settings.get("model_group_alias") if isinstance(settings, Mapping) else None
|
|
source: Final = claude_code_requested_group(
|
|
requested, llm_router, valid_token.team_id, (valid_token.aliases, valid_token.team_model_aliases, aliases)
|
|
)
|
|
if request is not None:
|
|
request.scope[_CLAUDE_MODEL_NORMALIZED] = True
|
|
if source is None:
|
|
return
|
|
request_data["model"] = source
|
|
_safe_set_request_parsed_body(request=request, parsed_body=request_data)
|
|
if request is not None:
|
|
request._json = request_data
|
|
request._body = orjson.dumps(request_data)
|
|
|
|
|
|
def _get_model_names_for_budget_checks(
|
|
model: str | list[str] | None,
|
|
) -> list[str]:
|
|
if model is None:
|
|
return []
|
|
if isinstance(model, str):
|
|
return [model]
|
|
return model
|
|
|
|
|
|
class _KeyModelBudgetLimiter(Protocol):
|
|
async def is_key_within_model_budget(self, user_api_key_dict: UserAPIKeyAuth, model: str) -> bool: ...
|
|
|
|
async def get_fallback_model_within_budget(self, user_api_key_dict: UserAPIKeyAuth, model: str) -> str | None: ...
|
|
|
|
|
|
class _UserModelBudgetLimiter(Protocol):
|
|
async def is_user_within_model_budget(
|
|
self, user_id: str, user_model_max_budget: Mapping[str, object], model: str
|
|
) -> bool: ...
|
|
|
|
|
|
class _TokenTeamModels(Protocol):
|
|
@property
|
|
def team_models(self) -> list[str]: ...
|
|
|
|
|
|
def _token_team_models(valid_token: _TokenTeamModels) -> list[str]:
|
|
return valid_token.team_models
|
|
|
|
|
|
async def _read_user_model_max_budget(
|
|
user_id: str | None,
|
|
prisma_client: PrismaClient | None,
|
|
user_api_key_cache: UserApiKeyCache,
|
|
parent_otel_span: Span | None,
|
|
proxy_logging_obj: ProxyLogging,
|
|
) -> Mapping[str, object] | None:
|
|
"""The user row's `model_max_budget`, or None when the row cannot be read.
|
|
|
|
A user whose row is missing must not be refused: this is a budget lookup,
|
|
and the main auth path likewise treats an unreadable user as no user.
|
|
"""
|
|
if user_id is None or prisma_client is None:
|
|
return None
|
|
try:
|
|
user_obj: Final = await get_user_object(
|
|
user_id=user_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
user_id_upsert=False,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
except Exception as e: # noqa: BLE001 # mirrors the main path's tolerance
|
|
verbose_logger.debug("Unable to read user for the per-model budget check: %s", e)
|
|
return None
|
|
return user_obj.model_max_budget if user_obj is not None else None
|
|
|
|
|
|
async def _check_user_model_budget(
|
|
valid_token: UserAPIKeyAuth,
|
|
model_max_budget_limiter: _UserModelBudgetLimiter,
|
|
models: list[str],
|
|
) -> None:
|
|
"""Enforce the internal user's own `model_max_budget` across the request's models.
|
|
|
|
Separate from the key check: a user's per-model budget caps every key they
|
|
own, so a caller cannot escape it by minting another key.
|
|
"""
|
|
user_model_max_budget: Final = valid_token.user_model_max_budget
|
|
if valid_token.user_id is None or not isinstance(user_model_max_budget, Mapping) or not user_model_max_budget:
|
|
return
|
|
for model_name in models:
|
|
await model_max_budget_limiter.is_user_within_model_budget(
|
|
user_id=valid_token.user_id,
|
|
user_model_max_budget=user_model_max_budget,
|
|
model=model_name,
|
|
)
|
|
|
|
|
|
async def _check_key_model_budget_with_fallback(
|
|
valid_token: UserAPIKeyAuth,
|
|
model_max_budget_limiter: _KeyModelBudgetLimiter,
|
|
model_name: str,
|
|
request_data: dict,
|
|
request: Request,
|
|
llm_model_list: list | None = None,
|
|
llm_router: litellm.Router | None = None,
|
|
) -> None:
|
|
"""
|
|
Enforce the key's per-model budget for `model_name`. If exceeded and the
|
|
key has a `budget_fallbacks` chain configured for `model_name`, reroute
|
|
the request to the first fallback model still within its own budget
|
|
instead of rejecting the request.
|
|
|
|
The selected fallback is validated against the key's model-access
|
|
allowlist and the team's model restrictions so that budget_fallbacks
|
|
cannot bypass model authorization. The rewrite is persisted to the
|
|
parsed-body cache, Starlette's JSON cache (``request._json``), and
|
|
path parameters so that downstream handlers see the final model
|
|
regardless of whether they consume ``_read_request_body()``,
|
|
``request.json()``, or the path ``model`` parameter.
|
|
|
|
Fallback is only attempted when ``model_name`` matches the top-level
|
|
``request_data["model"]``; models extracted from nested fields
|
|
(``session.model``, ``completion.model``, etc.) are not rewritable
|
|
and raise immediately.
|
|
|
|
Raises:
|
|
BudgetExceededError: if `model_name` is over budget and no configured
|
|
fallback is within budget either (or the fallback is not authorized).
|
|
"""
|
|
try:
|
|
await model_max_budget_limiter.is_key_within_model_budget(
|
|
user_api_key_dict=valid_token,
|
|
model=model_name,
|
|
)
|
|
except litellm.BudgetExceededError as e:
|
|
if request_data.get("model") != model_name:
|
|
raise e
|
|
fallback_model: Final = await model_max_budget_limiter.get_fallback_model_within_budget(
|
|
user_api_key_dict=valid_token,
|
|
model=model_name,
|
|
)
|
|
if fallback_model is None:
|
|
raise e
|
|
try:
|
|
await can_key_call_model(
|
|
model=fallback_model,
|
|
llm_model_list=llm_model_list,
|
|
valid_token=valid_token,
|
|
llm_router=llm_router,
|
|
)
|
|
if valid_token.team_models:
|
|
_can_object_call_model(
|
|
model=fallback_model,
|
|
llm_router=llm_router,
|
|
models=valid_token.team_models,
|
|
team_model_aliases=valid_token.team_model_aliases,
|
|
team_id=valid_token.team_id,
|
|
object_type="team",
|
|
)
|
|
except ProxyException:
|
|
raise e
|
|
request_data["model"] = fallback_model
|
|
_safe_set_request_parsed_body(request=request, parsed_body=request_data)
|
|
request._json = request_data
|
|
request._body = orjson.dumps(request_data)
|
|
path_params: Final = request.scope.get("path_params")
|
|
if isinstance(path_params, dict) and "model" in path_params:
|
|
path_params["model"] = fallback_model
|
|
|
|
|
|
def _get_bearer_token_or_received_api_key(api_key: str) -> str:
|
|
if api_key.startswith("Bearer "): # ensure Bearer token passed in
|
|
api_key = api_key.replace("Bearer ", "") # extract the token
|
|
elif api_key.startswith("Basic "):
|
|
api_key = api_key.replace("Basic ", "") # handle langfuse input
|
|
elif api_key.startswith("bearer "):
|
|
api_key = api_key.replace("bearer ", "")
|
|
elif api_key.startswith("AWS4-HMAC-SHA256"):
|
|
# Handle AWS Signature V4 format from LangChain
|
|
# Format: AWS4-HMAC-SHA256 Credential=Bearer sk-12345/date/region/service/aws4_request, SignedHeaders=..., Signature=...
|
|
# Extract the Bearer token from the Credential field
|
|
match = re.search(r"Credential=Bearer\s+([^/\s,]+)", api_key)
|
|
if match:
|
|
api_key = match.group(1)
|
|
else:
|
|
# If no Bearer token found in Credential, try to extract just the credential value
|
|
match = re.search(r"Credential=([^/\s,]+)", api_key)
|
|
if match:
|
|
api_key = match.group(1)
|
|
|
|
return api_key
|
|
|
|
|
|
def _routing_selector_matches_claim(
|
|
selector_value: Any | None,
|
|
claim_value: Any | None,
|
|
*,
|
|
split_space_delimited: bool = False,
|
|
) -> bool:
|
|
if selector_value is None:
|
|
return True
|
|
|
|
selector_list: Final[list[str]] = (
|
|
[str(v) for v in selector_value] if isinstance(selector_value, list) else [str(selector_value)]
|
|
)
|
|
|
|
if claim_value is None:
|
|
return False
|
|
|
|
if isinstance(claim_value, list):
|
|
claim_list = [str(v) for v in claim_value]
|
|
elif split_space_delimited and isinstance(claim_value, str) and " " in claim_value.strip():
|
|
# OAuth/OIDC often sends scope as a single space-delimited string. Only split
|
|
# for the scope selector: iss/aud/client_id must stay exact full-string match
|
|
# on unverified claims (see routing override security review). The elif guard
|
|
# (`" " in claim_value.strip()`) ensures at least two non-empty tokens survive.
|
|
claim_list = [v for v in claim_value.strip().split(" ") if v]
|
|
else:
|
|
claim_list = [str(claim_value)]
|
|
|
|
def _selector_matches_claim(selector: str, claim: str) -> bool:
|
|
# NOTE: wildcard matching is case-sensitive (fnmatch.fnmatchcase).
|
|
if "*" in selector or "?" in selector:
|
|
# Without scope splitting, do not let `*` span whitespace: a malformed
|
|
# iss like "trusted.example.com evil.com" must not match "trusted.*".
|
|
# Scope uses split_space_delimited so each claim token is checked separately.
|
|
if not split_space_delimited and any(ch.isspace() for ch in claim):
|
|
return False
|
|
return fnmatch.fnmatchcase(claim, selector)
|
|
return selector == claim
|
|
|
|
return any(_selector_matches_claim(selector=s, claim=c) for s in selector_list for c in claim_list)
|
|
|
|
|
|
def _matches_routing_override(token_claims: dict, override: "JWTRoutingOverride") -> bool:
|
|
return (
|
|
_routing_selector_matches_claim(override.iss, token_claims.get("iss"))
|
|
and _routing_selector_matches_claim(override.client_id, token_claims.get("client_id"))
|
|
and _routing_selector_matches_claim(
|
|
override.scope,
|
|
token_claims.get("scope"),
|
|
split_space_delimited=True,
|
|
)
|
|
and _routing_selector_matches_claim(override.aud, token_claims.get("aud"))
|
|
)
|
|
|
|
|
|
def _should_route_jwt_to_oauth2_override(token: str, jwt_handler: JWTHandler) -> bool:
|
|
routing_overrides: Final = jwt_handler.litellm_jwtauth.routing_overrides
|
|
if not routing_overrides:
|
|
return False
|
|
|
|
token_claims: Final = jwt_handler.get_unverified_claims(token=token)
|
|
if token_claims is None:
|
|
return False
|
|
|
|
for override in routing_overrides:
|
|
if override.path == "oauth2" and _matches_routing_override(token_claims=token_claims, override=override):
|
|
verbose_proxy_logger.debug("JWT routing override matched. Routing token to OAuth2 introspection.")
|
|
return True
|
|
|
|
return False
|
|
|
|
|
|
def _get_bearer_token(
|
|
api_key: str,
|
|
):
|
|
if api_key.startswith("Bearer "): # ensure Bearer token passed in
|
|
api_key = api_key.replace("Bearer ", "") # extract the token
|
|
elif api_key.startswith("Basic "):
|
|
api_key = api_key.replace("Basic ", "") # handle langfuse input
|
|
elif api_key.startswith("bearer "):
|
|
api_key = api_key.replace("bearer ", "")
|
|
elif api_key.startswith("AWS4-HMAC-SHA256"):
|
|
# Handle AWS Signature V4 format from LangChain
|
|
# Format: AWS4-HMAC-SHA256 Credential=Bearer sk-12345/date/region/service/aws4_request, SignedHeaders=..., Signature=...
|
|
# Extract the Bearer token from the Credential field
|
|
match = re.search(r"Credential=Bearer\s+([^/\s,]+)", api_key)
|
|
if match:
|
|
api_key = match.group(1)
|
|
else:
|
|
# If no Bearer token found in Credential, try to extract just the credential value
|
|
match = re.search(r"Credential=([^/\s,]+)", api_key)
|
|
if match:
|
|
api_key = match.group(1)
|
|
else:
|
|
api_key = ""
|
|
else:
|
|
api_key = ""
|
|
return api_key
|
|
|
|
|
|
def _apply_budget_limits_to_end_user_params(
|
|
end_user_params: dict,
|
|
budget_info: LiteLLM_BudgetTable,
|
|
end_user_id: str | None,
|
|
) -> None:
|
|
"""
|
|
Helper function to apply budget limits to end user parameters.
|
|
|
|
Args:
|
|
end_user_params: Dictionary to update with budget parameters
|
|
budget_info: Budget table object containing limits
|
|
end_user_id: ID of the end user for logging
|
|
"""
|
|
if budget_info.tpm_limit is not None:
|
|
end_user_params["end_user_tpm_limit"] = budget_info.tpm_limit
|
|
|
|
if budget_info.rpm_limit is not None:
|
|
end_user_params["end_user_rpm_limit"] = budget_info.rpm_limit
|
|
|
|
if budget_info.max_budget is not None:
|
|
end_user_params["end_user_max_budget"] = budget_info.max_budget
|
|
|
|
if budget_info.model_max_budget is not None:
|
|
end_user_params["end_user_model_max_budget"] = budget_info.model_max_budget
|
|
|
|
verbose_proxy_logger.debug("Applied budget limits to end user %s", end_user_id)
|
|
|
|
|
|
async def user_api_key_auth_websocket(websocket: WebSocket):
|
|
# Accept the WebSocket connection
|
|
|
|
ws_scope: Final = websocket.scope or {}
|
|
scope_headers: Final = list(ws_scope.get("headers") or [])
|
|
# ``get_request_route`` falls back to ``request.url.path`` when
|
|
# ``scope["path"]`` is absent. On WebSockets that fallback reads
|
|
# ``websocket.url``, which Starlette reconstructs from the (poisonable)
|
|
# Host header. Carry the ASGI scope's path / root_path so the lookup
|
|
# never reaches the fallback.
|
|
synthetic_scope: Final[dict[str, Any]] = {
|
|
"type": "http",
|
|
"headers": scope_headers,
|
|
"path": ws_scope.get("path", ""),
|
|
}
|
|
for key in ("root_path", "app_root_path"):
|
|
if key in ws_scope:
|
|
synthetic_scope[key] = ws_scope[key]
|
|
request: Final = Request(scope=synthetic_scope)
|
|
|
|
request._url = websocket.url
|
|
|
|
query_params: Final = websocket.query_params
|
|
|
|
model: Final = query_params.get("model")
|
|
|
|
async def return_body():
|
|
return _realtime_request_body(model)
|
|
|
|
request.body = return_body
|
|
|
|
authorization: Final = websocket.headers.get("authorization")
|
|
# If no Authorization header, try the api-key header
|
|
if not authorization:
|
|
api_key = websocket.headers.get("api-key")
|
|
if not api_key:
|
|
# Try extracting from WebSocket subprotocol (browser clients)
|
|
for protocol in websocket.headers.get("sec-websocket-protocol", "").split(","):
|
|
protocol = protocol.strip()
|
|
if protocol.startswith("openai-insecure-api-key."):
|
|
api_key = protocol[len("openai-insecure-api-key.") :]
|
|
break
|
|
if not api_key:
|
|
await websocket.close(code=status.WS_1008_POLICY_VIOLATION)
|
|
raise HTTPException(status_code=403, detail="No API key provided")
|
|
else:
|
|
# Extract the API key from the Bearer token
|
|
if not authorization.startswith("Bearer "):
|
|
await websocket.close(code=status.WS_1008_POLICY_VIOLATION)
|
|
raise HTTPException(status_code=403, detail="Invalid Authorization header format")
|
|
|
|
api_key = authorization[len("Bearer ") :].strip()
|
|
|
|
# Call user_api_key_auth with the extracted API key
|
|
# Note: You'll need to modify this to work with WebSocket context if needed
|
|
try:
|
|
return await user_api_key_auth(request=request, api_key=f"Bearer {api_key}")
|
|
except Exception as e:
|
|
if is_invalid_virtual_key_error(e):
|
|
raise WebSocketException(code=status.WS_1008_POLICY_VIOLATION)
|
|
verbose_proxy_logger.exception(e)
|
|
await websocket.close(code=status.WS_1008_POLICY_VIOLATION)
|
|
raise HTTPException(status_code=403, detail=str(e))
|
|
|
|
|
|
def update_valid_token_with_end_user_params(valid_token: UserAPIKeyAuth, end_user_params: dict) -> UserAPIKeyAuth:
|
|
valid_token.end_user_id = end_user_params.get("end_user_id")
|
|
# Only overwrite token fields when the DB-derived value is not None.
|
|
# This prevents DB lookups (where the budget table has no value set)
|
|
# from silently clearing values that a custom auth function may have
|
|
# already set on the token.
|
|
if end_user_params.get("end_user_tpm_limit") is not None:
|
|
valid_token.end_user_tpm_limit = end_user_params["end_user_tpm_limit"]
|
|
if end_user_params.get("end_user_rpm_limit") is not None:
|
|
valid_token.end_user_rpm_limit = end_user_params["end_user_rpm_limit"]
|
|
if end_user_params.get("allowed_model_region") is not None:
|
|
valid_token.allowed_model_region = end_user_params["allowed_model_region"]
|
|
if end_user_params.get("end_user_model_max_budget") is not None:
|
|
valid_token.end_user_model_max_budget = end_user_params["end_user_model_max_budget"]
|
|
return valid_token
|
|
|
|
|
|
# Reusable coordinator for global spend to prevent cache stampede
|
|
_global_spend_coordinator: Final = EventDrivenCacheCoordinator(log_prefix="[GLOBAL SPEND]")
|
|
|
|
|
|
async def _fetch_global_spend_with_event_coordination(
|
|
cache_key: str,
|
|
user_api_key_cache: UserApiKeyCache,
|
|
prisma_client: PrismaClient,
|
|
) -> float | None:
|
|
"""
|
|
Fetch global spend with event-driven coordination to prevent cache stampede.
|
|
Uses EventDrivenCacheCoordinator: first request queries DB and signals others when done.
|
|
|
|
Reads the proxy budget aggregate user row, which accrues proxy-wide spend
|
|
per request and is zeroed by ResetBudgetJob every ``litellm.budget_duration``.
|
|
"""
|
|
|
|
async def _load_global_spend() -> float | None:
|
|
proxy_budget_row: Final = await prisma_client.db.litellm_usertable.find_unique(
|
|
where={"user_id": LITELLM_PROXY_BUDGET_NAME}
|
|
)
|
|
return float(proxy_budget_row.spend) if proxy_budget_row is not None else None
|
|
|
|
return await _global_spend_coordinator.get_or_load(
|
|
cache_key=cache_key,
|
|
cache=user_api_key_cache, # pyright: ignore[reportArgumentType]
|
|
load_fn=_load_global_spend,
|
|
)
|
|
|
|
|
|
async def get_global_proxy_spend(
|
|
litellm_proxy_admin_name: str,
|
|
user_api_key_cache: UserApiKeyCache,
|
|
prisma_client: PrismaClient | None,
|
|
token: str,
|
|
proxy_logging_obj: ProxyLogging,
|
|
) -> float | None:
|
|
global_proxy_spend = None
|
|
if litellm.max_budget > 0 and prisma_client is not None: # user set proxy max budget
|
|
# Use event-driven coordination to prevent cache stampede
|
|
cache_key: Final = GLOBAL_PROXY_SPEND_CACHE_KEY
|
|
global_proxy_spend = await _fetch_global_spend_with_event_coordination(
|
|
cache_key=cache_key,
|
|
user_api_key_cache=user_api_key_cache,
|
|
prisma_client=prisma_client,
|
|
)
|
|
if global_proxy_spend is not None:
|
|
user_info: Final = CallInfo(
|
|
user_id=litellm_proxy_admin_name,
|
|
max_budget=litellm.max_budget,
|
|
spend=global_proxy_spend,
|
|
token=token,
|
|
event_group=Litellm_EntityType.PROXY,
|
|
)
|
|
asyncio.create_task(
|
|
proxy_logging_obj.budget_alerts(
|
|
type="proxy_budget",
|
|
user_info=user_info,
|
|
)
|
|
)
|
|
return global_proxy_spend
|
|
|
|
|
|
def get_rbac_role(jwt_handler: JWTHandler, scopes: list[str]) -> str:
|
|
is_admin: Final = jwt_handler.is_admin(scopes=scopes)
|
|
if is_admin:
|
|
return LitellmUserRoles.PROXY_ADMIN
|
|
else:
|
|
return LitellmUserRoles.TEAM
|
|
|
|
|
|
def get_api_key(
|
|
custom_litellm_key_header: str | None,
|
|
api_key: str,
|
|
azure_api_key_header: str | None,
|
|
anthropic_api_key_header: str | None,
|
|
google_ai_studio_api_key_header: str | None,
|
|
azure_apim_header: str | None,
|
|
pass_through_endpoints: list[dict] | None,
|
|
route: str,
|
|
request: Request,
|
|
) -> tuple[str, str | None]:
|
|
"""
|
|
Returns:
|
|
Tuple[Optional[str], Optional[str]]: Tuple of the api_key and the passed_in_key
|
|
"""
|
|
from litellm.proxy.auth.route_checks import RouteChecks
|
|
from litellm.proxy.common_utils.http_parsing_utils import (
|
|
_safe_get_request_query_params,
|
|
)
|
|
|
|
api_key = api_key
|
|
passed_in_key: str | None = None
|
|
if isinstance(custom_litellm_key_header, str):
|
|
passed_in_key = custom_litellm_key_header
|
|
api_key = _get_bearer_token_or_received_api_key(custom_litellm_key_header)
|
|
elif isinstance(api_key, str) and len(api_key) > 0:
|
|
passed_in_key = api_key
|
|
api_key = _get_bearer_token(api_key=api_key)
|
|
elif isinstance(azure_api_key_header, str):
|
|
passed_in_key = azure_api_key_header
|
|
api_key = azure_api_key_header
|
|
elif isinstance(anthropic_api_key_header, str):
|
|
passed_in_key = anthropic_api_key_header
|
|
api_key = anthropic_api_key_header
|
|
elif isinstance(google_ai_studio_api_key_header, str):
|
|
passed_in_key = google_ai_studio_api_key_header
|
|
api_key = google_ai_studio_api_key_header
|
|
elif isinstance(azure_apim_header, str):
|
|
passed_in_key = azure_apim_header
|
|
api_key = azure_apim_header
|
|
elif (
|
|
RouteChecks.is_generate_content_route(route=route)
|
|
and request is not None
|
|
and _safe_get_request_query_params(request).get("key")
|
|
):
|
|
google_auth_key: Final[str] = _safe_get_request_query_params(request).get("key") or ""
|
|
passed_in_key = google_auth_key
|
|
api_key = google_auth_key
|
|
elif pass_through_endpoints is not None:
|
|
for endpoint in pass_through_endpoints:
|
|
if endpoint.get("path", "") == route:
|
|
headers: dict | None = endpoint.get("headers", None)
|
|
if headers is not None:
|
|
header_key: str = headers.get("litellm_user_api_key", "")
|
|
if request.headers.get(header_key) is not None:
|
|
api_key = request.headers.get(header_key) or ""
|
|
passed_in_key = api_key
|
|
return api_key, passed_in_key
|
|
|
|
|
|
async def check_api_key_for_custom_headers_or_pass_through_endpoints(
|
|
request: Request,
|
|
route: str,
|
|
pass_through_endpoints: list[dict] | None,
|
|
api_key: str,
|
|
) -> UserAPIKeyAuth | str:
|
|
is_mapped_pass_through_route: bool = False
|
|
normalized_route: Final = normalize_route_for_root_path(route)
|
|
if normalized_route is not None:
|
|
for mapped_route in LiteLLMRoutes.mapped_pass_through_routes.value:
|
|
if normalized_route.startswith(mapped_route):
|
|
is_mapped_pass_through_route = True
|
|
break
|
|
if is_mapped_pass_through_route:
|
|
if request.headers.get("litellm_user_api_key") is not None:
|
|
api_key = request.headers.get("litellm_user_api_key") or ""
|
|
if pass_through_endpoints is not None:
|
|
for endpoint in pass_through_endpoints:
|
|
if isinstance(endpoint, dict) and endpoint.get("path", "") == route:
|
|
## IF AUTH DISABLED
|
|
# Default to True: a config dict with no ``auth`` key
|
|
# otherwise produced an unauthenticated forwarder. The
|
|
# Pydantic ``PassThroughGenericEndpoint.auth`` default
|
|
# is also True, but raw config dicts skip that path —
|
|
# so this runtime check has to default to True too.
|
|
if endpoint.get("auth", True) is not True:
|
|
return UserAPIKeyAuth()
|
|
## IF AUTH ENABLED
|
|
### IF CUSTOM PARSER REQUIRED
|
|
if endpoint.get("custom_auth_parser") is not None and endpoint.get("custom_auth_parser") == "langfuse":
|
|
# langfuse returns {'Authorization': 'Basic <base64(username:password)>'}
|
|
# check the langfuse public key if it contains the litellm api key
|
|
import base64
|
|
|
|
api_key = api_key.replace("Basic ", "").strip()
|
|
decoded_bytes = base64.b64decode(api_key)
|
|
decoded_str = decoded_bytes.decode("utf-8")
|
|
api_key = decoded_str.split(":")[0]
|
|
else:
|
|
headers = endpoint.get("headers", None)
|
|
if headers is not None:
|
|
header_key = headers.get("litellm_user_api_key", "")
|
|
if isinstance(request.headers, dict) and request.headers.get(key=header_key) is not None:
|
|
api_key = request.headers.get(key=header_key)
|
|
return api_key
|
|
|
|
|
|
# Cache sentinel written when a JWT under AUTO_REGISTER resolved to a proxy
|
|
# admin via auth_builder. Proxy admins don't need a mapped virtual key (they
|
|
# have full access via auth_builder anyway), but without a cache entry every
|
|
# subsequent request from the same JWT identity would re-query the DB for a
|
|
# non-existent mapping. Sentinel tells _resolve_jwt_to_virtual_key to skip
|
|
# the lookup and return None (caller proceeds to auth_builder).
|
|
_JWT_PROXY_ADMIN_SENTINEL: Final = "__JWT_PROXY_ADMIN__"
|
|
|
|
_JWT_AUTH_DISABLED_HINT = (
|
|
" This key has the structure of a JWT, but JWT auth is not enabled on this proxy, so it was treated as a"
|
|
" virtual key. Set `enable_jwt_auth: true` under `general_settings` in your proxy config to authenticate"
|
|
" with JWTs."
|
|
)
|
|
|
|
|
|
class _PendingAutoRegister(NamedTuple):
|
|
"""
|
|
Signal returned by ``_resolve_jwt_to_virtual_key`` when the JWT's claim is
|
|
unmapped and ``unregistered_jwt_client_behavior`` is AUTO_REGISTER.
|
|
|
|
The caller MUST run standard ``JWTAuthManager.auth_builder`` to apply RBAC,
|
|
scope mappings, ``custom_validate``, and ``user_allowed_email_domain``
|
|
policy BEFORE calling ``_auto_register_jwt_mapping`` with the validated
|
|
``team_id`` / ``user_id`` from the auth_builder result. Auto-registering
|
|
purely on a signature-valid JWT (the old behavior) bypassed every JWT
|
|
policy beyond signature verification.
|
|
"""
|
|
|
|
claim_field: str
|
|
claim_value: str
|
|
cache_key: str
|
|
|
|
|
|
async def _auto_register_jwt_mapping(
|
|
virtual_key_claim_field: str,
|
|
claim_value: str,
|
|
jwt_handler: JWTHandler,
|
|
prisma_client: PrismaClient,
|
|
user_api_key_cache: UserApiKeyCache,
|
|
parent_otel_span: Span | None,
|
|
proxy_logging_obj: ProxyLogging,
|
|
cache_key: str,
|
|
team_id: str | None = None,
|
|
user_id: str | None = None,
|
|
org_id: str | None = None,
|
|
end_user_id: str | None = None,
|
|
) -> UserAPIKeyAuth | None:
|
|
"""
|
|
Auto-register: create a new virtual key + mapping for an unrecognised JWT
|
|
claim value. ``team_id`` and ``user_id`` must come from a successful
|
|
``JWTAuthManager.auth_builder`` run — they encode the JWT identity AFTER
|
|
RBAC/scope/custom_validate/email-domain policy has been enforced. The key
|
|
is stamped with those values so the cached future-request path inherits
|
|
the same team/user/org limits the auth_builder path would have applied.
|
|
|
|
Race safety: if two concurrent requests both reach here simultaneously (both
|
|
saw no mapping in the DB), one will win the unique-constraint race on
|
|
litellm_jwtkeymapping. The loser catches the conflict, deletes its orphaned
|
|
key, fetches the winner's mapping, and proceeds — no error surfaced.
|
|
"""
|
|
# Inline import required: key_management_endpoints imports user_api_key_auth
|
|
# (line 51) so a module-level import here would create a circular dependency.
|
|
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
|
generate_key_helper_fn,
|
|
)
|
|
|
|
# ``table_name="key"`` is required: without it, generate_key_helper_fn
|
|
# falls into the user-upsert branch (`table_name is None or "user"`) and
|
|
# attempts to insert into LiteLLM_UserTable with user_id=None, which fails
|
|
# the NOT NULL @id constraint. Every successful key-creation caller (e.g.
|
|
# /key/generate) passes table_name="key" explicitly.
|
|
key_data: Final = await generate_key_helper_fn(
|
|
llm_router=None,
|
|
request_type="key",
|
|
table_name="key",
|
|
team_id=team_id,
|
|
user_id=user_id,
|
|
organization_id=org_id,
|
|
metadata={
|
|
"auto_registered": True,
|
|
"jwt_claim_field": virtual_key_claim_field,
|
|
"jwt_claim_value": claim_value,
|
|
},
|
|
)
|
|
# generate_key_helper_fn returns the plaintext key in "token"; the persisted
|
|
# row in LiteLLM_VerificationToken uses its hash, so hash here to get the FK
|
|
# value referenced by LiteLLM_JWTKeyMapping.token.
|
|
token_hash = hash_token(key_data["token"])
|
|
|
|
try:
|
|
await prisma_client.db.litellm_jwtkeymapping.create(
|
|
data={
|
|
"jwt_claim_name": virtual_key_claim_field,
|
|
"jwt_claim_value": claim_value,
|
|
"token": token_hash,
|
|
"created_by": "auto_register",
|
|
"updated_by": "auto_register",
|
|
}
|
|
)
|
|
except Exception as e:
|
|
error_str: Final = str(e).lower()
|
|
if "unique" in error_str or "p2002" in error_str:
|
|
# A concurrent request won the race. The key generate_key_helper_fn
|
|
# just persisted to LiteLLM_VerificationToken is orphaned — nothing
|
|
# maps to it, but it's a fully valid unrestricted API key sitting in
|
|
# the DB and the cleartext is in memory on this request. Delete it
|
|
# so orphans don't accumulate under sustained concurrency.
|
|
verbose_proxy_logger.debug(
|
|
"JWT Key Mapping (auto_register): unique conflict on create — "
|
|
"deleting orphaned virtual key and fetching winner's mapping for %s='%s'.",
|
|
virtual_key_claim_field,
|
|
claim_value,
|
|
)
|
|
try:
|
|
await prisma_client.db.litellm_verificationtoken.delete(where={"token": token_hash})
|
|
except Exception as delete_err:
|
|
# Don't fail the request if cleanup fails — the orphan is
|
|
# unmapped and inert. Log so an operator can prune it later.
|
|
verbose_proxy_logger.warning(
|
|
"JWT Key Mapping (auto_register): failed to delete orphaned key after race: %s",
|
|
delete_err,
|
|
)
|
|
token_hash = await get_jwt_key_mapping_object(
|
|
jwt_claim_name=virtual_key_claim_field,
|
|
jwt_claim_value=claim_value,
|
|
prisma_client=prisma_client,
|
|
)
|
|
if token_hash is None:
|
|
# The winner's mapping vanished between the unique-constraint
|
|
# conflict and our re-fetch (concurrent delete). Returning None
|
|
# here would silently fall through to team-based JWT auth —
|
|
# a less-restrictive path than the operator configured. Raise
|
|
# 503 so the caller retries against a stable state instead.
|
|
raise HTTPException(
|
|
status_code=503,
|
|
detail=(
|
|
"JWT Key Mapping: AUTO_REGISTER race resolution failed — "
|
|
"winner's mapping was concurrently removed. Retry the request."
|
|
),
|
|
)
|
|
else:
|
|
raise
|
|
|
|
await user_api_key_cache.async_set_cache(
|
|
key=cache_key,
|
|
value=token_hash,
|
|
ttl=jwt_handler.litellm_jwtauth.virtual_key_mapping_cache_ttl,
|
|
)
|
|
|
|
verbose_proxy_logger.info(
|
|
"JWT Key Mapping (auto_register): created new virtual key for %s='%s'.",
|
|
virtual_key_claim_field,
|
|
claim_value,
|
|
)
|
|
|
|
auto_registered_key: Final = IdentityStore.key_from_principal(
|
|
await IdentityStore(
|
|
prisma_client,
|
|
user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
).resolve(hashed_token=token_hash)
|
|
)
|
|
if auto_registered_key is not None:
|
|
auto_registered_key.org_id = org_id
|
|
auto_registered_key.end_user_id = end_user_id
|
|
auto_registered_key.api_key = auto_registered_key.token
|
|
return auto_registered_key
|
|
|
|
|
|
async def _resolve_jwt_to_virtual_key(
|
|
jwt_claims: dict,
|
|
jwt_handler: JWTHandler,
|
|
prisma_client: PrismaClient | None,
|
|
user_api_key_cache: UserApiKeyCache,
|
|
parent_otel_span: Span | None,
|
|
proxy_logging_obj: ProxyLogging,
|
|
) -> Union[UserAPIKeyAuth | None, "_PendingAutoRegister"]:
|
|
"""
|
|
Returns:
|
|
- ``UserAPIKeyAuth``: a resolved virtual key (cache hit or DB hit). The
|
|
caller may use this directly; JWT policy has been enforced previously
|
|
(at key-creation time or, for cached results, before caching).
|
|
- ``_PendingAutoRegister``: claim is unmapped and behavior is AUTO_REGISTER.
|
|
The caller MUST run ``JWTAuthManager.auth_builder`` to enforce JWT
|
|
policy (RBAC, scope, custom_validate, email-domain), then invoke
|
|
``_auto_register_jwt_mapping`` with the validated team_id/user_id.
|
|
- ``None``: claim is unmapped and behavior is FALLBACK_TEAM_MAPPING.
|
|
The caller falls through to standard team-based JWT auth (which itself
|
|
enforces full JWT policy via auth_builder).
|
|
- Raises HTTPException: REJECT policy hit, missing claim under
|
|
REJECT/AUTO_REGISTER, or other policy violations.
|
|
"""
|
|
raw_issuer: Final = jwt_claims.get(JWTHandler.LITELLM_JWT_ISSUER_CLAIM)
|
|
normalized_issuer: Final = raw_issuer if isinstance(raw_issuer, str) else None
|
|
virtual_key_claim_field: Final = jwt_handler.litellm_jwtauth.get_virtual_key_claim_field(normalized_issuer)
|
|
if virtual_key_claim_field is None:
|
|
return None
|
|
behavior: Final = jwt_handler.litellm_jwtauth.get_unregistered_jwt_client_behavior(normalized_issuer)
|
|
|
|
claim_value: Final = get_nested_value(
|
|
data=jwt_claims,
|
|
key_path=virtual_key_claim_field,
|
|
default=None,
|
|
)
|
|
|
|
if claim_value is None:
|
|
verbose_proxy_logger.debug(
|
|
"JWT Key Mapping: Claim field '%s' not found in JWT claims.", virtual_key_claim_field
|
|
)
|
|
# A missing claim is an unmapped client — apply the no-match policy
|
|
# rather than returning early. Otherwise a caller can bypass REJECT
|
|
# simply by presenting a JWT that omits the configured field. For
|
|
# AUTO_REGISTER there is no stable identity to map without a claim
|
|
# value, so we deny rather than create a sentinel-keyed record.
|
|
if behavior in (
|
|
UnregisteredJWTClientBehavior.REJECT,
|
|
UnregisteredJWTClientBehavior.AUTO_REGISTER,
|
|
):
|
|
raise HTTPException(
|
|
status_code=403,
|
|
detail=(
|
|
f"JWT Key Mapping: Required claim '{virtual_key_claim_field}' "
|
|
"is missing from the JWT. Access denied."
|
|
),
|
|
)
|
|
return None
|
|
|
|
cache_key: Final = jwt_key_mapping_cache_key(virtual_key_claim_field, str(claim_value))
|
|
raw_cached_mapping: Final = await user_api_key_cache.async_get_cache(cache_key)
|
|
sentinel_written_by_this_policy: Final = behavior == UnregisteredJWTClientBehavior.AUTO_REGISTER
|
|
cached_mapping: Final = (
|
|
None
|
|
if raw_cached_mapping == _JWT_PROXY_ADMIN_SENTINEL and not sentinel_written_by_this_policy
|
|
else raw_cached_mapping
|
|
)
|
|
|
|
if cached_mapping == _JWT_PROXY_ADMIN_SENTINEL:
|
|
# Previously resolved to a proxy admin via auth_builder; skip the
|
|
# mapping lookup and let the caller re-run auth_builder. Avoids a
|
|
# repeated DB hit on every proxy-admin request under AUTO_REGISTER.
|
|
return None
|
|
|
|
if cached_mapping == "__NO_MAPPING__":
|
|
if behavior == UnregisteredJWTClientBehavior.REJECT:
|
|
raise HTTPException(
|
|
status_code=403,
|
|
detail=f"JWT Key Mapping: No registered mapping for {virtual_key_claim_field}='{claim_value}'. Access denied.",
|
|
)
|
|
if behavior == UnregisteredJWTClientBehavior.AUTO_REGISTER:
|
|
# Stale sentinel written under a prior fallback_team_mapping config —
|
|
# evict it and defer auto-register to after auth_builder runs. Raise
|
|
# the same 500 as the fresh-path AUTO_REGISTER branch when there is
|
|
# no DB, so behavior is consistent regardless of whether the cache
|
|
# happens to hold the sentinel.
|
|
if prisma_client is None:
|
|
raise HTTPException(
|
|
status_code=500,
|
|
detail=(
|
|
"JWT Key Mapping: AUTO_REGISTER requires a database connection. "
|
|
"Configure a database or change unregistered_jwt_client_behavior."
|
|
),
|
|
)
|
|
await user_api_key_cache.async_delete_cache(cache_key)
|
|
return _PendingAutoRegister(
|
|
claim_field=virtual_key_claim_field,
|
|
claim_value=str(claim_value),
|
|
cache_key=cache_key,
|
|
)
|
|
return None
|
|
elif cached_mapping is not None:
|
|
return IdentityStore.key_from_principal(
|
|
await IdentityStore(
|
|
prisma_client,
|
|
user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
).resolve(hashed_token=cached_mapping)
|
|
)
|
|
|
|
# Resolve the mapping from DB, or treat prisma_client=None as a definitive
|
|
# miss (no DB → no mapping can exist → apply no-match policy below).
|
|
token_hash: str | None = None
|
|
if prisma_client is not None:
|
|
token_hash = await get_jwt_key_mapping_object(
|
|
jwt_claim_name=virtual_key_claim_field,
|
|
jwt_claim_value=str(claim_value),
|
|
prisma_client=prisma_client,
|
|
)
|
|
|
|
if token_hash is not None:
|
|
await user_api_key_cache.async_set_cache(
|
|
key=cache_key,
|
|
value=token_hash,
|
|
ttl=jwt_handler.litellm_jwtauth.virtual_key_mapping_cache_ttl,
|
|
)
|
|
return IdentityStore.key_from_principal(
|
|
await IdentityStore(
|
|
prisma_client,
|
|
user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
).resolve(hashed_token=token_hash)
|
|
)
|
|
|
|
# No mapping found (DB miss or no DB) — apply no-match policy.
|
|
if behavior == UnregisteredJWTClientBehavior.REJECT:
|
|
# Cache the miss before raising so repeated rejections are served from
|
|
# cache and don't re-query the DB on every request.
|
|
await user_api_key_cache.async_set_cache(
|
|
key=cache_key,
|
|
value="__NO_MAPPING__",
|
|
ttl=jwt_handler.litellm_jwtauth.virtual_key_mapping_cache_ttl,
|
|
)
|
|
raise HTTPException(
|
|
status_code=403,
|
|
detail=f"JWT Key Mapping: No registered mapping for {virtual_key_claim_field}='{claim_value}'. Access denied.",
|
|
)
|
|
|
|
if behavior == UnregisteredJWTClientBehavior.AUTO_REGISTER:
|
|
if prisma_client is None:
|
|
raise HTTPException(
|
|
status_code=500,
|
|
detail=(
|
|
"JWT Key Mapping: AUTO_REGISTER requires a database connection. "
|
|
"Configure a database or change unregistered_jwt_client_behavior."
|
|
),
|
|
)
|
|
# Defer: caller runs JWTAuthManager.auth_builder to enforce RBAC, scope,
|
|
# custom_validate, and email-domain policy, then auto-registers using
|
|
# the validated identity. Auto-registering here on a signature-only
|
|
# JWT would bypass every JWT policy beyond signature verification.
|
|
return _PendingAutoRegister(
|
|
claim_field=virtual_key_claim_field,
|
|
claim_value=str(claim_value),
|
|
cache_key=cache_key,
|
|
)
|
|
|
|
# FALLBACK_TEAM_MAPPING (default): cache the miss and return None so the
|
|
# caller falls through to standard team-based JWT auth.
|
|
await user_api_key_cache.async_set_cache(
|
|
key=cache_key,
|
|
value="__NO_MAPPING__",
|
|
ttl=jwt_handler.litellm_jwtauth.virtual_key_mapping_cache_ttl,
|
|
)
|
|
return None
|
|
|
|
|
|
def _ensure_litellm_received_at_on_request_state(request: Request) -> datetime:
|
|
"""Idempotently stamp ``request.state.litellm_received_at`` with the moment
|
|
litellm's own code started handling this request -- the first line of
|
|
``user_api_key_auth``, before any auth/pre-call work runs. This is the
|
|
basis for the request-latency Prometheus metrics (see
|
|
``litellm/integrations/prometheus.py``), and unlike the OTEL SERVER span
|
|
below, it is set unconditionally so those metrics don't depend on OTEL
|
|
being configured.
|
|
"""
|
|
existing_received_at: Final[datetime | None] = getattr(request.state, "litellm_received_at", None)
|
|
if existing_received_at is not None:
|
|
return existing_received_at
|
|
received_at: Final = datetime.now(timezone.utc)
|
|
try:
|
|
request.state.litellm_received_at = received_at
|
|
except Exception:
|
|
pass
|
|
return received_at
|
|
|
|
|
|
def _ensure_parent_otel_span_on_request_state(request: Request) -> None:
|
|
"""Idempotently create the OTEL SERVER span and stash it on
|
|
``request.state.parent_otel_span``. Safe to call multiple times.
|
|
|
|
Called both at the top of ``user_api_key_auth`` (so body-parse failures
|
|
have a span to close) and inside ``_user_api_key_auth_builder`` (for
|
|
callers that bypass ``user_api_key_auth``, e.g. MCP).
|
|
"""
|
|
from litellm.proxy.proxy_server import open_telemetry_logger
|
|
|
|
start_time: Final = _ensure_litellm_received_at_on_request_state(request)
|
|
|
|
if open_telemetry_logger is None:
|
|
return
|
|
if getattr(request.state, "parent_otel_span", None) is not None:
|
|
return
|
|
parent_otel_span: Final = open_telemetry_logger.create_litellm_proxy_request_started_span(
|
|
start_time=start_time,
|
|
headers=_safe_get_request_headers(request),
|
|
)
|
|
# Under V2 the FastAPI instrumentor stamps http.route / url.path on the server
|
|
# span; only the legacy logger needs these set explicitly.
|
|
set_route_attrs: Final = getattr(open_telemetry_logger, "set_proxy_request_route_attributes", None)
|
|
if not is_otel_v2_enabled() and set_route_attrs is not None:
|
|
set_route_attrs(
|
|
parent_otel_span,
|
|
url_path=get_request_route(request=request),
|
|
http_route=get_request_route_template(request),
|
|
)
|
|
request.state.parent_otel_span = parent_otel_span
|
|
|
|
|
|
async def _read_request_body_deferring_parse_failure(
|
|
request: Request,
|
|
) -> tuple[dict, ProxyException | None]:
|
|
"""Parse the body, returning a parse failure instead of raising it.
|
|
|
|
A body that fails to parse is still a request from a known caller, so auth
|
|
must run (resolving identity onto the request's trace) before the 400 goes
|
|
out; the caller re-raises the returned exception once identity is seeded.
|
|
"""
|
|
try:
|
|
parsed_body: Final = await _read_request_body(request=request)
|
|
except ProxyException as parse_exception:
|
|
return {}, parse_exception # mutable-ok: request_data is a plain dict across the whole auth path
|
|
return populate_request_with_path_params(request_data=parsed_body, request=request), None
|
|
|
|
|
|
async def _record_unparsable_body_failure(
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
body_parse_exception: ProxyException,
|
|
route: str,
|
|
) -> None:
|
|
"""Record the 400 an unparsable body earns as a failed request log.
|
|
|
|
The endpoint never runs for these, so no downstream failure hook writes the
|
|
spend log row the Admin UI reads. Logging must not change what the caller
|
|
sees, so a failure here is swallowed and the 400 is raised either way.
|
|
"""
|
|
from litellm.proxy.proxy_server import proxy_logging_obj
|
|
|
|
try:
|
|
await proxy_logging_obj.post_call_failure_hook( # pyright: ignore[reportUnknownMemberType] # bare dict in sig
|
|
request_data={}, # mutable-ok: the failure hook seeds the call id and metadata onto this dict
|
|
original_exception=body_parse_exception,
|
|
user_api_key_dict=user_api_key_dict,
|
|
error_type=ProxyErrorTypes.bad_request_error,
|
|
route=route,
|
|
)
|
|
except Exception as e: # noqa: BLE001 # any logging failure must leave the caller's 400 untouched
|
|
verbose_proxy_logger.exception("Failed to log the request rejected for an unparsable body: %s", e)
|
|
|
|
|
|
async def _refresh_session_token_grants(
|
|
valid_token: UserAPIKeyAuth,
|
|
prisma_client: PrismaClient,
|
|
user_api_key_cache: UserApiKeyCache,
|
|
parent_otel_span: Span | None,
|
|
proxy_logging_obj: ProxyLogging,
|
|
) -> UserAPIKeyAuth:
|
|
"""Rebuild a ``lite login`` session token's grants from the live user and team rows.
|
|
|
|
The blob only proves who logged in and which team they picked. Team models, aliases, the user's own model
|
|
list, and their role are re-read every request, so a `/team/update` or a demotion shows up without a
|
|
re-login, and a user removed from the team or deleted outright is refused. When a row cannot be read for
|
|
a reason unrelated to the caller, the minted grants stand in exactly as they did before this refresh.
|
|
"""
|
|
outcome: Final = await GrantResolver(
|
|
prisma_client,
|
|
user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
load_user=get_user_object,
|
|
load_team=get_team_object,
|
|
load_membership=get_team_membership,
|
|
).resolve(UserLookup(user_id=valid_token.user_id), team_id=valid_token.team_id)
|
|
match outcome:
|
|
case ResolvedGrants(
|
|
user_object=LiteLLM_UserTable() as user_object, team_object=team_object, team_membership=team_membership
|
|
):
|
|
return UserAPIKeyAuth.model_validate(
|
|
MappingProxyType(
|
|
{
|
|
**valid_token.model_dump(exclude_none=True),
|
|
**team_grants(team_object, team_membership, user_object.user_id),
|
|
"user_role": _get_user_role(user_object),
|
|
"models": () if team_object is not None else user_models(user_object),
|
|
}
|
|
)
|
|
)
|
|
case ResolvedGrants():
|
|
return valid_token
|
|
case LookupDegraded(error=error):
|
|
verbose_proxy_logger.debug("Session token grants not refreshed, keeping minted grants: %s", error)
|
|
return valid_token
|
|
case _:
|
|
raise_public(outcome)
|
|
|
|
|
|
async def _resolve_object_permission_for_unresolvable_team(
|
|
object_permission_id: str | None,
|
|
prisma_client: PrismaClient | None,
|
|
user_api_key_cache: UserApiKeyCache,
|
|
parent_otel_span: Span | None,
|
|
proxy_logging_obj: ProxyLogging,
|
|
) -> LiteLLM_ObjectPermissionTable | None:
|
|
"""Re-resolve a team's object permission by id when the team row itself is unreadable, so the
|
|
token-derived fallback doesn't silently drop it."""
|
|
if object_permission_id is None or prisma_client is None:
|
|
return None
|
|
return await get_object_permission(
|
|
object_permission_id=object_permission_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
|
|
|
|
async def _user_api_key_auth_builder(
|
|
request: Request,
|
|
api_key: str,
|
|
azure_api_key_header: str,
|
|
anthropic_api_key_header: str | None,
|
|
google_ai_studio_api_key_header: str | None,
|
|
azure_apim_header: str | None,
|
|
request_data: dict,
|
|
custom_litellm_key_header: str | None = None,
|
|
) -> UserAPIKeyAuth:
|
|
from litellm.proxy.proxy_server import (
|
|
general_settings,
|
|
jwt_handler,
|
|
litellm_proxy_admin_name,
|
|
llm_model_list,
|
|
llm_router,
|
|
master_key,
|
|
model_max_budget_limiter,
|
|
open_telemetry_logger,
|
|
prisma_client,
|
|
proxy_logging_obj,
|
|
user_api_key_cache,
|
|
user_custom_auth,
|
|
)
|
|
|
|
parent_otel_span: Span | None = None
|
|
# Prefer the receive-instant stamped by the early helper in
|
|
# user_api_key_auth (before body parse) — overwriting it would shorten
|
|
# the preprocessing-duration measurement by the body-parse window.
|
|
start_time: Final = getattr(request.state, "litellm_received_at", None) or datetime.now(timezone.utc)
|
|
try:
|
|
request.state.litellm_received_at = start_time
|
|
except Exception:
|
|
pass
|
|
route: Final[str] = get_request_route(request=request)
|
|
valid_token: UserAPIKeyAuth | None = None
|
|
custom_auth_api_key: bool = False
|
|
|
|
try:
|
|
with tracer.trace("litellm.proxy.auth.pre_db_read_auth_checks"):
|
|
await pre_db_read_auth_checks(
|
|
request_data=request_data,
|
|
request=request,
|
|
route=route,
|
|
)
|
|
pass_through_endpoints: Final[list[dict] | None] = general_settings.get("pass_through_endpoints", None)
|
|
## CHECK IF X-LITELM-API-KEY IS PASSED IN - supercedes Authorization header
|
|
api_key, passed_in_key = get_api_key(
|
|
custom_litellm_key_header=custom_litellm_key_header,
|
|
api_key=api_key,
|
|
azure_api_key_header=azure_api_key_header,
|
|
anthropic_api_key_header=anthropic_api_key_header,
|
|
google_ai_studio_api_key_header=google_ai_studio_api_key_header,
|
|
azure_apim_header=azure_apim_header,
|
|
pass_through_endpoints=pass_through_endpoints,
|
|
route=route,
|
|
request=request,
|
|
)
|
|
# if user wants to pass LiteLLM_Master_Key as a custom header, example pass litellm keys as X-LiteLLM-Key: Bearer sk-1234
|
|
custom_litellm_key_header_name: Final = general_settings.get("litellm_key_header_name")
|
|
if custom_litellm_key_header_name is not None:
|
|
api_key = get_api_key_from_custom_header(
|
|
request=request,
|
|
custom_litellm_key_header_name=custom_litellm_key_header_name,
|
|
)
|
|
|
|
if open_telemetry_logger is not None:
|
|
# Reuse the span created by user_api_key_auth (before body parse)
|
|
# so it survives _read_request_body failures. For callers that
|
|
# bypass user_api_key_auth (e.g. MCP), create it lazily.
|
|
_ensure_parent_otel_span_on_request_state(request)
|
|
parent_otel_span = getattr(request.state, "parent_otel_span", None)
|
|
|
|
### USER-DEFINED AUTH FUNCTION ###
|
|
if enterprise_custom_auth is not None:
|
|
with tracer.trace("litellm.proxy.auth.enterprise_custom_auth"):
|
|
response = await enterprise_custom_auth(
|
|
request=request, api_key=api_key, user_custom_auth=user_custom_auth
|
|
)
|
|
if response is not None and isinstance(response, UserAPIKeyAuth):
|
|
validated = UserAPIKeyAuth.model_validate(response)
|
|
if getattr(litellm, "enable_post_custom_auth_checks", False):
|
|
validated = await _run_post_custom_auth_checks(
|
|
valid_token=validated,
|
|
request=request,
|
|
request_data=request_data,
|
|
route=route,
|
|
parent_otel_span=parent_otel_span,
|
|
)
|
|
return validated
|
|
elif response is not None and isinstance(response, str):
|
|
api_key = response
|
|
custom_auth_api_key = True
|
|
elif user_custom_auth is not None:
|
|
response = await user_custom_auth(request=request, api_key=api_key)
|
|
validated = UserAPIKeyAuth.model_validate(response)
|
|
if getattr(litellm, "enable_post_custom_auth_checks", False):
|
|
validated = await _run_post_custom_auth_checks(
|
|
valid_token=validated,
|
|
request=request,
|
|
request_data=request_data,
|
|
route=route,
|
|
parent_otel_span=parent_otel_span,
|
|
)
|
|
return validated
|
|
|
|
### LITELLM-DEFINED AUTH FUNCTION ###
|
|
#### IF JWT ####
|
|
"""
|
|
LiteLLM supports using JWTs.
|
|
|
|
Enable this in proxy config, by setting
|
|
```
|
|
general_settings:
|
|
enable_jwt_auth: true
|
|
```
|
|
"""
|
|
|
|
######## Route Checks Before Reading DB / Cache for "token" ################
|
|
if not _route_requires_auth_despite_public(route=route, general_settings=general_settings) and (
|
|
route in LiteLLMRoutes.public_routes.value or route_in_additonal_public_routes(current_route=route)
|
|
):
|
|
# check if public endpoint
|
|
return UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER_VIEW_ONLY)
|
|
|
|
########## End of Route Checks Before Reading DB / Cache for "token" ########
|
|
|
|
enable_oauth2_auth: Final = general_settings.get("enable_oauth2_auth", False) is True
|
|
enable_jwt_auth: Final = general_settings.get("enable_jwt_auth", False) is True
|
|
is_jwt = jwt_handler.is_jwt(token=api_key) if enable_jwt_auth else False
|
|
|
|
# Routing uses unverified JWT claims only to choose auth path.
|
|
# Final authentication is enforced by the selected validator.
|
|
route_jwt_to_oauth2 = is_jwt and _should_route_jwt_to_oauth2_override(token=api_key, jwt_handler=jwt_handler)
|
|
|
|
# OAuth2 applies for:
|
|
# 1) when global OAuth2 auth is enabled on LLM + info routes
|
|
# 2) JWT tokens that explicitly match routing_overrides on LLM + info routes
|
|
should_apply_override_oauth2: Final = route_jwt_to_oauth2 and (
|
|
RouteChecks.is_llm_api_route(route=route) or RouteChecks.is_info_route(route=route)
|
|
)
|
|
should_apply_global_oauth2: Final = enable_oauth2_auth and (
|
|
RouteChecks.is_llm_api_route(route=route) or RouteChecks.is_info_route(route=route)
|
|
)
|
|
if (should_apply_global_oauth2 and not is_jwt) or should_apply_override_oauth2:
|
|
from litellm.proxy.proxy_server import premium_user
|
|
|
|
if premium_user is not True:
|
|
raise ProxyException(
|
|
message="Oauth2 token validation is only available for premium users. "
|
|
+ CommonProxyErrors.not_premium_user.value,
|
|
type=ProxyErrorTypes.auth_error,
|
|
param="premium_user",
|
|
code=status.HTTP_403_FORBIDDEN,
|
|
)
|
|
|
|
return await Oauth2Handler.check_oauth2_token(token=api_key)
|
|
|
|
if general_settings.get("enable_oauth2_proxy_auth", False) is True:
|
|
return await handle_oauth2_proxy_request(request=request)
|
|
|
|
if general_settings.get("enable_jwt_auth", False) is True:
|
|
is_jwt = jwt_handler.is_jwt(token=api_key)
|
|
verbose_proxy_logger.debug("is_jwt: %s", is_jwt)
|
|
if is_jwt:
|
|
from litellm.proxy.proxy_server import premium_user
|
|
|
|
if premium_user is not True:
|
|
raise ProxyException(
|
|
message=f"JWT Auth is an enterprise only feature. {CommonProxyErrors.not_premium_user.value}",
|
|
type=ProxyErrorTypes.auth_error,
|
|
param="premium_user",
|
|
code=status.HTTP_403_FORBIDDEN,
|
|
)
|
|
# Try JWT-to-Virtual-Key mapping first to avoid
|
|
# unnecessary DB queries in auth_builder
|
|
do_standard_jwt_auth = True
|
|
pending_auto_register: _PendingAutoRegister | None = None
|
|
if jwt_handler.litellm_jwtauth.is_virtual_key_mapping_configured():
|
|
# Decode JWT to get claims without running full auth_builder
|
|
jwt_claims: dict | None
|
|
if jwt_handler.litellm_jwtauth.oidc_userinfo_enabled and not is_jwt:
|
|
jwt_claims = await jwt_handler.get_oidc_userinfo(token=api_key)
|
|
else:
|
|
jwt_claims = await jwt_handler.auth_jwt(token=api_key)
|
|
|
|
resolve_result: Final = await _resolve_jwt_to_virtual_key(
|
|
jwt_claims=jwt_claims,
|
|
jwt_handler=jwt_handler,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
if isinstance(resolve_result, UserAPIKeyAuth):
|
|
valid_token = resolve_result
|
|
api_key = valid_token.token or ""
|
|
valid_token.jwt_claims = jwt_claims
|
|
do_standard_jwt_auth = False
|
|
# Fall through to virtual key checks
|
|
if valid_token.user_id is not None and valid_token.user_email is None:
|
|
mapped_claims = jwt_claims or {} # mutable-ok: empty-dict fallback for the None-claims case
|
|
mapped_user_email = jwt_handler.get_user_email(token=mapped_claims, default_value=None)
|
|
mapped_jwt_user_id: Final = jwt_handler.get_user_id(token=mapped_claims, default_value=None)
|
|
if mapped_user_email is not None and mapped_jwt_user_id == valid_token.user_id:
|
|
try:
|
|
mapped_user_obj: Final = await get_user_object(
|
|
user_id=valid_token.user_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
user_id_upsert=False,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
user_email=mapped_user_email,
|
|
)
|
|
except Exception as e:
|
|
verbose_proxy_logger.debug("JWT mapped-key user_email backfill skipped: %s", e)
|
|
else:
|
|
if mapped_user_obj is not None:
|
|
valid_token.user_email = mapped_user_obj.user_email
|
|
elif isinstance(resolve_result, _PendingAutoRegister):
|
|
# Run full JWT policy (RBAC, scope, custom_validate,
|
|
# email-domain) via auth_builder, then create the key
|
|
# from the validated identity below.
|
|
pending_auto_register = resolve_result
|
|
# else: None → FALLBACK_TEAM_MAPPING, falls through to
|
|
# standard JWT auth_builder below
|
|
|
|
if do_standard_jwt_auth:
|
|
with tracer.trace("litellm.proxy.auth.jwt_auth_builder"):
|
|
result: Final = await JWTAuthManager.auth_builder(
|
|
request_data=request_data,
|
|
general_settings=general_settings,
|
|
api_key=api_key,
|
|
jwt_handler=jwt_handler,
|
|
route=route,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
parent_otel_span=parent_otel_span,
|
|
request_headers=_safe_get_request_headers(request),
|
|
request_method=RouteChecks._get_request_method(request=request),
|
|
)
|
|
|
|
is_proxy_admin: Final = result["is_proxy_admin"]
|
|
team_id: Final = result["team_id"]
|
|
team_object: Final = result["team_object"]
|
|
user_id: Final = result["user_id"]
|
|
user_email: Final = result["user_email"]
|
|
user_object: Final = result["user_object"]
|
|
end_user_id = result["end_user_id"]
|
|
org_id: Final = result["org_id"]
|
|
team_membership: Final[LiteLLM_TeamMembership | None] = result.get("team_membership", None)
|
|
jwt_claims = result.get("jwt_claims", None)
|
|
|
|
if is_proxy_admin:
|
|
# Proxy admins authenticate via auth_builder (full
|
|
# access), not via a mapped virtual key. If
|
|
# AUTO_REGISTER was pending, cache a sentinel so
|
|
# future requests from this JWT identity skip the
|
|
# DB mapping lookup in _resolve_jwt_to_virtual_key.
|
|
# Without this, every proxy-admin request under
|
|
# AUTO_REGISTER re-hits get_jwt_key_mapping_object.
|
|
if pending_auto_register is not None:
|
|
await user_api_key_cache.async_set_cache(
|
|
key=pending_auto_register.cache_key,
|
|
value=_JWT_PROXY_ADMIN_SENTINEL,
|
|
ttl=jwt_handler.litellm_jwtauth.virtual_key_mapping_cache_ttl,
|
|
)
|
|
return UserAPIKeyAuth(
|
|
api_key=None,
|
|
user_role=LitellmUserRoles.PROXY_ADMIN,
|
|
user_id=user_id,
|
|
user_email=user_email,
|
|
team_id=team_id,
|
|
org_id=org_id,
|
|
end_user_id=end_user_id,
|
|
parent_otel_span=parent_otel_span,
|
|
jwt_claims=jwt_claims,
|
|
**team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id),
|
|
)
|
|
|
|
valid_token = UserAPIKeyAuth(
|
|
api_key=None,
|
|
team_id=team_id,
|
|
user_role=(
|
|
LitellmUserRoles(user_object.user_role)
|
|
if user_object is not None and user_object.user_role is not None
|
|
else LitellmUserRoles.INTERNAL_USER
|
|
),
|
|
user_id=user_id,
|
|
user_email=user_email,
|
|
org_id=org_id,
|
|
parent_otel_span=parent_otel_span,
|
|
end_user_id=end_user_id,
|
|
user_tpm_limit=(user_object.tpm_limit if user_object is not None else None),
|
|
user_rpm_limit=(user_object.rpm_limit if user_object is not None else None),
|
|
user_model_max_budget=(user_object.model_max_budget if user_object is not None else None),
|
|
jwt_claims=jwt_claims,
|
|
**team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id),
|
|
)
|
|
|
|
# AUTO_REGISTER deferred from _resolve_jwt_to_virtual_key.
|
|
# JWT policy (RBAC, scope, custom_validate, email-domain)
|
|
# has now been enforced by auth_builder above. Create the
|
|
# mapping + virtual key from the *validated* identity, then
|
|
# replace valid_token with the new key so downstream checks
|
|
# use the key-scoped path.
|
|
if pending_auto_register is not None and prisma_client is not None:
|
|
auto_registered: Final = await _auto_register_jwt_mapping(
|
|
virtual_key_claim_field=pending_auto_register.claim_field,
|
|
claim_value=pending_auto_register.claim_value,
|
|
jwt_handler=jwt_handler,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
cache_key=pending_auto_register.cache_key,
|
|
team_id=team_id,
|
|
user_id=user_id,
|
|
org_id=org_id,
|
|
end_user_id=end_user_id,
|
|
)
|
|
if auto_registered is not None:
|
|
auto_registered.jwt_claims = jwt_claims
|
|
auto_registered.user_email = user_email
|
|
# The auto-registered token is built from the new key's
|
|
# columns, which carry no user budget. Carry over the
|
|
# already-loaded user row rather than re-reading it, or
|
|
# the budget check below has nothing to enforce.
|
|
auto_registered.user_model_max_budget = (
|
|
user_object.model_max_budget if user_object is not None else None
|
|
)
|
|
valid_token = auto_registered
|
|
api_key = valid_token.token or ""
|
|
|
|
# Check if model has zero cost - if so, skip all budget checks
|
|
model = _get_model_from_request_context(
|
|
request_data=request_data,
|
|
route=route,
|
|
request=request,
|
|
llm_router=llm_router,
|
|
team_id=valid_token.team_id,
|
|
)
|
|
skip_budget_checks = False
|
|
if model is not None and llm_router is not None:
|
|
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
|
|
|
|
skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router)
|
|
if skip_budget_checks:
|
|
verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model)
|
|
|
|
# Fetch project object for JWT path if project_id is set
|
|
_jwt_project_obj = None
|
|
if valid_token.project_id is not None:
|
|
_jwt_project_obj = await get_project_object(
|
|
project_id=valid_token.project_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
if _jwt_project_obj is not None:
|
|
valid_token.project_metadata = _jwt_project_obj.metadata
|
|
valid_token.project_alias = _jwt_project_obj.project_alias
|
|
|
|
# JWT auth returns here rather than falling through to the
|
|
# virtual-key checks below, so the user's per-model budget
|
|
# has to be enforced on this path too. Without it the
|
|
# post-call increment still charges the counter and nothing
|
|
# ever reads it, which is worse than not tracking at all.
|
|
# Guarded by the same flag the virtual-key path uses, or a
|
|
# zero-cost model would be refused here and allowed there,
|
|
# while the log above claims all budget checks were skipped.
|
|
if not skip_budget_checks:
|
|
await _check_user_model_budget(
|
|
valid_token=cast(UserAPIKeyAuth, valid_token),
|
|
model_max_budget_limiter=model_max_budget_limiter,
|
|
models=_get_model_names_for_budget_checks(
|
|
model=_get_model_from_request_context(
|
|
request_data=request_data,
|
|
route=route,
|
|
request=request,
|
|
llm_router=llm_router,
|
|
team_id=valid_token.team_id,
|
|
)
|
|
),
|
|
)
|
|
|
|
return cast(UserAPIKeyAuth, valid_token)
|
|
|
|
#### ELSE ####
|
|
## CHECK PASS-THROUGH ENDPOINTS ##
|
|
if not custom_auth_api_key:
|
|
response = await check_api_key_for_custom_headers_or_pass_through_endpoints(
|
|
request=request,
|
|
route=route,
|
|
pass_through_endpoints=pass_through_endpoints,
|
|
api_key=api_key,
|
|
)
|
|
if isinstance(response, str):
|
|
api_key = response
|
|
elif isinstance(response, UserAPIKeyAuth):
|
|
return response
|
|
if master_key is None:
|
|
if isinstance(api_key, str):
|
|
return UserAPIKeyAuth(
|
|
api_key=api_key,
|
|
user_role=LitellmUserRoles.INTERNAL_USER,
|
|
parent_otel_span=parent_otel_span,
|
|
)
|
|
else:
|
|
return UserAPIKeyAuth(
|
|
user_role=LitellmUserRoles.INTERNAL_USER,
|
|
parent_otel_span=parent_otel_span,
|
|
)
|
|
elif api_key is None: # only require api key if master key is set
|
|
raise Exception("No api key passed in.")
|
|
elif api_key == "":
|
|
# missing 'Bearer ' prefix
|
|
raise Exception("Malformed API Key passed in. Ensure Key has `Bearer ` prefix.")
|
|
|
|
if route == "/user/auth":
|
|
if general_settings.get("allow_user_auth", False) is True:
|
|
return UserAPIKeyAuth()
|
|
else:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
|
detail="'allow_user_auth' not set or set to False",
|
|
)
|
|
|
|
## Check END-USER OBJECT
|
|
_end_user_object = None
|
|
end_user_params: Final = {}
|
|
|
|
raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, _safe_get_request_headers(request))
|
|
end_user_id = await resolve_and_validate_end_user_id(
|
|
raw_end_user_id=raw_end_user_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
route=route,
|
|
)
|
|
if end_user_id:
|
|
try:
|
|
end_user_params["end_user_id"] = end_user_id
|
|
|
|
with tracer.trace("litellm.proxy.auth.get_end_user_object"):
|
|
_end_user_object = await get_end_user_object(
|
|
end_user_id=end_user_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
route=route,
|
|
)
|
|
if _end_user_object is not None:
|
|
end_user_params["allowed_model_region"] = _end_user_object.allowed_model_region
|
|
if _end_user_object.litellm_budget_table is not None:
|
|
_apply_budget_limits_to_end_user_params(
|
|
end_user_params=end_user_params,
|
|
budget_info=_end_user_object.litellm_budget_table,
|
|
end_user_id=end_user_id,
|
|
)
|
|
elif litellm.max_end_user_budget_id is not None:
|
|
# End user doesn't exist yet, but apply default budget limits if configured
|
|
from litellm.proxy.auth.auth_checks import (
|
|
get_default_end_user_budget,
|
|
)
|
|
|
|
default_budget: Final = await get_default_end_user_budget(
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
)
|
|
if default_budget is not None:
|
|
_apply_budget_limits_to_end_user_params(
|
|
end_user_params=end_user_params,
|
|
budget_info=default_budget,
|
|
end_user_id=end_user_id,
|
|
)
|
|
except Exception as e:
|
|
if isinstance(e, litellm.BudgetExceededError):
|
|
raise e
|
|
verbose_proxy_logger.debug("Unable to find user in db. Error - %s", e)
|
|
|
|
### CHECK IF ADMIN ###
|
|
# note: never string compare api keys, this is vulenerable to a time attack. Use secrets.compare_digest instead
|
|
### CHECK IF ADMIN ###
|
|
# note: never string compare api keys, this is vulenerable to a time attack. Use secrets.compare_digest instead
|
|
if valid_token is None:
|
|
## Check CACHE
|
|
try:
|
|
with tracer.trace("litellm.proxy.auth.get_key_object_check_cache"):
|
|
valid_token = IdentityStore.key_from_principal(
|
|
await IdentityStore(
|
|
prisma_client,
|
|
user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
check_cache_only=True,
|
|
).resolve(hashed_token=hash_token(api_key))
|
|
)
|
|
# Key-cache entries are written only after the proxy validated a
|
|
# virtual key or the master key, but via_virtual_key is exclude=True
|
|
# so serialization drops it; restore it at this trusted boundary.
|
|
# The UI-login JWT fallback below constructs its token from a
|
|
# decrypted blob, not this cache, and stays unmarked.
|
|
if isinstance(valid_token, UserAPIKeyAuth):
|
|
valid_token.via_virtual_key = True
|
|
except Exception:
|
|
verbose_logger.debug("api key not found in cache.")
|
|
valid_token = None
|
|
|
|
## Check UI/CLI Hash Key
|
|
# Attempt decryption for non-sk- tokens unless the operator has
|
|
# explicitly set EXPERIMENTAL_UI_LOGIN=false to disable it.
|
|
# Unset (None) keeps the new default of always attempting decryption;
|
|
# decryption fails closed for anything that is not a genuine blob.
|
|
if (
|
|
valid_token is None
|
|
and not api_key.startswith("sk-")
|
|
and get_secret_bool("EXPERIMENTAL_UI_LOGIN") is not False
|
|
):
|
|
valid_token = ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(api_key)
|
|
|
|
if valid_token is not None and valid_token.is_session_token and prisma_client is not None:
|
|
valid_token = await _refresh_session_token_grants( # rebind-ok: later checks read this name
|
|
valid_token=valid_token,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
|
|
if (
|
|
valid_token is not None
|
|
and isinstance(valid_token, UserAPIKeyAuth)
|
|
and valid_token.user_role == LitellmUserRoles.PROXY_ADMIN
|
|
):
|
|
if valid_token.expires is not None:
|
|
current_time = datetime.now(timezone.utc)
|
|
if isinstance(valid_token.expires, datetime):
|
|
expiry_time = valid_token.expires
|
|
else:
|
|
expiry_time = datetime.fromisoformat(valid_token.expires)
|
|
if expiry_time.tzinfo is None or expiry_time.tzinfo.utcoffset(expiry_time) is None:
|
|
expiry_time = expiry_time.replace(tzinfo=timezone.utc)
|
|
if expiry_time < current_time:
|
|
await _delete_cache_key_object(
|
|
hashed_token=hash_token(api_key),
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
raise ProxyException(
|
|
message=f"Authentication Error - Expired Key. Key Expiry time {expiry_time} and current time {current_time}",
|
|
type=ProxyErrorTypes.expired_key,
|
|
code=status.HTTP_401_UNAUTHORIZED,
|
|
param=abbreviate_api_key(api_key=api_key),
|
|
)
|
|
valid_token = update_valid_token_with_end_user_params(
|
|
valid_token=valid_token, end_user_params=end_user_params
|
|
)
|
|
valid_token.parent_otel_span = parent_otel_span
|
|
if _end_user_object is not None:
|
|
valid_token.end_user_object_permission = _end_user_object.object_permission
|
|
|
|
return valid_token
|
|
|
|
if (
|
|
valid_token is not None
|
|
and isinstance(valid_token, UserAPIKeyAuth)
|
|
and valid_token.team_id is not None
|
|
and valid_token.team_id != UI_TEAM_ID
|
|
):
|
|
## UPDATE TEAM VALUES BASED ON CACHED TEAM OBJECT - allows `/team/update` values to work for cached token
|
|
try:
|
|
team_obj: Final[LiteLLM_TeamTableCachedObj] = await get_team_object(
|
|
team_id=valid_token.team_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
check_cache_only=True,
|
|
)
|
|
|
|
if (
|
|
team_obj.last_refreshed_at is not None
|
|
and valid_token.last_refreshed_at is not None
|
|
and team_obj.last_refreshed_at > valid_token.last_refreshed_at
|
|
):
|
|
team_obj_dict: Final = team_obj.__dict__
|
|
|
|
for k, v in team_obj_dict.items():
|
|
field_name = f"team_{k}"
|
|
if field_name in valid_token.__fields__:
|
|
setattr(valid_token, field_name, v)
|
|
except Exception as e:
|
|
verbose_logger.debug(e) # moving from .warning to .debug as it spams logs when team missing from cache.
|
|
|
|
try:
|
|
is_master_key_valid = secrets.compare_digest(api_key, master_key)
|
|
except Exception:
|
|
is_master_key_valid = False
|
|
|
|
## VALIDATE MASTER KEY ##
|
|
if not isinstance(master_key, str):
|
|
raise HTTPException(
|
|
status_code=500,
|
|
detail={f"Master key must be a valid string. Current type={type(master_key)}"},
|
|
)
|
|
|
|
if is_master_key_valid:
|
|
# Substitute a stable alias for the raw master key so neither the
|
|
# master key nor its hash propagates into spend logs, Prometheus
|
|
# /metrics labels, audit trails, rate-limit buckets, or any other
|
|
# downstream consumer of UserAPIKeyAuth.api_key.
|
|
_user_api_key_obj = await _return_user_api_key_auth_obj(
|
|
user_obj=None,
|
|
user_role=LitellmUserRoles.PROXY_ADMIN,
|
|
api_key=LITELLM_PROXY_MASTER_KEY_ALIAS,
|
|
parent_otel_span=parent_otel_span,
|
|
valid_token_dict={
|
|
**end_user_params,
|
|
"user_id": litellm_proxy_admin_name,
|
|
},
|
|
route=route,
|
|
start_time=start_time,
|
|
)
|
|
asyncio.create_task(
|
|
_cache_key_object(
|
|
hashed_token=hash_token(master_key),
|
|
user_api_key_obj=_user_api_key_obj,
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
)
|
|
|
|
_user_api_key_obj = update_valid_token_with_end_user_params(
|
|
valid_token=_user_api_key_obj, end_user_params=end_user_params
|
|
)
|
|
_user_api_key_obj.via_virtual_key = True
|
|
|
|
return _user_api_key_obj
|
|
|
|
## IF it's not a master key
|
|
## Route should not be in master_key_only_routes
|
|
if route in LiteLLMRoutes.master_key_only_routes.value:
|
|
raise Exception(f"Tried to access route={route}, which is only for MASTER KEY")
|
|
|
|
## Check DB
|
|
|
|
if (
|
|
prisma_client is None
|
|
): # if both master key + user key submitted, and user key != master key, and no db connected, raise an error
|
|
raise ProxyException(
|
|
message="No connected db.",
|
|
type=ProxyErrorTypes.no_db_connection,
|
|
code=400,
|
|
param=None,
|
|
)
|
|
|
|
if valid_token is None:
|
|
if isinstance(api_key, str): # if generated token, make sure it starts with sk-.
|
|
_masked_key: Final = f"{api_key[:4]}****{api_key[-4:]}" if len(api_key) > 8 else "****"
|
|
if not api_key.startswith("sk-"):
|
|
_hint = _JWT_AUTH_DISABLED_HINT if not enable_jwt_auth and JWTHandler.is_jwt(token=api_key) else ""
|
|
_malformed_key_error = HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail=(
|
|
f"{INVALID_VIRTUAL_KEY_ERROR_MESSAGE}. Received={_masked_key}, "
|
|
f"expected to start with 'sk-'.{_hint}"
|
|
),
|
|
) # prevent token hashes from being used
|
|
# Stamp provenance here so log routing classifies this 401 by
|
|
# where it was raised, never by its message text.
|
|
setattr(_malformed_key_error, INVALID_VIRTUAL_KEY_ERROR_MARKER, True)
|
|
raise _malformed_key_error
|
|
else:
|
|
verbose_logger.warning(
|
|
"litellm.proxy.proxy_server.user_api_key_auth(): Warning - Key is not a string. Got type={}".format(
|
|
type(api_key) if api_key is not None else "None"
|
|
)
|
|
)
|
|
abbreviated_api_key: Final = abbreviate_api_key(api_key=api_key)
|
|
if api_key.startswith("sk-"):
|
|
api_key = hash_token(token=api_key)
|
|
|
|
try:
|
|
with tracer.trace("litellm.proxy.auth.get_key_object_from_db"):
|
|
valid_token = IdentityStore.key_from_principal(
|
|
await IdentityStore(
|
|
prisma_client,
|
|
user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
).resolve(hashed_token=api_key)
|
|
)
|
|
except ProxyException as e:
|
|
if e.code == 401 or e.code == "401":
|
|
e.message = f"Authentication Error, Invalid proxy server token passed. Received API Key = {abbreviated_api_key}, Key Hash (Token) ={api_key}. Unable to find token in cache or `LiteLLM_VerificationTokenTable`"
|
|
raise e
|
|
# update end-user params on valid token
|
|
# These can change per request - it's important to update them here
|
|
valid_token.end_user_id = end_user_params.get("end_user_id")
|
|
valid_token.end_user_tpm_limit = end_user_params.get("end_user_tpm_limit")
|
|
valid_token.end_user_rpm_limit = end_user_params.get("end_user_rpm_limit")
|
|
valid_token.allowed_model_region = end_user_params.get("allowed_model_region")
|
|
|
|
if valid_token is not None:
|
|
valid_token = _update_key_budget_with_temp_budget_increase(valid_token)
|
|
|
|
user_obj: LiteLLM_UserTable | None = None
|
|
valid_token_dict: dict = {}
|
|
if valid_token is not None:
|
|
# Got Valid Token from Cache, DB
|
|
# Run checks for
|
|
# 1. If token can call model
|
|
## 1a. If token can call fallback models (if client-side fallbacks given)
|
|
# 2. If user_id for this token is in budget
|
|
# 3. If the user spend within their own team is within budget
|
|
# 4. If 'user' passed to /chat/completions, /embeddings endpoint is in budget
|
|
# 5. If token is expired
|
|
# 6. If token spend is under Budget for the token
|
|
# 7. If token spend per model is under budget per model
|
|
# 8. If token spend is under team budget
|
|
# 9. If team spend is under team budget
|
|
|
|
## base case ## key is disabled
|
|
if valid_token.blocked is True:
|
|
raise Exception("Key is blocked. Update via `/key/unblock` if you're an admin.")
|
|
await _enforce_key_and_fallback_model_access(
|
|
valid_token=valid_token,
|
|
request_data=request_data,
|
|
route=route,
|
|
request=request,
|
|
llm_model_list=llm_model_list,
|
|
llm_router=llm_router,
|
|
)
|
|
await _prefetch_referenced_auth_objects(
|
|
valid_token, end_user_id=end_user_id, user_api_key_cache=user_api_key_cache, prisma_client=prisma_client
|
|
)
|
|
|
|
# Check 2. If user_id for this token is in budget - done in common_checks()
|
|
if valid_token.user_id is not None:
|
|
try:
|
|
with tracer.trace("litellm.proxy.auth.get_user_object"):
|
|
user_obj = await get_user_object(
|
|
user_id=valid_token.user_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
user_id_upsert=False,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
except Exception as e:
|
|
verbose_logger.debug(
|
|
"litellm.proxy.auth.user_api_key_auth.py::user_api_key_auth() - Unable to get user from db/cache. Setting user_obj to None. Exception received - %s",
|
|
e,
|
|
)
|
|
user_obj = None
|
|
|
|
if user_obj is not None:
|
|
# The joint verification-token view carries the key's columns only, so the
|
|
# user's own per-model budget reaches enforcement and the post-call
|
|
# increment through the row fetched here.
|
|
valid_token.user_model_max_budget = user_obj.model_max_budget
|
|
|
|
if (
|
|
user_obj is not None
|
|
and isinstance(user_obj.metadata, dict)
|
|
and user_obj.metadata.get("scim_active") is False
|
|
):
|
|
raise Exception(
|
|
f"User={valid_token.user_id} has been deactivated via SCIM. Keys owned by this user cannot be used."
|
|
)
|
|
|
|
# Check 2a. Check if model has zero cost - if so, skip all budget checks
|
|
model = _get_model_from_request_context(
|
|
request_data=request_data,
|
|
route=route,
|
|
request=request,
|
|
llm_router=llm_router,
|
|
team_id=valid_token.team_id,
|
|
)
|
|
skip_budget_checks = False
|
|
if model is not None and llm_router is not None:
|
|
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
|
|
|
|
skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router)
|
|
if skip_budget_checks:
|
|
verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model)
|
|
|
|
# Check 3. Check if user is in their team budget
|
|
if not skip_budget_checks and valid_token.team_member_spend is not None:
|
|
_user_id: Final = valid_token.user_id
|
|
_team_id: Final = valid_token.team_id
|
|
if prisma_client is not None and _user_id is not None and _team_id is not None:
|
|
_cache_key: Final = team_membership_auth_cache_key(team_id=_team_id, user_id=_user_id)
|
|
|
|
team_member_info = await user_api_key_cache.async_get_cache(
|
|
key=_cache_key,
|
|
model_type=LiteLLM_TeamMembership,
|
|
)
|
|
if team_member_info is None:
|
|
# read from DB
|
|
_db_member: Final = await TeamMembershipRepository(prisma_client).table.find_first(
|
|
where={
|
|
"user_id": _user_id,
|
|
"team_id": _team_id,
|
|
},
|
|
include={"litellm_budget_table": True},
|
|
)
|
|
if _db_member is not None:
|
|
team_member_info = LiteLLM_TeamMembership(**_db_member.model_dump())
|
|
await user_api_key_cache.async_set_cache(
|
|
key=_cache_key,
|
|
value=team_member_info,
|
|
model_type=LiteLLM_TeamMembership,
|
|
ttl=5,
|
|
)
|
|
|
|
if team_member_info is not None and team_member_info.litellm_budget_table is not None:
|
|
team_member_budget: Final = team_member_info.litellm_budget_table.max_budget
|
|
if team_member_budget is not None and team_member_budget > 0:
|
|
# Read from cross-pod counter (Redis-first) if available
|
|
from litellm.proxy.proxy_server import get_current_spend
|
|
|
|
team_member_spend = valid_token.team_member_spend
|
|
if valid_token.user_id is not None and valid_token.team_id is not None:
|
|
team_member_spend = await get_current_spend(
|
|
counter_key=f"spend:team_member:{valid_token.user_id}:{valid_token.team_id}",
|
|
fallback_spend=team_member_spend,
|
|
max_budget=team_member_budget,
|
|
)
|
|
if team_member_spend >= team_member_budget:
|
|
_entity_id: Final = f"{valid_token.user_id}:{valid_token.team_id}"
|
|
raise litellm.BudgetExceededError(
|
|
current_cost=team_member_spend,
|
|
max_budget=team_member_budget,
|
|
message=(
|
|
f"Budget has been exceeded! TeamMember={_entity_id} "
|
|
f"Current cost: {team_member_spend}, Max budget: {team_member_budget}"
|
|
),
|
|
entity_type=Litellm_EntityType.TEAM_MEMBER.value,
|
|
entity_id=_entity_id,
|
|
)
|
|
|
|
# Check 3. If token is expired
|
|
if valid_token.expires is not None:
|
|
current_time = datetime.now(timezone.utc)
|
|
if isinstance(valid_token.expires, datetime):
|
|
expiry_time = valid_token.expires
|
|
else:
|
|
expiry_time = datetime.fromisoformat(valid_token.expires)
|
|
if expiry_time.tzinfo is None or expiry_time.tzinfo.utcoffset(expiry_time) is None:
|
|
expiry_time = expiry_time.replace(tzinfo=timezone.utc)
|
|
verbose_proxy_logger.debug(
|
|
"Checking if token expired, expiry time %s and current time %s", expiry_time, current_time
|
|
)
|
|
if expiry_time < current_time:
|
|
# Token exists but is expired.
|
|
raise ProxyException(
|
|
message=f"Authentication Error - Expired Key. Key Expiry time {expiry_time} and current time {current_time}",
|
|
type=ProxyErrorTypes.expired_key,
|
|
code=status.HTTP_401_UNAUTHORIZED,
|
|
param=abbreviate_api_key(api_key=api_key),
|
|
)
|
|
|
|
if not skip_budget_checks:
|
|
with tracer.trace("litellm.proxy.auth.budget_checks"):
|
|
# Check 4. Max Budget Alert Check (runs before budget enforcement
|
|
# so multi-threshold 100% alerts fire on the request that crosses
|
|
# max_budget, before BudgetExceededError is raised below)
|
|
await _virtual_key_max_budget_alert_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
user_obj=user_obj,
|
|
)
|
|
|
|
# Check 5. Token Spend is under budget
|
|
if RouteChecks.is_llm_api_route(route=route):
|
|
await _virtual_key_max_budget_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
user_obj=user_obj,
|
|
)
|
|
|
|
# Check 6. Soft Budget Check
|
|
await _virtual_key_soft_budget_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
user_obj=user_obj,
|
|
)
|
|
|
|
# Check 5. Token Model Spend is under Model budget
|
|
max_budget_per_model: Final = valid_token.model_max_budget
|
|
current_model = _get_model_from_request_context(
|
|
request_data=request_data,
|
|
route=route,
|
|
request=request,
|
|
llm_router=llm_router,
|
|
team_id=valid_token.team_id,
|
|
)
|
|
current_models = _get_model_names_for_budget_checks(model=current_model)
|
|
|
|
if (
|
|
max_budget_per_model is not None
|
|
and isinstance(max_budget_per_model, dict)
|
|
and len(max_budget_per_model) > 0
|
|
and prisma_client is not None
|
|
and current_models
|
|
and valid_token.token is not None
|
|
):
|
|
## GET THE SPEND FOR THIS MODEL
|
|
for model_name in current_models:
|
|
await _check_key_model_budget_with_fallback(
|
|
valid_token=valid_token,
|
|
model_max_budget_limiter=model_max_budget_limiter,
|
|
model_name=model_name,
|
|
request_data=request_data,
|
|
request=request,
|
|
llm_model_list=llm_model_list,
|
|
llm_router=llm_router,
|
|
)
|
|
|
|
# Recompute after a potential budget-fallback rewrite so
|
|
# the end-user check below validates the final model
|
|
current_model = _get_model_from_request_context(
|
|
request_data=request_data,
|
|
route=route,
|
|
request=request,
|
|
llm_router=llm_router,
|
|
team_id=valid_token.team_id,
|
|
)
|
|
current_models = _get_model_names_for_budget_checks(model=current_model)
|
|
|
|
# Check 5a. Internal user model_max_budget
|
|
if current_models:
|
|
await _check_user_model_budget(
|
|
valid_token=valid_token,
|
|
model_max_budget_limiter=model_max_budget_limiter,
|
|
models=current_models,
|
|
)
|
|
|
|
# Check 5b. End-user model max budget
|
|
end_user_mmb: Final = valid_token.end_user_model_max_budget
|
|
if (
|
|
end_user_mmb is not None
|
|
and isinstance(end_user_mmb, dict)
|
|
and len(end_user_mmb) > 0
|
|
and current_models
|
|
and valid_token.end_user_id is not None
|
|
):
|
|
for model_name in current_models:
|
|
await model_max_budget_limiter.is_end_user_within_model_budget(
|
|
end_user_id=valid_token.end_user_id,
|
|
end_user_model_max_budget=end_user_mmb,
|
|
model=model_name,
|
|
)
|
|
|
|
# Check 6: Additional Common Checks across jwt + key auth
|
|
if valid_token.team_id is not None:
|
|
try:
|
|
if valid_token.team_id == UI_TEAM_ID:
|
|
raise TeamNotFoundError(team_id=UI_TEAM_ID)
|
|
with tracer.trace("litellm.proxy.auth.get_team_object"):
|
|
_team_obj = await get_team_object(
|
|
team_id=valid_token.team_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
except HTTPException:
|
|
token_team_models: Final = _token_team_models(valid_token)
|
|
_team_obj = LiteLLM_TeamTableCachedObj(
|
|
team_id=valid_token.team_id,
|
|
max_budget=valid_token.team_max_budget,
|
|
soft_budget=valid_token.team_soft_budget,
|
|
spend=valid_token.team_spend,
|
|
tpm_limit=valid_token.team_tpm_limit,
|
|
rpm_limit=valid_token.team_rpm_limit,
|
|
blocked=valid_token.team_blocked,
|
|
models=token_team_models,
|
|
metadata=valid_token.team_metadata,
|
|
object_permission_id=valid_token.team_object_permission_id,
|
|
object_permission=await _resolve_object_permission_for_unresolvable_team(
|
|
object_permission_id=valid_token.team_object_permission_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
),
|
|
)
|
|
else:
|
|
_team_obj = None
|
|
|
|
if _team_obj is not None:
|
|
valid_token.team_object_permission = _team_obj.object_permission
|
|
# Keep team_metadata in sync with the freshly fetched team so that
|
|
# guardrails (or any other metadata) added after the key was cached
|
|
# are picked up on subsequent requests without a cache eviction.
|
|
valid_token.team_metadata = _team_obj.metadata
|
|
else:
|
|
valid_token.team_object_permission = None
|
|
|
|
# Fetch project object if key belongs to a project
|
|
_project_obj = None
|
|
if valid_token.project_id is not None:
|
|
_project_obj = await get_project_object(
|
|
project_id=valid_token.project_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
if _project_obj is not None:
|
|
valid_token.project_metadata = _project_obj.metadata
|
|
valid_token.project_alias = _project_obj.project_alias
|
|
|
|
global_proxy_spend = None
|
|
if litellm.max_budget > 0 and prisma_client is not None: # user set proxy max budget
|
|
cache_key: Final = GLOBAL_PROXY_SPEND_CACHE_KEY
|
|
with tracer.trace("litellm.proxy.auth.get_global_proxy_spend"):
|
|
global_proxy_spend = await _fetch_global_spend_with_event_coordination(
|
|
cache_key=cache_key,
|
|
user_api_key_cache=user_api_key_cache,
|
|
prisma_client=prisma_client,
|
|
)
|
|
|
|
if global_proxy_spend is not None:
|
|
call_info: Final = CallInfo(
|
|
token=valid_token.token,
|
|
spend=global_proxy_spend,
|
|
max_budget=litellm.max_budget,
|
|
user_id=litellm_proxy_admin_name,
|
|
team_id=valid_token.team_id,
|
|
event_group=Litellm_EntityType.PROXY,
|
|
)
|
|
asyncio.create_task(
|
|
proxy_logging_obj.budget_alerts(
|
|
type="proxy_budget",
|
|
user_info=call_info,
|
|
)
|
|
)
|
|
# Token passed all checks
|
|
if valid_token is None:
|
|
raise HTTPException(401, detail="Invalid API key")
|
|
if valid_token.token is None:
|
|
raise HTTPException(401, detail="Invalid API key, no token associated")
|
|
api_key = valid_token.token
|
|
|
|
valid_token_dict = valid_token.model_dump(exclude_none=True)
|
|
valid_token_dict.pop("token", None)
|
|
# budget_throttle_pct is excluded from model_dump (it must not leak
|
|
# into serialized responses), so carry the request-scoped decision
|
|
# forward by hand to the auth object the rate limiter receives.
|
|
if valid_token.budget_throttle_pct is not None:
|
|
valid_token_dict["budget_throttle_pct"] = valid_token.budget_throttle_pct
|
|
|
|
if _end_user_object is not None:
|
|
valid_token_dict.update(end_user_params)
|
|
valid_token_dict["end_user_object_permission"] = _end_user_object.object_permission
|
|
|
|
# check if token is from litellm-ui, litellm ui makes keys to allow users to login with sso. These keys can only be used for LiteLLM UI functions
|
|
# sso/login, ui/login, /key functions and /user functions
|
|
# this will never be allowed to call /chat/completions
|
|
|
|
if valid_token is None:
|
|
# No token was found when looking up in the DB
|
|
raise Exception("Invalid proxy server token passed")
|
|
if valid_token_dict is not None:
|
|
virtual_key_auth_obj: Final = await _return_user_api_key_auth_obj(
|
|
user_obj=user_obj,
|
|
api_key=api_key,
|
|
parent_otel_span=parent_otel_span,
|
|
valid_token_dict=valid_token_dict,
|
|
route=route,
|
|
start_time=start_time,
|
|
)
|
|
virtual_key_auth_obj.via_virtual_key = True
|
|
return virtual_key_auth_obj
|
|
except Exception as e:
|
|
return await UserAPIKeyAuthExceptionHandler._handle_authentication_error(
|
|
e=e,
|
|
request=request,
|
|
request_data=request_data,
|
|
route=route,
|
|
parent_otel_span=parent_otel_span,
|
|
api_key=api_key,
|
|
resolved_identity=valid_token,
|
|
)
|
|
|
|
|
|
async def _safe_fetch(label: str, awaitable):
|
|
"""Run an awaitable and return its result. Re-raises authentication /
|
|
authorization failures (HTTPException, ProxyException,
|
|
BudgetExceededError) so they propagate to the caller.
|
|
Other exceptions (e.g. transient DB errors fetching context) are
|
|
swallowed with a debug log and ``None`` is returned so
|
|
``common_checks`` can still run against whatever limits are recorded
|
|
directly on the token.
|
|
"""
|
|
try:
|
|
return await awaitable
|
|
except (HTTPException, ProxyException, litellm.BudgetExceededError) as e:
|
|
verbose_proxy_logger.debug(
|
|
"centralized auth: %s fetch failed (%s: %s)",
|
|
label,
|
|
type(e).__name__,
|
|
e,
|
|
)
|
|
raise
|
|
except Exception as e:
|
|
verbose_proxy_logger.debug(
|
|
"centralized auth: %s fetch swallowed (%s: %s)",
|
|
label,
|
|
type(e).__name__,
|
|
e,
|
|
)
|
|
return None
|
|
|
|
|
|
def _team_obj_from_token(valid_token: UserAPIKeyAuth) -> LiteLLM_TeamTableCachedObj:
|
|
"""Reconstruct a cached team object from the fields already on the
|
|
UserAPIKeyAuth. Only called when valid_token.team_id is known to be
|
|
non-None (the caller gates on it)."""
|
|
assert valid_token.team_id is not None
|
|
token_team_models: Final = _token_team_models(valid_token)
|
|
return LiteLLM_TeamTableCachedObj(
|
|
team_id=valid_token.team_id,
|
|
max_budget=valid_token.team_max_budget,
|
|
soft_budget=valid_token.team_soft_budget,
|
|
spend=valid_token.team_spend,
|
|
tpm_limit=valid_token.team_tpm_limit,
|
|
rpm_limit=valid_token.team_rpm_limit,
|
|
blocked=valid_token.team_blocked,
|
|
models=token_team_models,
|
|
metadata=valid_token.team_metadata,
|
|
object_permission_id=valid_token.team_object_permission_id,
|
|
)
|
|
|
|
|
|
def _token_can_vouch_for_team(valid_token: UserAPIKeyAuth, lookup_error: BaseException) -> bool:
|
|
"""Whether the token's own team fields may stand in for a team that failed to
|
|
resolve, without widening access.
|
|
|
|
The UI dashboard mints every session key against the ``UI_TEAM_ID`` sentinel,
|
|
which by design never has a team row, so a failed lookup for it is not a
|
|
degraded read to be treated with suspicion; it always vouches, exactly as it
|
|
always safely has (these keys are restricted elsewhere to UI-only routes).
|
|
|
|
For every other team, a team that is provably gone is a definitive answer,
|
|
not a degraded read, so nothing may stand in for it and no setting may
|
|
override that.
|
|
|
|
Otherwise the team's grant is merely unknown. A token carrying one may vouch,
|
|
since replaying a recorded grant cannot widen it and denying every team key
|
|
while the row is briefly unreadable would trade the widening for an outage. A
|
|
token carrying none may not: ``team_models=[]`` reads as every model and
|
|
``team_blocked=False`` as unblocked. ``allow_requests_on_db_unavailable`` opts
|
|
back out, and is only consulted here because the failure is known by this
|
|
point to be a degraded read.
|
|
"""
|
|
if valid_token.team_id == UI_TEAM_ID:
|
|
return True
|
|
if isinstance(lookup_error, TeamNotFoundError):
|
|
return False
|
|
if valid_token.team_models:
|
|
return True
|
|
return PrismaDBExceptionHandler.should_allow_request_on_db_unavailable()
|
|
|
|
|
|
@tracer.wrap()
|
|
async def _run_centralized_common_checks(
|
|
user_api_key_auth_obj: UserAPIKeyAuth,
|
|
request: Request,
|
|
request_data: dict,
|
|
route: str,
|
|
) -> None:
|
|
"""Run ``common_checks`` once at the ``user_api_key_auth`` wrapper
|
|
boundary, regardless of which ``_user_api_key_auth_builder`` path
|
|
returned. This is the single invariant enforcement point for key
|
|
model-access, budgets, guardrails, org, and vector-store checks.
|
|
|
|
Invariants:
|
|
- ``user_custom_auth`` with ``custom_auth_run_common_checks`` unset
|
|
skips the gate — matches the existing custom-auth RPS guarantee.
|
|
Custom-auth deployments don't use OAuth2 / DB-fallback paths, so
|
|
the skip does not re-open any bypass.
|
|
- ``PROXY_ADMIN`` tokens still run through ``common_checks`` so
|
|
team-blocked / team-budget / end-user-budget / tag-budget /
|
|
vector-store / tool-allowlist enforcement applies to admin keys
|
|
too. Admin status is honored where the underlying check exempts it
|
|
(``_is_api_route_allowed``, ``organization_role_based_access_check``).
|
|
"""
|
|
from litellm.proxy.proxy_server import (
|
|
general_settings,
|
|
litellm_proxy_admin_name,
|
|
llm_router,
|
|
master_key,
|
|
prisma_client,
|
|
proxy_logging_obj,
|
|
user_api_key_cache,
|
|
user_custom_auth,
|
|
)
|
|
|
|
# Public routes (e.g. /health/liveness) are exempt from
|
|
# auth in the builder — the wrapper must not retroactively apply
|
|
# authz on top, or k8s readiness probes and other unauthenticated
|
|
# callers get 401.
|
|
if route in LiteLLMRoutes.public_routes.value or route_in_additonal_public_routes(current_route=route):
|
|
return
|
|
|
|
# User-configured pass-through endpoints with ``auth: false`` are
|
|
# explicitly unauthenticated — the builder returns an empty
|
|
# UserAPIKeyAuth() and the request is forwarded as-is. Running
|
|
# common_checks on the empty token would reject the request as
|
|
# admin-only. The "auth" flag on the endpoint config is the
|
|
# contract; honor it.
|
|
pass_through_endpoints: Final = general_settings.get("pass_through_endpoints", None)
|
|
if pass_through_endpoints is not None:
|
|
for endpoint in pass_through_endpoints:
|
|
if isinstance(endpoint, dict) and endpoint.get("path", "") == route and endpoint.get("auth") is not True:
|
|
return
|
|
|
|
# No-auth dev mode: master_key unset AND no JWT/OAuth2 auth
|
|
# configured. The builder returns an INTERNAL_USER token for any
|
|
# api_key; the proxy is unauthenticated by configuration.
|
|
# Running common_checks would block every admin route on these
|
|
# deployments where that was previously not the contract. If any
|
|
# authn is enabled (JWT, OAuth2, OAuth2-proxy), authz must run.
|
|
if master_key is None and not (
|
|
general_settings.get("enable_jwt_auth", False)
|
|
or general_settings.get("enable_oauth2_auth", False)
|
|
or general_settings.get("enable_oauth2_proxy_auth", False)
|
|
):
|
|
return
|
|
|
|
if user_custom_auth is not None and not general_settings.get("custom_auth_run_common_checks", False):
|
|
return
|
|
|
|
parent_otel_span: Final = user_api_key_auth_obj.parent_otel_span
|
|
# In the integrated auth flow ``_user_api_key_auth_builder`` has already
|
|
# resolved the end-user id and attached it here. Reuse that to avoid a
|
|
# second extraction pass; fall back to extracting locally when the
|
|
# function is invoked in isolation (e.g. in direct unit tests).
|
|
end_user_id = user_api_key_auth_obj.end_user_id
|
|
if end_user_id is None:
|
|
raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, _safe_get_request_headers(request))
|
|
end_user_id = await resolve_and_validate_end_user_id(
|
|
raw_end_user_id=raw_end_user_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
route=route,
|
|
)
|
|
|
|
fetch_coros: Final = []
|
|
if user_api_key_auth_obj.team_id is not None and user_api_key_auth_obj.team_id != UI_TEAM_ID:
|
|
fetch_coros.append(
|
|
_safe_fetch(
|
|
"team",
|
|
get_team_object(
|
|
team_id=user_api_key_auth_obj.team_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
),
|
|
)
|
|
)
|
|
else:
|
|
fetch_coros.append(_safe_fetch("team", _noop_none()))
|
|
|
|
if user_api_key_auth_obj.user_id is not None:
|
|
fetch_coros.append(
|
|
_safe_fetch(
|
|
"user",
|
|
get_user_object(
|
|
user_id=user_api_key_auth_obj.user_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
user_id_upsert=False,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
),
|
|
)
|
|
)
|
|
else:
|
|
fetch_coros.append(_safe_fetch("user", _noop_none()))
|
|
|
|
if user_api_key_auth_obj.project_id is not None:
|
|
fetch_coros.append(
|
|
_safe_fetch(
|
|
"project",
|
|
get_project_object(
|
|
project_id=user_api_key_auth_obj.project_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
),
|
|
)
|
|
)
|
|
else:
|
|
fetch_coros.append(_safe_fetch("project", _noop_none()))
|
|
|
|
if end_user_id:
|
|
fetch_coros.append(
|
|
_safe_fetch(
|
|
"end_user",
|
|
get_end_user_object(
|
|
end_user_id=end_user_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
route=route,
|
|
token_end_user_max_budget=user_api_key_auth_obj.end_user_max_budget,
|
|
),
|
|
)
|
|
)
|
|
else:
|
|
fetch_coros.append(_safe_fetch("end_user", _noop_none()))
|
|
|
|
fetch_coros.append(
|
|
_safe_fetch(
|
|
"global_spend",
|
|
get_global_proxy_spend(
|
|
litellm_proxy_admin_name=litellm_proxy_admin_name,
|
|
user_api_key_cache=user_api_key_cache,
|
|
prisma_client=prisma_client,
|
|
token=user_api_key_auth_obj.token or "",
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
),
|
|
)
|
|
)
|
|
|
|
# Per-fetch error isolation. ``_safe_fetch`` lets HTTPException,
|
|
# ProxyException, and BudgetExceededError escape (everything else is
|
|
# already swallowed to None). A bare ``except`` over ``gather`` would
|
|
# let one fetch's HTTPException null out every other context — e.g.
|
|
# a 404 from ``get_team_object`` (token references a deleted team)
|
|
# would silently skip the user, end-user, project, and global-spend
|
|
# checks. Use ``return_exceptions=True`` and apply per-fetch fallback
|
|
# so a missing team only zeros out the team object.
|
|
(
|
|
team_result,
|
|
user_result,
|
|
project_result,
|
|
end_user_result,
|
|
global_spend_result,
|
|
) = await asyncio.gather(*fetch_coros, return_exceptions=True)
|
|
|
|
# ProxyException / BudgetExceededError are authorization failures —
|
|
# propagate so the wrapper renders them. HTTPException is fallback
|
|
# material (404 from get_team_object is the only known producer).
|
|
for r in (
|
|
team_result,
|
|
user_result,
|
|
project_result,
|
|
end_user_result,
|
|
global_spend_result,
|
|
):
|
|
if isinstance(r, (ProxyException, litellm.BudgetExceededError)):
|
|
raise r
|
|
|
|
# Use BaseException (not HTTPException) in the narrowing checks so
|
|
# mypy can narrow ``Any | BaseException`` to the typed object in the
|
|
# else branch. After the for-loop above, the only BaseException that
|
|
# can still appear here is HTTPException (other listed re-raises were
|
|
# propagated; non-listed exceptions were already swallowed to None).
|
|
team_object: LiteLLM_TeamTableCachedObj | None
|
|
if isinstance(team_result, BaseException):
|
|
# Token-derived fallback only valid when a team_id is set;
|
|
# _team_obj_from_token asserts that precondition.
|
|
if user_api_key_auth_obj.team_id is None:
|
|
team_object = None
|
|
elif _token_can_vouch_for_team(user_api_key_auth_obj, team_result):
|
|
team_object = _team_obj_from_token(user_api_key_auth_obj)
|
|
else:
|
|
raise team_result
|
|
else:
|
|
team_object = (
|
|
_team_obj_from_token(user_api_key_auth_obj) if user_api_key_auth_obj.team_id == UI_TEAM_ID else team_result
|
|
)
|
|
|
|
user_object: LiteLLM_UserTable | None = None if isinstance(user_result, BaseException) else user_result
|
|
project_object: Final[LiteLLM_ProjectTableCachedObj | None] = (
|
|
None if isinstance(project_result, BaseException) else project_result
|
|
)
|
|
end_user_object: Final[LiteLLM_EndUserTable | None] = (
|
|
None if isinstance(end_user_result, BaseException) else end_user_result
|
|
)
|
|
global_proxy_spend: float | None = None if isinstance(global_spend_result, BaseException) else global_spend_result
|
|
carry_team_and_user_budget_state(
|
|
valid_token=user_api_key_auth_obj,
|
|
team_object=team_object,
|
|
user_object=user_object,
|
|
)
|
|
|
|
if user_api_key_auth_obj.org_id is None and team_object is not None and team_object.organization_id is not None:
|
|
user_api_key_auth_obj.org_id = team_object.organization_id
|
|
|
|
# common_checks identifies admin via user_object, not the token
|
|
# (non_proxy_admin_allowed_routes_check). JWT admin shortcut and
|
|
# master_key tokens get admin from the token; the DB row for the
|
|
# same user_id (e.g. litellm_proxy_admin_name = "default_user_id")
|
|
# may have a non-admin user_role and would otherwise demote the
|
|
# caller. The token is the source of truth for these paths — force
|
|
# the admin user_object whenever the token says PROXY_ADMIN, even
|
|
# if a DB row was fetched.
|
|
if user_api_key_auth_obj.user_role == LitellmUserRoles.PROXY_ADMIN:
|
|
user_object = LiteLLM_UserTable(
|
|
user_id=user_api_key_auth_obj.user_id or litellm_proxy_admin_name,
|
|
user_role=LitellmUserRoles.PROXY_ADMIN,
|
|
spend=user_object.spend if user_object is not None else 0.0,
|
|
)
|
|
|
|
if project_object is not None:
|
|
user_api_key_auth_obj.project_metadata = project_object.metadata
|
|
user_api_key_auth_obj.project_alias = project_object.project_alias
|
|
|
|
skip_budget_checks: Final = _should_skip_budget_checks(
|
|
request_data=request_data,
|
|
route=route,
|
|
request=request,
|
|
llm_router=llm_router,
|
|
team_id=user_api_key_auth_obj.team_id,
|
|
)
|
|
|
|
# Pin the metadata variable name (litellm_metadata vs metadata) before
|
|
# any tag merge runs. Without this, header tags from
|
|
# apply_client_tag_policy_pre_auth would land in `metadata` while the
|
|
# later seed in common_checks pushes key tags and the
|
|
# _tag_max_budget_check read into `litellm_metadata`, hiding header
|
|
# tags from per-tag budget enforcement on LITELLM_METADATA_ROUTES.
|
|
LiteLLMProxyRequestSetup.pre_seed_litellm_metadata_for_route(
|
|
request_data=request_data,
|
|
route=route,
|
|
)
|
|
|
|
# Merge x-litellm-tags into request_data BEFORE common_checks runs.
|
|
# _tag_max_budget_check inside common_checks only inspects request_data;
|
|
# without this pre-merge, header-supplied tags bypass tag-budget
|
|
# enforcement.
|
|
LiteLLMProxyRequestSetup.apply_client_tag_policy_pre_auth(
|
|
request=request,
|
|
request_data=request_data,
|
|
user_api_key_dict=user_api_key_auth_obj,
|
|
)
|
|
|
|
bind_admission_counter_keys(user_api_key_auth_obj, end_user_id=end_user_id)
|
|
try:
|
|
_ = await common_checks(
|
|
request=request,
|
|
request_body=request_data,
|
|
team_object=team_object,
|
|
user_object=user_object,
|
|
end_user_object=end_user_object,
|
|
general_settings=general_settings,
|
|
global_proxy_spend=global_proxy_spend,
|
|
route=route,
|
|
llm_router=llm_router,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
valid_token=user_api_key_auth_obj,
|
|
skip_budget_checks=skip_budget_checks,
|
|
project_object=project_object,
|
|
)
|
|
finally:
|
|
release_spend_counter_batch()
|
|
|
|
await _reserve_budget_after_common_checks(
|
|
user_api_key_auth_obj=user_api_key_auth_obj,
|
|
request=request,
|
|
request_data=request_data,
|
|
route=route,
|
|
llm_router=llm_router,
|
|
team_object=team_object,
|
|
user_object=user_object,
|
|
end_user_id=end_user_id,
|
|
end_user_object=end_user_object,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
skip_budget_checks=skip_budget_checks,
|
|
general_settings=general_settings,
|
|
)
|
|
|
|
|
|
async def _noop_none() -> None:
|
|
"""Sentinel coroutine for asyncio.gather when a fetch is unnecessary
|
|
(e.g. token has no team_id). Keeps the result tuple positional."""
|
|
return
|
|
|
|
|
|
async def _reserve_budget_after_common_checks(
|
|
user_api_key_auth_obj: UserAPIKeyAuth,
|
|
request_data: dict,
|
|
route: str,
|
|
llm_router: Any | None,
|
|
team_object: LiteLLM_TeamTableCachedObj | None,
|
|
user_object: LiteLLM_UserTable | None,
|
|
prisma_client: PrismaClient | None,
|
|
user_api_key_cache: UserApiKeyCache,
|
|
proxy_logging_obj: ProxyLogging,
|
|
skip_budget_checks: bool,
|
|
general_settings: dict,
|
|
end_user_id: str | None = None,
|
|
end_user_object: LiteLLM_EndUserTable | None = None,
|
|
request: Request | None = None,
|
|
) -> None:
|
|
user_api_key_auth_obj.budget_reservation = None
|
|
if skip_budget_checks:
|
|
return
|
|
if general_settings.get("disable_budget_reservation") is True:
|
|
return
|
|
|
|
from litellm.proxy.spend_tracking.budget_reservation import (
|
|
reserve_budget_for_request,
|
|
)
|
|
|
|
user_api_key_auth_obj.budget_reservation = await reserve_budget_for_request(
|
|
request_body=request_data,
|
|
route=route,
|
|
llm_router=llm_router,
|
|
valid_token=user_api_key_auth_obj,
|
|
team_object=team_object,
|
|
user_object=user_object,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
end_user_id=end_user_id,
|
|
end_user_object=end_user_object,
|
|
apply_user_budget_to_team_keys=general_settings.get("apply_user_budget_to_team_keys") is True,
|
|
fail_closed_budget_enforcement=general_settings.get("fail_closed_budget_enforcement") is True,
|
|
raw_body=await read_raw_json_body(request=request),
|
|
)
|
|
|
|
|
|
def _should_skip_budget_checks(
|
|
request_data: dict,
|
|
route: str,
|
|
request: Request | None,
|
|
llm_router: Any | None,
|
|
team_id: str | None = None,
|
|
) -> bool:
|
|
model: Final = _get_model_from_request_context(
|
|
request_data=request_data,
|
|
route=route,
|
|
request=request,
|
|
llm_router=llm_router,
|
|
team_id=team_id,
|
|
)
|
|
if model is not None and llm_router is not None:
|
|
return _is_model_cost_zero(model=model, llm_router=llm_router)
|
|
return False
|
|
|
|
|
|
def _resolve_request_principal(request: Request, valid_token: UserAPIKeyAuth) -> Principal:
|
|
"""Project the resolved identity into one per-request Principal, off the key
|
|
object the builder already fetched, and stamp the request network context
|
|
onto it once. X-Forwarded-For is only trusted when the operator configured
|
|
``trusted_proxy_ranges``; otherwise the direct peer is authoritative.
|
|
|
|
credential_ref and a stable subject fallback are always set off the token so
|
|
the Principal can never be anonymous, even for a keyless service-account key
|
|
with no user or alias."""
|
|
cidrs: Final = get_trusted_proxy_cidrs()
|
|
network: Final = resolve_network_context(
|
|
request,
|
|
TrustedProxyConfig(use_forwarded_for=bool(cidrs), trusted_proxy_cidrs=cidrs),
|
|
)
|
|
auth_method: Final = AuthMethod.BEARER_JWT if valid_token.jwt_claims else AuthMethod.API_KEY
|
|
return IdentityStore._principal_from_key(
|
|
valid_token,
|
|
auth_method=auth_method,
|
|
network=network,
|
|
subject_fallback=valid_token.token,
|
|
credential_ref=CredentialRef(token_id=valid_token.token),
|
|
)
|
|
|
|
|
|
async def _authorize_authenticated_request(
|
|
user_api_key_auth_obj: UserAPIKeyAuth,
|
|
request: Request,
|
|
request_data: dict,
|
|
route: str,
|
|
api_key: str,
|
|
) -> UserAPIKeyAuth | None:
|
|
"""Authorize an already-authenticated request: disabled-route check, the single
|
|
``common_checks`` gate (which also reserves budget), and end-user fallback
|
|
resolution. Returns the auth object the exception handler recovered when a check
|
|
failed but the request may proceed anyway, else ``None``.
|
|
"""
|
|
## ENSURE DISABLE ROUTE WORKS ACROSS ALL USER AUTH FLOWS ##
|
|
RouteChecks.should_call_route(route=route, valid_token=user_api_key_auth_obj, request=request)
|
|
await _normalize_claude_model(request_data, user_api_key_auth_obj, request, route)
|
|
|
|
# Single authorization point. Builder paths MUST NOT call common_checks.
|
|
# Route through the same exception handler the builder uses so
|
|
# authorization failures (ProxyException, or plain Exception from
|
|
# admin-only-route / model-access / budget checks) surface as
|
|
# ProxyException consistently with pre-refactor behavior.
|
|
try:
|
|
await _run_centralized_common_checks(
|
|
user_api_key_auth_obj=user_api_key_auth_obj,
|
|
request=request,
|
|
request_data=request_data,
|
|
route=route,
|
|
)
|
|
except Exception as e:
|
|
return await UserAPIKeyAuthExceptionHandler._handle_authentication_error(
|
|
e=e,
|
|
request=request,
|
|
request_data=request_data,
|
|
route=route,
|
|
parent_otel_span=user_api_key_auth_obj.parent_otel_span,
|
|
api_key=api_key,
|
|
resolved_identity=user_api_key_auth_obj,
|
|
)
|
|
|
|
# Defense-in-depth: ``_user_api_key_auth_builder`` has multiple early-return
|
|
# paths (no master key, /user/auth route, JWT short-circuits) that bypass
|
|
# the end-user resolution block. If those paths produced an auth obj
|
|
# without an ``end_user_id`` set, fall back to extracting from the request
|
|
# body so spend logs are still attributed correctly. Validation honours
|
|
# ``litellm.validate_end_user_id_in_db``.
|
|
if user_api_key_auth_obj.end_user_id is None:
|
|
from litellm.proxy.proxy_server import (
|
|
prisma_client,
|
|
proxy_logging_obj,
|
|
user_api_key_cache,
|
|
)
|
|
|
|
raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, _safe_get_request_headers(request))
|
|
if raw_end_user_id is not None:
|
|
resolved_end_user_id: Final = await resolve_and_validate_end_user_id(
|
|
raw_end_user_id=raw_end_user_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
parent_otel_span=user_api_key_auth_obj.parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
route=route,
|
|
)
|
|
if resolved_end_user_id is not None:
|
|
user_api_key_auth_obj.end_user_id = resolved_end_user_id
|
|
return None
|
|
|
|
|
|
def _spend_counter_redis_cache() -> RedisCache | None:
|
|
from litellm.proxy.proxy_server import spend_counter_cache
|
|
|
|
return spend_counter_cache.redis_cache
|
|
|
|
|
|
async def _prefetch_referenced_auth_objects(
|
|
valid_token: UserAPIKeyAuth,
|
|
end_user_id: str | None,
|
|
user_api_key_cache: UserApiKeyCache,
|
|
prisma_client: PrismaClient | None,
|
|
) -> None:
|
|
"""Warm every object and spend counter the checks below will read, in one MGET each (one DB query when cold).
|
|
Runs after the key's model access check so a denied request costs no more than it did before."""
|
|
bind_admission_counter_keys(valid_token, end_user_id=end_user_id or None)
|
|
await prefetch_auth_objects(
|
|
refs=AuthObjectRefs.from_token(valid_token),
|
|
user_api_key_cache=user_api_key_cache,
|
|
prisma_client=prisma_client,
|
|
)
|
|
|
|
|
|
def _seed_request_destinations(user_api_key_dict: UserAPIKeyAuth, request: Request | None = None) -> None:
|
|
"""Anchor the OTLP destinations this key or team overrides its traces to.
|
|
|
|
Called inside the ``auth`` phase span so that span reaches the tenant's account
|
|
as well, and on the request task so the ``ContextVar`` is inherited by the logging
|
|
tasks that close the LLM span. Best-effort: trace routing must never fail auth.
|
|
|
|
``request`` carries the headers, so a backend this request disabled with
|
|
``x-litellm-disable-callbacks`` resolves to no destination.
|
|
|
|
Only destinations the published fan-out can build are anchored. Anchoring one is
|
|
what tells the operator's exporter to hold that backend's spans back under
|
|
``override``, so an unbuildable one would leave the span with nowhere to go.
|
|
|
|
The ``postgres`` spans under ``auth`` close before this runs, because they are the
|
|
reads that resolve the identity being read here. They never reach the tenant's
|
|
account, and they are never withheld from the operator's backend, whichever mode
|
|
is set.
|
|
"""
|
|
try:
|
|
from litellm.integrations.otel.logger import fan_out_provider
|
|
from litellm.integrations.otel.plumbing.context import set_request_destinations
|
|
from litellm.integrations.otel.plumbing.providers import deliverable_destinations
|
|
from litellm.proxy.litellm_pre_call_utils import (
|
|
resolve_tenant_otel_destinations,
|
|
)
|
|
|
|
set_request_destinations(
|
|
deliverable_destinations(
|
|
resolve_tenant_otel_destinations(user_api_key_dict, _safe_get_request_headers(request)),
|
|
fan_out_provider(),
|
|
)
|
|
)
|
|
except Exception as exc: # noqa: BLE001 # telemetry routing is best-effort and must never break authentication
|
|
verbose_proxy_logger.debug("OTel V2: tenant destination resolution failed: %s", exc)
|
|
|
|
|
|
@tracer.wrap()
|
|
async def user_api_key_auth(
|
|
request: Request,
|
|
api_key: str = fastapi.Security(api_key_header),
|
|
azure_api_key_header: str = fastapi.Security(azure_api_key_header),
|
|
anthropic_api_key_header: str | None = fastapi.Security(anthropic_api_key_header),
|
|
google_ai_studio_api_key_header: str | None = fastapi.Security(google_ai_studio_api_key_header),
|
|
azure_apim_header: str | None = fastapi.Security(azure_apim_header),
|
|
custom_litellm_key_header: str | None = fastapi.Security(custom_litellm_key_header),
|
|
) -> UserAPIKeyAuth:
|
|
"""
|
|
Parent function to authenticate user api key / jwt token.
|
|
"""
|
|
|
|
# Create the SERVER span and stash it on request.state BEFORE reading the
|
|
# body. _read_request_body can raise ProxyException for malformed JSON;
|
|
# without this, that path leaves no span for the exception handler to
|
|
# close, and the trace never reaches the backend.
|
|
_ensure_parent_otel_span_on_request_state(request)
|
|
|
|
request_data, body_parse_exception = await _read_request_body_deferring_parse_failure(request=request)
|
|
route: Final[str] = get_request_route(request=request)
|
|
## CHECK IF ROUTE IS ALLOWED
|
|
|
|
# Run the whole auth phase inside a live ``auth`` span so the DB lookups it
|
|
# triggers (key/user/team object reads) nest under it instead of flattening
|
|
# onto the server span. No-op when OTel V2 isn't active.
|
|
with phase_span(f"auth {route}"), spend_counter_batch_scope(_spend_counter_redis_cache()):
|
|
try:
|
|
user_api_key_auth_obj: Final = await _user_api_key_auth_builder(
|
|
request=request,
|
|
api_key=api_key,
|
|
azure_api_key_header=azure_api_key_header,
|
|
anthropic_api_key_header=anthropic_api_key_header,
|
|
google_ai_studio_api_key_header=google_ai_studio_api_key_header,
|
|
azure_apim_header=azure_apim_header,
|
|
request_data=request_data,
|
|
custom_litellm_key_header=custom_litellm_key_header,
|
|
)
|
|
except Exception:
|
|
# The body was read first, so a caller who sent both a malformed body and
|
|
# a rejected key used to get the 400; the response is unchanged, and the
|
|
# auth failure is still recorded on the trace by the handler that ran.
|
|
if body_parse_exception is not None:
|
|
raise body_parse_exception
|
|
raise
|
|
user_api_key_auth_obj.budget_reservation = None
|
|
_seed_request_destinations(user_api_key_auth_obj, request)
|
|
|
|
# A body that never parsed is authenticated (so the trace carries identity
|
|
# and this ``auth`` span) but not authorized: there is no model to check it
|
|
# against, and budget reservation would increment live spend counters that
|
|
# only the endpoint's post-call path releases; the endpoint never runs, since
|
|
# the parse failure is raised below.
|
|
if body_parse_exception is None:
|
|
recovered_auth_obj: Final = await _authorize_authenticated_request(
|
|
user_api_key_auth_obj=user_api_key_auth_obj,
|
|
request=request,
|
|
request_data=request_data,
|
|
route=route,
|
|
api_key=api_key,
|
|
)
|
|
if recovered_auth_obj is not None:
|
|
return recovered_auth_obj
|
|
|
|
# Identity is now resolved. Seed it AFTER the auth span closes so the Baggage
|
|
# persists on the request task (detaching the span's context token inside the
|
|
# ``with`` would unwind a Baggage attach made within it) and every post-auth
|
|
# span — pre-call, LLM call, guardrail, spend write — inherits team/key/user.
|
|
seed_request_identity(
|
|
user_api_key_auth_obj,
|
|
model=request_data.get("model") if isinstance(request_data, dict) else None,
|
|
)
|
|
user_api_key_auth_obj.request_route = normalize_request_route(route)
|
|
|
|
if body_parse_exception is not None:
|
|
await _record_unparsable_body_failure(
|
|
user_api_key_dict=user_api_key_auth_obj,
|
|
body_parse_exception=body_parse_exception,
|
|
route=route,
|
|
)
|
|
raise body_parse_exception
|
|
|
|
# Resolve caller identity once, here at the seam, into a single per-request
|
|
# Principal projected off the key object the builder already fetched (no
|
|
# second lookup). Downstream consumers read identity off this instead of
|
|
# re-resolving it. Additive and defensive: a projection failure must never
|
|
# reject an already-authenticated request, so it is left unset on failure;
|
|
# any future consumer must treat a missing principal as deny, not allow.
|
|
try:
|
|
request.state.principal = _resolve_request_principal(request, user_api_key_auth_obj)
|
|
except Exception as e:
|
|
verbose_proxy_logger.warning("Principal projection at auth seam failed (non-fatal): %s", e)
|
|
|
|
return user_api_key_auth_obj
|
|
|
|
|
|
async def _return_user_api_key_auth_obj(
|
|
user_obj: LiteLLM_UserTable | None,
|
|
api_key: str,
|
|
parent_otel_span: Span | None,
|
|
valid_token_dict: dict,
|
|
route: str,
|
|
start_time: datetime,
|
|
user_role: LitellmUserRoles | None = None,
|
|
) -> UserAPIKeyAuth:
|
|
end_time: Final = datetime.now(timezone.utc)
|
|
|
|
asyncio.create_task(
|
|
user_api_key_service_logger_obj.async_service_success_hook(
|
|
service=ServiceTypes.AUTH,
|
|
call_type=route,
|
|
start_time=start_time,
|
|
end_time=end_time,
|
|
duration=end_time.timestamp() - start_time.timestamp(),
|
|
parent_otel_span=parent_otel_span,
|
|
)
|
|
)
|
|
|
|
retrieved_user_role: Final = user_role or _get_user_role(user_obj=user_obj) or LitellmUserRoles.INTERNAL_USER
|
|
|
|
user_api_key_kwargs: Final = {
|
|
"api_key": api_key,
|
|
"parent_otel_span": parent_otel_span,
|
|
"user_role": retrieved_user_role,
|
|
**valid_token_dict,
|
|
}
|
|
if user_obj is not None:
|
|
user_api_key_kwargs.update(
|
|
user_tpm_limit=user_obj.tpm_limit,
|
|
user_rpm_limit=user_obj.rpm_limit,
|
|
user_email=user_obj.user_email,
|
|
user_spend=getattr(user_obj, "spend", None),
|
|
user_max_budget=getattr(user_obj, "max_budget", None),
|
|
user_model_max_budget=getattr(user_obj, "model_max_budget", None),
|
|
)
|
|
if user_obj is not None and _is_user_proxy_admin(user_obj=user_obj):
|
|
user_api_key_kwargs.update(
|
|
user_role=LitellmUserRoles.PROXY_ADMIN,
|
|
)
|
|
return UserAPIKeyAuth.model_validate(user_api_key_kwargs)
|
|
else:
|
|
return UserAPIKeyAuth.model_validate(user_api_key_kwargs)
|
|
|
|
|
|
def get_api_key_from_custom_header(request: Request, custom_litellm_key_header_name: str) -> str:
|
|
"""
|
|
Get API key from custom header
|
|
|
|
Args:
|
|
request (Request): Request object
|
|
custom_litellm_key_header_name (str): Custom header name
|
|
|
|
Returns:
|
|
Optional[str]: API key
|
|
"""
|
|
api_key: str = ""
|
|
# use this as the virtual key passed to litellm proxy
|
|
custom_litellm_key_header_name = custom_litellm_key_header_name.lower()
|
|
_headers: Final = {k.lower(): v for k, v in request.headers.items()}
|
|
verbose_proxy_logger.debug(
|
|
"searching for custom_litellm_key_header_name= %s, in headers=%s",
|
|
custom_litellm_key_header_name,
|
|
_headers,
|
|
)
|
|
custom_api_key: Final = _headers.get(custom_litellm_key_header_name)
|
|
if custom_api_key:
|
|
api_key = _get_bearer_token(api_key=custom_api_key)
|
|
verbose_proxy_logger.debug(
|
|
"Found custom API key using header: %s, setting api_key=%s",
|
|
custom_litellm_key_header_name,
|
|
abbreviate_api_key(api_key),
|
|
)
|
|
else:
|
|
verbose_proxy_logger.exception(
|
|
"No LiteLLM Virtual Key pass. Please set header=%s: Bearer <api_key>", custom_litellm_key_header_name
|
|
)
|
|
return api_key
|
|
|
|
|
|
def _get_temp_budget_increase(valid_token: UserAPIKeyAuth):
|
|
valid_token_metadata: Final = valid_token.metadata
|
|
if "temp_budget_increase" in valid_token_metadata and "temp_budget_expiry" in valid_token_metadata:
|
|
expiry = datetime.fromisoformat(valid_token_metadata["temp_budget_expiry"])
|
|
if expiry.tzinfo is None:
|
|
expiry = expiry.replace(tzinfo=timezone.utc)
|
|
if expiry > datetime.now(timezone.utc):
|
|
return valid_token_metadata["temp_budget_increase"]
|
|
return None
|
|
|
|
|
|
def _update_key_budget_with_temp_budget_increase(
|
|
valid_token: UserAPIKeyAuth,
|
|
) -> UserAPIKeyAuth:
|
|
if valid_token.max_budget is None:
|
|
return valid_token
|
|
temp_budget_increase: Final = _get_temp_budget_increase(valid_token)
|
|
if not temp_budget_increase:
|
|
return valid_token
|
|
return valid_token.model_copy(update={"max_budget": valid_token.max_budget + temp_budget_increase})
|
|
|
|
|
|
async def _lookup_end_user_and_apply_budget(
|
|
valid_token: UserAPIKeyAuth,
|
|
route: str,
|
|
parent_otel_span: Span | None,
|
|
prisma_client,
|
|
user_api_key_cache,
|
|
proxy_logging_obj,
|
|
):
|
|
"""Look up end_user from DB and apply budget limits to valid_token."""
|
|
end_user_object = None
|
|
try:
|
|
end_user_object = await get_end_user_object(
|
|
end_user_id=valid_token.end_user_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
route=route,
|
|
token_end_user_max_budget=valid_token.end_user_max_budget,
|
|
)
|
|
if end_user_object is not None:
|
|
end_user_params = {
|
|
"end_user_id": valid_token.end_user_id,
|
|
"allowed_model_region": end_user_object.allowed_model_region,
|
|
}
|
|
if end_user_object.litellm_budget_table is not None:
|
|
_apply_budget_limits_to_end_user_params(
|
|
end_user_params=end_user_params,
|
|
budget_info=end_user_object.litellm_budget_table,
|
|
end_user_id=valid_token.end_user_id or "",
|
|
)
|
|
valid_token = update_valid_token_with_end_user_params(
|
|
valid_token=valid_token, end_user_params=end_user_params
|
|
)
|
|
elif litellm.max_end_user_budget_id is not None:
|
|
from litellm.proxy.auth.auth_checks import get_default_end_user_budget
|
|
|
|
default_budget: Final = await get_default_end_user_budget(
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
)
|
|
if default_budget is not None:
|
|
end_user_params = {"end_user_id": valid_token.end_user_id}
|
|
_apply_budget_limits_to_end_user_params(
|
|
end_user_params=end_user_params,
|
|
budget_info=default_budget,
|
|
end_user_id=valid_token.end_user_id or "",
|
|
)
|
|
valid_token = update_valid_token_with_end_user_params(
|
|
valid_token=valid_token, end_user_params=end_user_params
|
|
)
|
|
except Exception as e:
|
|
if isinstance(e, litellm.BudgetExceededError):
|
|
raise e
|
|
verbose_proxy_logger.debug("Unable to find user in db. Error - %s", e)
|
|
return valid_token, end_user_object
|
|
|
|
|
|
async def _enforce_key_and_fallback_model_access(
|
|
*,
|
|
valid_token: UserAPIKeyAuth,
|
|
request_data: dict,
|
|
route: str,
|
|
request: Request | None,
|
|
llm_model_list: list | None,
|
|
llm_router: Any | None,
|
|
) -> None:
|
|
"""
|
|
Key-level model allowlist and client fallbacks (same as standard auth).
|
|
Not included in common_checks — common_checks enforces team/user/project model access only.
|
|
"""
|
|
await _normalize_claude_model(request_data, valid_token, request, route)
|
|
config: Final = valid_token.config
|
|
|
|
if config != {}:
|
|
model_list: Final = config.get("model_list", [])
|
|
new_model_list: Final = model_list
|
|
verbose_proxy_logger.debug("\n new llm router model list %s", new_model_list)
|
|
elif isinstance(valid_token.models, list) and "all-team-models" in valid_token.models:
|
|
pass
|
|
else:
|
|
model: Final = _get_model_from_request_context(
|
|
request_data=request_data,
|
|
route=route,
|
|
request=request,
|
|
llm_router=llm_router,
|
|
team_id=valid_token.team_id,
|
|
)
|
|
|
|
if model is not None:
|
|
await can_key_call_model(
|
|
model=model,
|
|
llm_model_list=llm_model_list,
|
|
valid_token=valid_token,
|
|
llm_router=llm_router,
|
|
)
|
|
|
|
fallback_names: Final = tuple(
|
|
name
|
|
for target in iter_request_fallback_targets(request_data)
|
|
if (name := _fallback_target_model_name(target)) is not None
|
|
)
|
|
|
|
for _name in dict.fromkeys(fallback_names): # dedupe, preserve order
|
|
await can_key_call_model(
|
|
model=_name,
|
|
llm_model_list=llm_model_list,
|
|
valid_token=valid_token,
|
|
llm_router=llm_router,
|
|
)
|
|
await is_valid_fallback_model(
|
|
model=_name,
|
|
llm_router=llm_router,
|
|
user_model=None,
|
|
)
|
|
|
|
|
|
def _fallback_target_model_name(target: object) -> str | None:
|
|
if isinstance(target, str):
|
|
return target
|
|
if isinstance(target, dict):
|
|
model: Final = target.get("model")
|
|
if isinstance(model, str):
|
|
return model
|
|
return None
|
|
|
|
|
|
async def _run_post_custom_auth_checks(
|
|
valid_token: UserAPIKeyAuth,
|
|
request: Request,
|
|
request_data: dict,
|
|
route: str,
|
|
parent_otel_span: Span | None,
|
|
) -> UserAPIKeyAuth:
|
|
from litellm.proxy.proxy_server import (
|
|
general_settings,
|
|
llm_model_list,
|
|
llm_router,
|
|
model_max_budget_limiter,
|
|
prisma_client,
|
|
proxy_logging_obj,
|
|
user_api_key_cache,
|
|
)
|
|
|
|
# 1. Look up end_user object from DB if end_user_id is set
|
|
end_user_object = None
|
|
if valid_token.end_user_id is not None:
|
|
valid_token, end_user_object = await _lookup_end_user_and_apply_budget(
|
|
valid_token=valid_token,
|
|
route=route,
|
|
parent_otel_span=parent_otel_span,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
# common_checks() enforces the end-user budget, but the centralized
|
|
# gate skips it for custom-auth deployments unless
|
|
# custom_auth_run_common_checks is set. Enforce it here on that path
|
|
# so an over-budget end user can't keep making requests.
|
|
if end_user_object is not None and not general_settings.get("custom_auth_run_common_checks", False):
|
|
await _check_end_user_budget(end_user_obj=end_user_object, route=route)
|
|
|
|
# 2. Check token expiry
|
|
if valid_token.expires is not None:
|
|
current_time: Final = datetime.now(timezone.utc)
|
|
if isinstance(valid_token.expires, datetime):
|
|
expiry_time = valid_token.expires
|
|
else:
|
|
expiry_time = datetime.fromisoformat(valid_token.expires)
|
|
if expiry_time.tzinfo is None or expiry_time.tzinfo.utcoffset(expiry_time) is None:
|
|
expiry_time = expiry_time.replace(tzinfo=timezone.utc)
|
|
if expiry_time < current_time:
|
|
raise ProxyException(
|
|
message=f"Authentication Error - Expired Key. Key Expiry time {expiry_time} and current time {current_time}",
|
|
type=ProxyErrorTypes.expired_key,
|
|
code=status.HTTP_401_UNAUTHORIZED,
|
|
param=(abbreviate_api_key(api_key=valid_token.token) if valid_token.token else ""),
|
|
)
|
|
|
|
if general_settings.get("custom_auth_run_common_checks", False):
|
|
await _enforce_key_and_fallback_model_access(
|
|
valid_token=valid_token,
|
|
request_data=request_data,
|
|
route=route,
|
|
request=request,
|
|
llm_model_list=llm_model_list,
|
|
llm_router=llm_router,
|
|
)
|
|
|
|
current_model = _get_model_from_request_context(
|
|
request_data=request_data,
|
|
route=route,
|
|
request=request,
|
|
llm_router=llm_router,
|
|
team_id=valid_token.team_id,
|
|
)
|
|
current_models = _get_model_names_for_budget_checks(model=current_model)
|
|
|
|
# A zero-cost model cannot move any counter, so refusing it means refusing on
|
|
# spend some other model accrued. The JWT and virtual-key paths already skip
|
|
# every budget check for these; this path did not, so the same request could
|
|
# be refused under custom auth and served under the other two.
|
|
skip_budget_checks: Final = (
|
|
_is_model_cost_zero(model=current_model, llm_router=llm_router)
|
|
if current_model is not None and llm_router is not None
|
|
else False
|
|
)
|
|
|
|
# 3. Check key-level model_max_budget
|
|
max_budget_per_model: Final = valid_token.model_max_budget
|
|
if (
|
|
not skip_budget_checks
|
|
and max_budget_per_model is not None
|
|
and isinstance(max_budget_per_model, dict)
|
|
and len(max_budget_per_model) > 0
|
|
and current_models
|
|
and valid_token.token is not None
|
|
):
|
|
for model_name in current_models:
|
|
await _check_key_model_budget_with_fallback(
|
|
valid_token=valid_token,
|
|
model_max_budget_limiter=model_max_budget_limiter,
|
|
model_name=model_name,
|
|
request_data=request_data,
|
|
request=request,
|
|
llm_model_list=llm_model_list,
|
|
llm_router=llm_router,
|
|
)
|
|
|
|
# Recompute after a potential budget-fallback rewrite so
|
|
# the end-user check below validates the final model
|
|
current_model = _get_model_from_request_context(
|
|
request_data=request_data,
|
|
route=route,
|
|
request=request,
|
|
llm_router=llm_router,
|
|
team_id=valid_token.team_id,
|
|
)
|
|
current_models = _get_model_names_for_budget_checks(model=current_model)
|
|
|
|
# 3b. Attach and check the internal user's model_max_budget.
|
|
# Custom auth builds its own token, so unlike the main path nothing has
|
|
# loaded the user row yet. The attach is unconditional because the post-call
|
|
# spend hook reads this field off the token: gating it on the same condition
|
|
# as enforcement would leave the user's counter uncharged whenever this
|
|
# request was not itself enforceable, so its spend would go untracked.
|
|
user_budget: Final = await _read_user_model_max_budget(
|
|
user_id=valid_token.user_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
parent_otel_span=parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
valid_token.user_model_max_budget = user_budget # rebind-ok: the spend hook reads it off this token
|
|
if not skip_budget_checks and current_models:
|
|
await _check_user_model_budget(
|
|
valid_token=valid_token,
|
|
model_max_budget_limiter=model_max_budget_limiter,
|
|
models=current_models,
|
|
)
|
|
|
|
# 4. Check end-user model_max_budget
|
|
end_user_mmb: Final = valid_token.end_user_model_max_budget
|
|
if (
|
|
not skip_budget_checks
|
|
and end_user_mmb is not None
|
|
and isinstance(end_user_mmb, dict)
|
|
and len(end_user_mmb) > 0
|
|
and current_models
|
|
and valid_token.end_user_id is not None
|
|
):
|
|
for model_name in current_models:
|
|
await model_max_budget_limiter.is_end_user_within_model_budget(
|
|
end_user_id=valid_token.end_user_id,
|
|
end_user_model_max_budget=end_user_mmb,
|
|
model=model_name,
|
|
)
|
|
|
|
# team / user / end_user / project context objects are fetched by
|
|
# the centralized common_checks gate in user_api_key_auth after
|
|
# this helper returns. Keep only the project fetch here because it
|
|
# mutates the token (project_metadata / project_alias).
|
|
if valid_token.project_id is not None:
|
|
_project_obj: Final = await get_project_object(
|
|
project_id=valid_token.project_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
if _project_obj is not None:
|
|
valid_token.project_metadata = _project_obj.metadata
|
|
valid_token.project_alias = _project_obj.project_alias
|
|
|
|
return valid_token
|