diff --git a/litellm/__init__.py b/litellm/__init__.py index 55821012df9..3f8c742c5a2 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -211,6 +211,9 @@ filter_invalid_headers: Optional[bool] = False add_user_information_to_llm_headers: Optional[bool] = ( None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers ) +overwrite_user_with_key_hash: bool = ( + False # force the outgoing `user` param to the hashed api key, so providers see a stable, tamper-proof id +) store_audit_logs = False # Enterprise feature, allow users to see audit logs skip_system_message_in_guardrail: bool = False skip_tool_message_in_guardrail: bool = False diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 444e5ba0731..98efadc10a8 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2616,6 +2616,17 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob # key off. Server-only and stripped from validated input for the same reason as the marker # above: a forged entry would let a caller pick which team's rpm bucket it is charged against. mcp_source_team_rpm_limits: dict[str, dict[str, int]] | None = Field(default=None, exclude=True) + via_virtual_key: bool = Field( + default=False, + exclude=True, + description=( + "Server-only marker set exclusively by the DB virtual-key and master-key auth paths via " + "post-construction assignment. Stripped from validated input so custom auth handlers, JWT " + "claims, or key metadata cannot forge it. Gates overwrite_user_with_key_hash stamping: only " + "a credential the proxy itself validated as a key may be forwarded as the provider-facing " + "user id." + ), + ) budget_reservation: Optional[Dict[str, Any]] = Field(default=None, exclude=True) budget_throttle_pct: Optional[float] = Field(default=None, exclude=True) user: Optional[Any] = None # Expanded user object when expand=user is used @@ -2641,6 +2652,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob # kwargs, model_validate, a JWT/key claim splat) so it can never be forged from caller data. values.pop("mcp_admitted_user_subject", None) values.pop("mcp_source_team_rpm_limits", None) + values.pop("via_virtual_key", None) if values.get("api_key") is not None: values.update({"token": cls._safe_hash_litellm_api_key(values.get("api_key"))}) if isinstance(values.get("api_key"), str): diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 83a8a69511b..d776a626251 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1497,6 +1497,13 @@ async def _user_api_key_auth_builder( 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 @@ -1614,6 +1621,7 @@ async def _user_api_key_auth_builder( _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 @@ -2021,7 +2029,7 @@ async def _user_api_key_auth_builder( # No token was found when looking up in the DB raise Exception("Invalid proxy server token passed") if valid_token_dict is not None: - return await _return_user_api_key_auth_obj( + virtual_key_auth_obj = await _return_user_api_key_auth_obj( user_obj=user_obj, api_key=api_key, parent_otel_span=parent_otel_span, @@ -2029,6 +2037,8 @@ async def _user_api_key_auth_builder( 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, diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 9d9ef28ec9b..a4cc4a62009 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -13,7 +13,7 @@ from starlette.datastructures import Headers import litellm from litellm._logging import verbose_logger, verbose_proxy_logger from litellm._service_logger import ServiceLogging -from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY +from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS, PRE_CALL_EXECUTED_GUARDRAILS_KEY from litellm.litellm_core_utils.credential_accessor import CredentialAccessor from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( iter_client_callback_metadata_dicts, @@ -48,6 +48,24 @@ _EXPLICIT_SESSION_HEADERS = frozenset({"x-litellm-trace-id", "x-litellm-session- # Session-id values must be non-empty strings of alphanumerics, hyphens, or underscores # (covers UUIDs and most common session-id formats). _SESSION_ID_VALUE_RE = re.compile(r"^[a-zA-Z0-9_\-]{8,}$") + +_SHA256_HEX_RE = re.compile(r"^[0-9a-f]{64}$") + + +def _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None: + """Only proxy-validated keys are stamped, proven by the unforgeable + via_virtual_key marker AND a known non-secret shape: the sha256 hex digest + UserAPIKeyAuth stores virtual keys in, or the master key's stable alias. + Custom-auth credentials arrive raw (never forward auth material) and hashed + JWTs rotate on re-issue (useless as a stable ban id), so both are skipped.""" + api_key = user_api_key_dict.api_key + if not user_api_key_dict.via_virtual_key or api_key is None: + return None + if api_key == LITELLM_PROXY_MASTER_KEY_ALIAS or _SHA256_HEX_RE.fullmatch(api_key): + return api_key + return None + + _ANTHROPIC_SESSION_ID_VALUE_RE = re.compile(r"^[a-zA-Z0-9_\-]+$") @@ -1447,6 +1465,11 @@ async def add_litellm_data_to_request( if "user" not in data: data["user"] = user + if litellm.overwrite_user_with_key_hash is True: + stampable_hash = _stampable_key_hash(user_api_key_dict) + if stampable_hash is not None: + data["user"] = stampable_hash + data["secret_fields"] = SecretFields(raw_headers=_raw_headers) ## Dynamic api version (Azure OpenAI endpoints) ## diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 2c1948adca1..a2445ecb975 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -1290,6 +1290,250 @@ async def test_scim_deactivated_user_key_is_rejected(): setattr(_proxy_server_mod, attr, val) +@pytest.mark.asyncio +async def test_cached_proxy_admin_key_sets_via_virtual_key_marker(): + """Cached PROXY_ADMIN auth objects early-return before the marked DB and + master-key returns, and cache serialization drops the exclude=True marker; + the cache-hit boundary must restore it or cached admin traffic silently + bypasses overwrite_user_with_key_hash stamping.""" + from fastapi import Request + from starlette.datastructures import URL + + from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.proxy_server import hash_token + + api_key = "sk-cached-admin-marker-test" + hashed_key = hash_token(api_key) + + cached_token = UserAPIKeyAuth( + api_key=api_key, + token=hashed_key, + user_id="cached-admin-user", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + assert cached_token.via_virtual_key is False + + mock_cache = AsyncMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.delete_cache = MagicMock() + + mock_proxy_logging_obj = MagicMock() + mock_proxy_logging_obj.internal_usage_cache = MagicMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = ( + AsyncMock() + ) + mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + + import litellm.proxy.proxy_server as _proxy_server_mod + + _attrs_to_set = { + "prisma_client": MagicMock(), + "user_api_key_cache": mock_cache, + "proxy_logging_obj": mock_proxy_logging_obj, + "master_key": "sk-master-key", + "general_settings": {}, + "llm_model_list": [], + "llm_router": None, + "open_telemetry_logger": None, + "model_max_budget_limiter": MagicMock(), + "user_custom_auth": None, + "jwt_handler": None, + "litellm_proxy_admin_name": "admin", + } + _original_values = { + attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set + } + try: + for attr, val in _attrs_to_set.items(): + setattr(_proxy_server_mod, attr, val) + + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + + with patch( + "litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key", + new_callable=AsyncMock, + return_value=cached_token, + ): + result = await _user_api_key_auth_builder( + request=request, + api_key=f"Bearer {api_key}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={}, + ) + + assert isinstance(result, UserAPIKeyAuth) + assert result.user_role == LitellmUserRoles.PROXY_ADMIN + assert result.via_virtual_key is True + assert result.api_key == hashed_key + finally: + for attr, val in _original_values.items(): + setattr(_proxy_server_mod, attr, val) + + +@pytest.mark.asyncio +async def test_master_key_auth_sets_via_virtual_key_marker(): + """Master-key requests must also be stamped by overwrite_user_with_key_hash; + the auth path substitutes the stable alias for api_key and must mark the + result as proxy-validated.""" + from fastapi import Request + from starlette.datastructures import URL + + from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS + from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + + master_key = "sk-master-key" + + mock_cache = AsyncMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.delete_cache = MagicMock() + + mock_proxy_logging_obj = MagicMock() + mock_proxy_logging_obj.internal_usage_cache = MagicMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = ( + AsyncMock() + ) + mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + + import litellm.proxy.proxy_server as _proxy_server_mod + + _attrs_to_set = { + "prisma_client": MagicMock(), + "user_api_key_cache": mock_cache, + "proxy_logging_obj": mock_proxy_logging_obj, + "master_key": master_key, + "general_settings": {}, + "llm_model_list": [], + "llm_router": None, + "open_telemetry_logger": None, + "model_max_budget_limiter": MagicMock(), + "user_custom_auth": None, + "jwt_handler": None, + "litellm_proxy_admin_name": "admin", + } + _original_values = { + attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set + } + try: + for attr, val in _attrs_to_set.items(): + setattr(_proxy_server_mod, attr, val) + + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + + result = await _user_api_key_auth_builder( + request=request, + api_key=f"Bearer {master_key}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={}, + ) + + assert isinstance(result, UserAPIKeyAuth) + assert result.via_virtual_key is True + assert result.api_key == LITELLM_PROXY_MASTER_KEY_ALIAS + finally: + for attr, val in _original_values.items(): + setattr(_proxy_server_mod, attr, val) + + +@pytest.mark.asyncio +async def test_db_virtual_key_auth_sets_via_virtual_key_marker(): + """via_virtual_key gates overwrite_user_with_key_hash stamping and is + forge-stripped from validated input, so the DB auth path setting it by + post-construction assignment is the only thing that turns stamping on.""" + from fastapi import Request + from starlette.datastructures import URL + + from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.proxy_server import hash_token + + api_key = "sk-via-virtual-key-marker-test" + hashed_key = hash_token(api_key) + + valid_token = UserAPIKeyAuth( + api_key=api_key, + token=hashed_key, + user_id="marker-test-user", + ) + + mock_cache = AsyncMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.delete_cache = MagicMock() + + mock_proxy_logging_obj = MagicMock() + mock_proxy_logging_obj.internal_usage_cache = MagicMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = ( + AsyncMock() + ) + mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + + mock_prisma_client = MagicMock() + + import litellm.proxy.proxy_server as _proxy_server_mod + + _attrs_to_set = { + "prisma_client": mock_prisma_client, + "user_api_key_cache": mock_cache, + "proxy_logging_obj": mock_proxy_logging_obj, + "master_key": "sk-master-key", + "general_settings": {}, + "llm_model_list": [], + "llm_router": None, + "open_telemetry_logger": None, + "model_max_budget_limiter": MagicMock(), + "user_custom_auth": None, + "jwt_handler": None, + "litellm_proxy_admin_name": "admin", + } + _original_values = { + attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set + } + try: + for attr, val in _attrs_to_set.items(): + setattr(_proxy_server_mod, attr, val) + + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + + with ( + patch( + "litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key", + new_callable=AsyncMock, + return_value=valid_token, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.get_user_object", + new_callable=AsyncMock, + return_value=None, + ), + ): + result = await _user_api_key_auth_builder( + request=request, + api_key=f"Bearer {api_key}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={}, + ) + + assert isinstance(result, UserAPIKeyAuth) + assert result.via_virtual_key is True + assert result.api_key == hashed_key + finally: + for attr, val in _original_values.items(): + setattr(_proxy_server_mod, attr, val) + + @pytest.mark.asyncio async def test_return_user_api_key_auth_obj_user_spend_and_budget(): """ diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 8bee7e9f33b..1437899f561 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -5226,3 +5226,198 @@ async def test_add_litellm_data_to_request_unions_metadata_tags_with_header_tags tags = updated["litellm_metadata"]["tags"] assert "header-tag" in tags assert "body-tag" in tags + + +def _make_chat_request_mock() -> MagicMock: + return _make_request_mock("/v1/chat/completions", {"Content-Type": "application/json"}) + + +@pytest.mark.asyncio +async def test_overwrite_user_with_key_hash_clobbers_caller_supplied_user(monkeypatch): + """The flag exists so providers can ban by a tamper-proof id; a caller-chosen + `user` must never survive, and the raw sk- key must never be forwarded.""" + from litellm.proxy._types import hash_token + + monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True) + + raw_key = "sk-overwrite-user-test-1234" + user_api_key_dict = UserAPIKeyAuth(api_key=raw_key) + user_api_key_dict.via_virtual_key = True + data = {"model": "gpt-4o", "user": "attacker-chosen-id"} + + updated_data = await add_litellm_data_to_request( + data=data, + request=_make_chat_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated_data["user"] == hash_token(raw_key) + assert updated_data["user"] != "attacker-chosen-id" + assert raw_key not in updated_data["user"] + + +@pytest.mark.asyncio +async def test_overwrite_user_with_key_hash_sets_user_when_absent(monkeypatch): + from litellm.proxy._types import hash_token + + monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True) + + raw_key = "sk-overwrite-user-test-5678" + user_api_key_dict = UserAPIKeyAuth(api_key=raw_key) + user_api_key_dict.via_virtual_key = True + data = {"model": "gpt-4o"} + + updated_data = await add_litellm_data_to_request( + data=data, + request=_make_chat_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated_data["user"] == hash_token(raw_key) + + +@pytest.mark.asyncio +async def test_overwrite_user_with_key_hash_disabled_preserves_caller_user(): + assert litellm.overwrite_user_with_key_hash is False + + user_api_key_dict = UserAPIKeyAuth(api_key="sk-overwrite-user-test-9999") + user_api_key_dict.via_virtual_key = True + data = {"model": "gpt-4o", "user": "caller-chosen-id"} + + updated_data = await add_litellm_data_to_request( + data=data, + request=_make_chat_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated_data["user"] == "caller-chosen-id" + + +@pytest.mark.asyncio +async def test_overwrite_user_with_key_hash_skips_custom_auth_credential(monkeypatch): + """Custom-auth credentials are not sk-prefixed or JWTs, so UserAPIKeyAuth stores + them raw; the stamp must skip them entirely so auth material never leaks.""" + monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True) + + raw_credential = "my-custom-auth-credential-abc123" + user_api_key_dict = UserAPIKeyAuth(api_key=raw_credential) + assert user_api_key_dict.api_key == raw_credential + + updated_data = await add_litellm_data_to_request( + data={"model": "gpt-4o", "user": "caller-chosen-id"}, + request=_make_chat_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated_data["user"] == "caller-chosen-id" + + +@pytest.mark.asyncio +async def test_overwrite_user_with_key_hash_skips_jwt_auth(monkeypatch): + """A hashed JWT rotates on every token re-issue, so it is useless as a stable + ban id; JWT-authenticated requests are not stamped.""" + from litellm.proxy._types import hash_token + + monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True) + + hashed_jwt = f"hashed-jwt-{hash_token('some-jwt-token')}" + user_api_key_dict = UserAPIKeyAuth(api_key=hashed_jwt) + + updated_data = await add_litellm_data_to_request( + data={"model": "gpt-4o", "user": "caller-chosen-id"}, + request=_make_chat_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated_data["user"] == "caller-chosen-id" + + +@pytest.mark.asyncio +async def test_overwrite_user_with_key_hash_skips_hex_shaped_custom_credential(monkeypatch): + """A custom-auth credential that happens to be 64 hex chars is indistinguishable + from a key hash by shape alone; only the server-set via_virtual_key marker may + authorize stamping, so this raw credential must never be forwarded.""" + monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True) + + hex_shaped_credential = "a" * 64 + user_api_key_dict = UserAPIKeyAuth(api_key=hex_shaped_credential) + assert user_api_key_dict.api_key == hex_shaped_credential + assert user_api_key_dict.via_virtual_key is False + + updated_data = await add_litellm_data_to_request( + data={"model": "gpt-4o", "user": "caller-chosen-id"}, + request=_make_chat_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated_data["user"] == "caller-chosen-id" + + +def test_via_virtual_key_cannot_be_forged_from_validated_input(): + from_kwargs = UserAPIKeyAuth(api_key="b" * 64, via_virtual_key=True) + assert from_kwargs.via_virtual_key is False + + from_dict = UserAPIKeyAuth.model_validate({"api_key": "b" * 64, "via_virtual_key": True}) + assert from_dict.via_virtual_key is False + + +@pytest.mark.asyncio +async def test_overwrite_user_with_key_hash_stamps_master_key_alias(monkeypatch): + """Master-key requests carry the stable alias instead of a hash (so the master + key never propagates anywhere); the alias is the stampable id for them.""" + from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS + + monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True) + + user_api_key_dict = UserAPIKeyAuth(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS) + user_api_key_dict.via_virtual_key = True + + updated_data = await add_litellm_data_to_request( + data={"model": "gpt-4o", "user": "attacker-chosen-id"}, + request=_make_chat_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated_data["user"] == LITELLM_PROXY_MASTER_KEY_ALIAS + + +@pytest.mark.asyncio +async def test_overwrite_user_with_key_hash_rejects_alias_without_marker(monkeypatch): + from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS + + monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True) + + user_api_key_dict = UserAPIKeyAuth(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS) + assert user_api_key_dict.via_virtual_key is False + + updated_data = await add_litellm_data_to_request( + data={"model": "gpt-4o", "user": "caller-chosen-id"}, + request=_make_chat_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated_data["user"] == "caller-chosen-id"