Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 12 additions & 1 deletion litellm/litellm_core_utils/realtime_streaming.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,9 @@
from litellm._logging import redact_internal_details_from_client_message, verbose_logger
from litellm.constants import REALTIME_SESSION_FAILURE_LOGGED_KEY, REALTIME_SESSION_SUCCESS_LOGGED_KEY
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig, RealtimeBackend
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.llms.openai import (
OpenAIRealtimeEvents,
OpenAIRealtimeOutputItemDone,
Expand Down Expand Up @@ -123,6 +125,10 @@ def patch_outgoing_session(self, session: dict) -> dict: ...
]


def _as_user_api_key_auth(user_api_key_dict: object) -> UserAPIKeyAuth | None:
return user_api_key_dict if isinstance(user_api_key_dict, UserAPIKeyAuth) else None


class RealTimeStreaming:
def __init__(
self,
Expand Down Expand Up @@ -852,7 +858,12 @@ async def run_realtime_guardrails(
try:
await callback.apply_guardrail(
inputs={"texts": [transcript], "images": []},
request_data={"user_api_key_dict": self.user_api_key_dict},
request_data={
"user_api_key_dict": self.user_api_key_dict,
"litellm_metadata": BaseTranslation.transform_user_api_key_dict_to_metadata(
_as_user_api_key_auth(self.user_api_key_dict)
),
},
input_type="request",
)
except Exception as e:
Expand Down
41 changes: 8 additions & 33 deletions litellm/llms/base_llm/guardrail_translation/base_translation.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,43 +81,18 @@ def post_call_hook_response(self, response: object) -> object:

@staticmethod
def transform_user_api_key_dict_to_metadata(
user_api_key_dict: Any | None,
user_api_key_dict: "UserAPIKeyAuth | None",
) -> dict[str, object]:
"""
Transform user_api_key_dict to a metadata dict with prefixed keys.

Converts keys like 'user_id' to 'user_api_key_user_id' to clearly indicate
the source of the metadata.

Args:
user_api_key_dict: UserAPIKeyAuth object or dict with user information

Returns:
Dict with keys prefixed with 'user_api_key_'
"""
"""The authenticated key's identity as prefixed metadata, an allowlist safe to hand to guardrail vendors."""
if user_api_key_dict is None:
return {}
# Lazy: `import litellm` loads this module before litellm.Router exists, and litellm_pre_call_utils imports it
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup

# Convert to dict if it's a Pydantic object
user_dict = user_api_key_dict.model_dump() if hasattr(user_api_key_dict, "model_dump") else user_api_key_dict

if not isinstance(user_dict, dict):
return {}

# Transform keys to be prefixed with 'user_api_key_'
transformed: Final[dict[str, object]] = {}
for key, value in user_dict.items():
# Skip None values and internal fields
if value is None or key.startswith("_"):
continue

# If key already has the prefix, use as-is, otherwise add prefix
if key.startswith("user_api_key_"):
transformed[key] = value
else:
transformed[f"user_api_key_{key}"] = value

return transformed
return {
**LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(user_api_key_dict),
"user_api_key_key_alias": user_api_key_dict.key_alias,
}

@staticmethod
def merge_user_api_key_metadata_into_request(
Expand Down
6 changes: 5 additions & 1 deletion litellm/llms/pass_through/guardrail_translation/handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,10 @@
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging

_PROXY_OWNED_PAYLOAD_KEYS: Final = frozenset(
{"metadata", "litellm_metadata", "litellm_logging_obj", "proxy_server_request"}
)


class PassThroughEndpointHandler(BaseTranslation):
"""
Expand Down Expand Up @@ -80,7 +84,7 @@ def _extract_text_for_guardrail(
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps

payload_to_check: Final = {
k: v for k, v in data.items() if not k.startswith("_") and k not in ("metadata", "litellm_logging_obj")
k: v for k, v in data.items() if not k.startswith("_") and k not in _PROXY_OWNED_PAYLOAD_KEYS
}
verbose_proxy_logger.debug("PassThroughEndpointHandler: Using full payload for guardrail")
return safe_dumps(payload_to_check)
Expand Down
8 changes: 7 additions & 1 deletion litellm/proxy/guardrails/guardrail_endpoints.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@
encrypt_guardrail_litellm_params,
)
from litellm.proxy.guardrails.usage_endpoints import router as guardrails_usage_router
from litellm.proxy.litellm_pre_call_utils import caller_metadata_with_authenticated_identity
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import GuardrailsRepository
Expand Down Expand Up @@ -2425,9 +2426,14 @@ async def apply_guardrail(
if litellm_logging_obj is not None:
_patch_logging_obj_for_guardrail(litellm_logging_obj, request)

processed_metadata: Final = data.get("metadata")
inbound_headers: Final = processed_metadata.get("headers") if isinstance(processed_metadata, dict) else None
request_data: Final[dict] = {
**({"messages": request.messages} if request.messages is not None else {}),
**({"metadata": request.metadata} if request.metadata is not None else {}),
"metadata": {
**caller_metadata_with_authenticated_identity(request.metadata, user_api_key_dict),
**({"headers": inbound_headers} if inbound_headers is not None else {}),
},
}
_input_type: Final = _resolve_guardrail_input_type(active_guardrail, request.input_type)
guardrailed_inputs: Final = await active_guardrail.apply_guardrail(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -292,9 +292,8 @@ def _extract_user_api_key_metadata(self, request_data: dict) -> GenericGuardrail
if value is not None:
result_metadata[field_name] = value

# handle user_api_key_token = user_api_key_hash
if metadata_dict.get("user_api_key_token") is not None:
result_metadata["user_api_key_hash"] = metadata_dict.get("user_api_key_token")
if litellm_metadata.get("user_api_key_token") is not None and "user_api_key_hash" not in result_metadata:
result_metadata["user_api_key_hash"] = litellm_metadata["user_api_key_token"]

verbose_proxy_logger.debug(
"Generic Guardrail API: Extracted user metadata: %s",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1810,16 +1810,6 @@ async def apply_guardrail(
call_id,
_mcp_tool,
)
elif not request_data and logging_obj is None and input_type == "request":
# Direct /apply_guardrail endpoint — empty request_data, no
# logging_obj. Existing behavior: synthesize UUID.
call_id = str(uuid.uuid4())
request_data["litellm_call_id"] = call_id
verbose_proxy_logger.warning(
"PANW Prisma AIRS: litellm_call_id missing from empty "
"request_data, synthesized %s (direct /apply_guardrail?)",
call_id,
)
else:
call_id = str(uuid.uuid4())
request_data["litellm_call_id"] = call_id
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,9 @@
from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket
from litellm.llms import get_guardrail_translation_mapping, load_guardrail_translation_mappings
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation, StreamingScanKey
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import (
MCP_GUARDRAIL_CALL_TYPES,
Expand All @@ -40,10 +42,6 @@
# Imported lazily at runtime (inside the streaming hook) to avoid a
# module-level cyclic import with litellm.integrations.custom_guardrail.
from litellm.integrations.custom_guardrail import ModifyResponseException
from litellm.llms.base_llm.guardrail_translation.base_translation import (
BaseTranslation,
StreamingScanKey,
)

# Call types that stream JSON-RPC events (A2A); guardrail HTTPException is emitted as in-stream error
A2A_CALL_TYPES: Final = (CallTypes.asend_message, CallTypes.send_message)
Expand All @@ -64,7 +62,7 @@ def process_output_response(self) -> "Callable[..., Awaitable[object]]": ...
def process_output_streaming_response(self) -> "Callable[..., Awaitable[object]]": ...

@property
def get_streaming_scan_key(self) -> "Callable[[Sequence[object]], StreamingScanKey | None]": ...
def get_streaming_scan_key(self) -> Callable[[Sequence[object]], StreamingScanKey | None]: ...

@property
def released_stream_as_ended(self) -> "Callable[[Sequence[object]], tuple[object, ...]]": ...
Expand All @@ -82,7 +80,7 @@ def _as_endpoint_translation(translation: _EndpointTranslation) -> _EndpointTran

def resolve_endpoint_translation(
user_api_key_dict: UserAPIKeyAuth, first_response_item: object | None
) -> "tuple[str, BaseTranslation] | None":
) -> tuple[str, BaseTranslation] | None:
"""
Resolve the endpoint guardrail translation for a streamed response: the
request route wins, falling back to inferring the call type from the first
Expand Down Expand Up @@ -131,7 +129,7 @@ def _recorded_guardrail_information(request_data: _RequestData) -> tuple[Standar
)


def _is_redundant_scan(scan_key: "StreamingScanKey | None", last_scan_key: "StreamingScanKey | None") -> bool:
def _is_redundant_scan(scan_key: StreamingScanKey | None, last_scan_key: StreamingScanKey | None) -> bool:
if scan_key is None:
return False
return scan_key == last_scan_key or scan_key.has_nothing_to_scan
Expand Down Expand Up @@ -178,16 +176,19 @@ def _a2a_jsonrpc_error_chunk(exc: HTTPException, request_id: str | None) -> Mapp
}


def _ensure_litellm_metadata(data: dict, user_api_key_dict: UserAPIKeyAuth) -> None:
"""Populate data['litellm_metadata'] from user_api_key_dict if absent."""
if "litellm_metadata" not in data:
from litellm.llms.base_llm.guardrail_translation.base_translation import (
BaseTranslation,
)
_PROXY_ENRICHED_IDENTITY_FIELDS: Final = frozenset({"user_api_key_auth_metadata"})


user_metadata: Final = BaseTranslation.transform_user_api_key_dict_to_metadata(user_api_key_dict)
if user_metadata:
data["litellm_metadata"] = user_metadata
def _apply_authenticated_identity_to_litellm_metadata(data: dict, user_api_key_dict: UserAPIKeyAuth) -> None:
existing: Final = data.get("litellm_metadata")
if isinstance(existing, dict):
identity: Final = LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(user_api_key_dict)
existing.update({key: value for key, value in identity.items() if key not in _PROXY_ENRICHED_IDENTITY_FIELDS})
existing.pop("user_api_key_token", None)
return
user_metadata: Final = BaseTranslation.transform_user_api_key_dict_to_metadata(user_api_key_dict)
if user_metadata:
data["litellm_metadata"] = user_metadata


class UnifiedLLMGuardrails(CustomLogger):
Expand Down Expand Up @@ -249,7 +250,7 @@ async def async_pre_call_hook(

endpoint_translation: Final = _as_endpoint_translation(mappings[CallTypes(call_type)]())

_ensure_litellm_metadata(data, user_api_key_dict)
_apply_authenticated_identity_to_litellm_metadata(data, user_api_key_dict)

data = await endpoint_translation.process_input_messages(
data=data,
Expand Down Expand Up @@ -295,7 +296,7 @@ async def async_moderation_hook(

endpoint_translation: Final = _as_endpoint_translation(mappings[CallTypes(call_type)]())

_ensure_litellm_metadata(data, user_api_key_dict)
_apply_authenticated_identity_to_litellm_metadata(data, user_api_key_dict)

return await endpoint_translation.process_input_messages(
data=data,
Expand Down Expand Up @@ -428,7 +429,7 @@ async def handle_streaming_block(
@staticmethod
def _resolve_transform_call_type(
user_api_key_dict: UserAPIKeyAuth,
mappings: Mapping[CallTypes, type["BaseTranslation"]],
mappings: Mapping[CallTypes, type[BaseTranslation]],
) -> str | None:
"""Resolve the call type for the incremental_diff path, or None if the
route is unresolvable / unsupported.
Expand Down Expand Up @@ -694,7 +695,7 @@ async def _run_incremental_transform_stream(
call_type: str,
sampling_rate: int,
end_of_stream_only: bool,
mappings: Mapping[CallTypes, type["BaseTranslation"]],
mappings: Mapping[CallTypes, type[BaseTranslation]],
) -> AsyncGenerator[object, None]:
"""Emit guardrail text transformations as new deltas on the stream.

Expand Down
Loading
Loading