mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/litellm-logs-ui-lag-0ca4b8
This commit is contained in:
commit
2a50b3a087
36 changed files with 1780 additions and 1213 deletions
|
|
@ -35,6 +35,7 @@ from litellm.responses.sse_output_recovery import (
|
|||
record_output_item_chunk,
|
||||
record_output_text_chunk,
|
||||
)
|
||||
from litellm.responses.utils import normalize_responses_api_stream_options
|
||||
from litellm.types.llms.openai import (
|
||||
ChatCompletionAnnotation,
|
||||
ChatCompletionReasoningItem,
|
||||
|
|
@ -320,6 +321,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
responses_api_request["tool_choice"] = ( # type: ignore[assignment]
|
||||
self._normalize_tool_choice_for_responses_api(value)
|
||||
)
|
||||
elif key == "stream_options":
|
||||
stream_options = normalize_responses_api_stream_options(value)
|
||||
if stream_options is not None:
|
||||
responses_api_request["stream_options"] = stream_options
|
||||
elif key in ResponsesAPIOptionalRequestParams.__annotations__.keys():
|
||||
responses_api_request[key] = value # type: ignore
|
||||
elif key == "previous_response_id":
|
||||
|
|
@ -360,8 +365,6 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
continue
|
||||
if key == "instructions" and instructions:
|
||||
request_data["instructions"] = instructions
|
||||
elif key == "stream_options" and isinstance(value, dict):
|
||||
request_data["stream_options"] = value.get("include_obfuscation")
|
||||
elif key == "user" and isinstance(value, str):
|
||||
# OpenAI API requires user param to be max 64 chars - truncate if longer
|
||||
if len(value) <= 64:
|
||||
|
|
|
|||
|
|
@ -17,7 +17,11 @@ from typing import (
|
|||
)
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.core_helpers import redact_nested_match_and_regex_keys
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
get_or_create_metadata_bucket,
|
||||
redact_nested_match_and_regex_keys,
|
||||
)
|
||||
from litellm.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.secret_managers.main import str_to_bool
|
||||
|
|
@ -107,6 +111,8 @@ class CustomGuardrail(CustomLogger):
|
|||
# If True, during_call runs async_moderation_hook instead of the unified apply_guardrail path.
|
||||
use_native_during_call_hook: ClassVar[bool] = False
|
||||
|
||||
records_own_guardrail_information: ClassVar[bool] = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
guardrail_name: Optional[str] = None,
|
||||
|
|
@ -954,17 +960,8 @@ class CustomGuardrail(CustomLogger):
|
|||
# should not happen
|
||||
container[key] = [existing, slg]
|
||||
|
||||
if "metadata" in request_data:
|
||||
if request_data["metadata"] is None:
|
||||
request_data["metadata"] = {}
|
||||
_append_guardrail_info(request_data["metadata"])
|
||||
elif "litellm_metadata" in request_data:
|
||||
_append_guardrail_info(request_data["litellm_metadata"])
|
||||
else:
|
||||
# Ensure guardrail info is always logged (e.g. proxy may not have set
|
||||
# metadata yet). Attach to "metadata" so spend log / standard logging see it.
|
||||
request_data["metadata"] = {}
|
||||
_append_guardrail_info(request_data["metadata"])
|
||||
_, metadata_bucket = get_or_create_metadata_bucket(request_data)
|
||||
_append_guardrail_info(metadata_bucket)
|
||||
|
||||
_guardrail_self_recorded.set(True)
|
||||
|
||||
|
|
@ -1223,7 +1220,7 @@ def _sync_guardrail_info_to_logging_obj(request_data: dict, logging_obj: object)
|
|||
"""
|
||||
if logging_obj is None:
|
||||
return
|
||||
meta_src = request_data.get("metadata") or request_data.get("litellm_metadata") or {}
|
||||
meta_src = request_data.get(get_metadata_variable_name_from_kwargs(request_data)) or {}
|
||||
slg_info = meta_src.get("standard_logging_guardrail_information")
|
||||
if not slg_info:
|
||||
return
|
||||
|
|
@ -1256,6 +1253,14 @@ def log_guardrail_information(func):
|
|||
so it stays correct when guardrails run concurrently (asyncio copies the
|
||||
context into each gathered task): counting shared entries would let one
|
||||
guardrail's append hide another guardrail's missing record.
|
||||
|
||||
A guardrail that only records an entry when it actually runs (e.g.
|
||||
``HeadroomGuardrail``, which returns the inputs untouched on an endpoint
|
||||
whose payload it cannot act on) sets ``records_own_guardrail_information =
|
||||
True`` so the auto-record is skipped even on the return paths where it
|
||||
recorded nothing; otherwise a no-op early return would be logged as an
|
||||
"allow"/"success" run even though the guardrail did nothing. The exception
|
||||
branch below still records so a genuine failure is not lost.
|
||||
"""
|
||||
import functools
|
||||
import inspect
|
||||
|
|
@ -1291,7 +1296,7 @@ def log_guardrail_information(func):
|
|||
self_recorded_token = _guardrail_self_recorded.set(False)
|
||||
try:
|
||||
response = await func(*args, **kwargs)
|
||||
if _guardrail_self_recorded.get():
|
||||
if self.records_own_guardrail_information or _guardrail_self_recorded.get():
|
||||
return response
|
||||
return self._process_response(
|
||||
response=response,
|
||||
|
|
@ -1333,7 +1338,7 @@ def log_guardrail_information(func):
|
|||
self_recorded_token = _guardrail_self_recorded.set(False)
|
||||
try:
|
||||
response = func(*args, **kwargs)
|
||||
if _guardrail_self_recorded.get():
|
||||
if self.records_own_guardrail_information or _guardrail_self_recorded.get():
|
||||
return response
|
||||
return self._process_response(
|
||||
response=response,
|
||||
|
|
|
|||
|
|
@ -883,8 +883,8 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
request_data: dict,
|
||||
parent_span: Optional[Any],
|
||||
) -> None:
|
||||
"""Emit ``guardrail`` spans from ``request_data["metadata"]
|
||||
["standard_logging_guardrail_information"]``.
|
||||
"""Emit ``guardrail`` spans from the request's proxy-internal metadata bucket
|
||||
(``standard_logging_guardrail_information``).
|
||||
|
||||
Routed through ``_create_guardrail_span`` so the dedupe state in
|
||||
``_otel_internal`` is honoured — if ``_handle_failure`` already
|
||||
|
|
@ -892,7 +892,12 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
"""
|
||||
from opentelemetry import trace as _trace
|
||||
|
||||
metadata = (request_data or {}).get("metadata") or {}
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
)
|
||||
|
||||
request_data = request_data or {}
|
||||
metadata = request_data.get(get_metadata_variable_name_from_kwargs(request_data)) or {}
|
||||
guardrail_information = metadata.get("standard_logging_guardrail_information")
|
||||
if not guardrail_information:
|
||||
return
|
||||
|
|
|
|||
|
|
@ -195,6 +195,25 @@ def get_metadata_variable_name_from_kwargs(
|
|||
return "litellm_metadata" if "litellm_metadata" in kwargs else "metadata"
|
||||
|
||||
|
||||
def get_or_create_metadata_bucket(
|
||||
request_data: dict,
|
||||
) -> tuple[Literal["metadata", "litellm_metadata"], dict]:
|
||||
"""
|
||||
Return the proxy-internal metadata bucket for this request, creating it if absent.
|
||||
|
||||
Batch/file routes store proxy state in ``litellm_metadata`` so the OpenAI
|
||||
``metadata`` field can remain provider-safe (string values only). Every writer and
|
||||
reader of proxy-internal metadata resolves the bucket through here, so a caller that
|
||||
supplies its own ``metadata`` field cannot split them across two dicts.
|
||||
"""
|
||||
metadata_key = get_metadata_variable_name_from_kwargs(request_data)
|
||||
metadata_bucket = request_data.get(metadata_key)
|
||||
if not isinstance(metadata_bucket, dict):
|
||||
metadata_bucket = {}
|
||||
request_data[metadata_key] = metadata_bucket
|
||||
return metadata_key, metadata_bucket
|
||||
|
||||
|
||||
def get_litellm_metadata_from_kwargs(kwargs: dict):
|
||||
"""
|
||||
Helper to get litellm metadata from all litellm request kwargs
|
||||
|
|
|
|||
|
|
@ -600,9 +600,15 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
guardrail_inputs["tool_calls"] = tool_calls_list
|
||||
|
||||
try:
|
||||
prepared_request_data = self._prepare_request_data(
|
||||
request_data,
|
||||
model_response,
|
||||
user_api_key_dict,
|
||||
key="response",
|
||||
)
|
||||
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs=guardrail_inputs,
|
||||
request_data=request_data if request_data is not None else {},
|
||||
request_data=prepared_request_data,
|
||||
input_type="response",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
|
@ -618,9 +624,15 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
string_so_far = self.get_streaming_string_so_far(responses_so_far)
|
||||
try:
|
||||
prepared_request_data = self._prepare_request_data(
|
||||
request_data,
|
||||
responses_so_far,
|
||||
user_api_key_dict,
|
||||
key="responses",
|
||||
)
|
||||
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs={"texts": [string_so_far]},
|
||||
request_data=request_data if request_data is not None else {},
|
||||
request_data=prepared_request_data,
|
||||
input_type="response",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,12 +1,16 @@
|
|||
import copy
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Literal, Optional
|
||||
from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Optional
|
||||
|
||||
import litellm
|
||||
from litellm import get_secret
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
get_or_create_metadata_bucket,
|
||||
)
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
|
||||
from litellm.proxy._types import CommonProxyErrors, LiteLLMPromptInjectionParams
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
||||
|
|
@ -406,23 +410,6 @@ def get_logging_caching_headers(request_data: Dict) -> Optional[Dict]:
|
|||
return headers
|
||||
|
||||
|
||||
def get_metadata_variable_name_from_kwargs(
|
||||
kwargs: dict,
|
||||
) -> Literal["metadata", "litellm_metadata"]:
|
||||
"""
|
||||
Helper to return what the "metadata" field should be called in the request data
|
||||
|
||||
- New endpoints return `litellm_metadata`
|
||||
- Old endpoints return `metadata`
|
||||
|
||||
Context:
|
||||
- LiteLLM used `metadata` as an internal field for storing metadata
|
||||
- OpenAI then started using this field for their metadata
|
||||
- LiteLLM is now moving to using `litellm_metadata` for our metadata
|
||||
"""
|
||||
return "litellm_metadata" if "litellm_metadata" in kwargs else "metadata"
|
||||
|
||||
|
||||
LITELLM_PROXY_INTERNAL_METADATA_KEYS = frozenset(
|
||||
{
|
||||
"applied_policies",
|
||||
|
|
@ -450,23 +437,6 @@ LITELLM_PROXY_INTERNAL_METADATA_KEYS = frozenset(
|
|||
)
|
||||
|
||||
|
||||
def _get_or_create_proxy_metadata_bucket(
|
||||
request_data: Dict,
|
||||
) -> tuple[Literal["metadata", "litellm_metadata"], dict]:
|
||||
"""
|
||||
Return the proxy-internal metadata bucket for this request.
|
||||
|
||||
Batch/file routes store proxy state in ``litellm_metadata`` so the OpenAI
|
||||
``metadata`` field can remain provider-safe (string values only).
|
||||
"""
|
||||
metadata_key = get_metadata_variable_name_from_kwargs(request_data)
|
||||
metadata_bucket = request_data.get(metadata_key)
|
||||
if not isinstance(metadata_bucket, dict):
|
||||
metadata_bucket = {}
|
||||
request_data[metadata_key] = metadata_bucket
|
||||
return metadata_key, metadata_bucket
|
||||
|
||||
|
||||
def sanitize_openai_provider_metadata(
|
||||
metadata: Optional[Dict[str, Any]],
|
||||
) -> Optional[Dict[str, str]]:
|
||||
|
|
@ -496,7 +466,7 @@ def sanitize_openai_provider_metadata(
|
|||
def add_guardrail_to_applied_guardrails_header(request_data: Dict, guardrail_name: Optional[str]):
|
||||
if guardrail_name is None:
|
||||
return
|
||||
_, _metadata = _get_or_create_proxy_metadata_bucket(request_data)
|
||||
_, _metadata = get_or_create_metadata_bucket(request_data)
|
||||
if "applied_guardrails" in _metadata:
|
||||
if guardrail_name not in _metadata["applied_guardrails"]:
|
||||
_metadata["applied_guardrails"].append(guardrail_name)
|
||||
|
|
@ -513,7 +483,7 @@ def add_policy_to_applied_policies_header(request_data: Dict, policy_name: Optio
|
|||
"""
|
||||
if policy_name is None:
|
||||
return
|
||||
_, _metadata = _get_or_create_proxy_metadata_bucket(request_data)
|
||||
_, _metadata = get_or_create_metadata_bucket(request_data)
|
||||
if "applied_policies" in _metadata:
|
||||
if policy_name not in _metadata["applied_policies"]:
|
||||
_metadata["applied_policies"].append(policy_name)
|
||||
|
|
@ -531,7 +501,7 @@ def add_policy_sources_to_metadata(request_data: Dict, policy_sources: Dict[str,
|
|||
"""
|
||||
if not policy_sources:
|
||||
return
|
||||
_, _metadata = _get_or_create_proxy_metadata_bucket(request_data)
|
||||
_, _metadata = get_or_create_metadata_bucket(request_data)
|
||||
existing = _metadata.get("policy_sources", {})
|
||||
if not isinstance(existing, dict):
|
||||
existing = {}
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import json
|
|||
import re
|
||||
import time
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING, Any, List, Literal, Optional
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, List, Literal, Optional
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -209,6 +209,8 @@ def _build_responses_followup_items(
|
|||
|
||||
|
||||
class HeadroomGuardrail(CustomGuardrail):
|
||||
records_own_guardrail_information: ClassVar[bool] = True
|
||||
|
||||
@classmethod
|
||||
def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]:
|
||||
return [
|
||||
|
|
@ -481,7 +483,21 @@ class HeadroomGuardrail(CustomGuardrail):
|
|||
)
|
||||
end_time = time.time()
|
||||
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_to_applied_guardrails_header,
|
||||
)
|
||||
|
||||
if not compression_succeeded:
|
||||
self.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_json_response={"error": "headroom compression unavailable; request forwarded uncompressed"},
|
||||
request_data=request_data,
|
||||
guardrail_status="guardrail_failed_to_respond",
|
||||
guardrail_provider=HEADROOM_GUARDRAIL_PROVIDER,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
duration=end_time - start_time,
|
||||
)
|
||||
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
|
||||
return {**inputs, "structured_messages": compressed} # pyright: ignore[reportReturnType]
|
||||
|
||||
self.add_standard_logging_guardrail_information_to_request_data(
|
||||
|
|
@ -493,6 +509,7 @@ class HeadroomGuardrail(CustomGuardrail):
|
|||
end_time=end_time,
|
||||
duration=end_time - start_time,
|
||||
)
|
||||
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
|
||||
|
||||
hashes = extract_hashes_from_messages(compressed)
|
||||
if not hashes:
|
||||
|
|
|
|||
|
|
@ -30,6 +30,10 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
get_or_create_metadata_bucket,
|
||||
)
|
||||
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails.guardrail_hooks.model_armor.file_scanning import (
|
||||
|
|
@ -432,7 +436,11 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
Override to store only the Model Armor API response, not the entire data dict.
|
||||
This prevents circular references in logging.
|
||||
"""
|
||||
metadata = (request_data.get("metadata") or {}) if isinstance(request_data, dict) else {}
|
||||
metadata = (
|
||||
request_data.get(get_metadata_variable_name_from_kwargs(request_data)) or {}
|
||||
if isinstance(request_data, dict)
|
||||
else {}
|
||||
)
|
||||
guardrail_response = metadata.get("_model_armor_response", {})
|
||||
|
||||
# Determine status – default to "success" but prefer the explicit value if present.
|
||||
|
|
@ -471,7 +479,6 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
blocking, while fail_on_error still governs real Model Armor API errors.
|
||||
"""
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
_get_or_create_proxy_metadata_bucket,
|
||||
add_guardrail_to_applied_guardrails_header,
|
||||
)
|
||||
|
||||
|
|
@ -491,7 +498,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
|
||||
# Use the same metadata bucket the header helper writes to, so the logged Model Armor
|
||||
# payload and status land where _process_response reads them on every route.
|
||||
_, metadata = _get_or_create_proxy_metadata_bucket(data)
|
||||
_, metadata = get_or_create_metadata_bucket(data)
|
||||
fail_on_error = bool(self.optional_params.get("fail_on_error", True))
|
||||
|
||||
if unscannable_references > 0:
|
||||
|
|
@ -607,7 +614,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
# overwritten by another coroutine.
|
||||
blocked = self._should_block_content(armor_response, allow_sanitization=self.mask_request_content)
|
||||
if isinstance(data, dict):
|
||||
metadata = data.setdefault("metadata", {}) # ensures metadata exists and is unique per request
|
||||
_, metadata = get_or_create_metadata_bucket(data) # ensures metadata exists and is unique per request
|
||||
# Accumulate so a prior file scan on the same request is not overwritten by this text scan.
|
||||
metadata["_model_armor_response"] = self._append_armor_response(
|
||||
metadata.get("_model_armor_response"),
|
||||
|
|
@ -702,7 +709,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
blocked = self._should_block_content(armor_response, allow_sanitization=self.mask_request_content)
|
||||
# Store the armor response for logging
|
||||
if isinstance(data, dict):
|
||||
metadata = data.setdefault("metadata", {})
|
||||
_, metadata = get_or_create_metadata_bucket(data)
|
||||
# Accumulate so a prior file scan on the same request is not overwritten by this text scan.
|
||||
metadata["_model_armor_response"] = self._append_armor_response(
|
||||
metadata.get("_model_armor_response"),
|
||||
|
|
@ -868,7 +875,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
|
||||
# Attach Model Armor response & status to this request's metadata to avoid race conditions
|
||||
if isinstance(request_data, dict):
|
||||
metadata = request_data.setdefault("metadata", {})
|
||||
_, metadata = get_or_create_metadata_bucket(request_data)
|
||||
metadata["_model_armor_response"] = self._build_logging_response(armor_response)
|
||||
metadata["_model_armor_status"] = (
|
||||
"blocked" if self._should_block_content(armor_response) else "success"
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ from typing import TYPE_CHECKING, Any, Literal, NoReturn
|
|||
from urllib.parse import urlsplit
|
||||
|
||||
import httpx
|
||||
from pydantic import ValidationError
|
||||
from pydantic import BaseModel, TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._version import version as litellm_version
|
||||
|
|
@ -24,11 +24,12 @@ from litellm.integrations.custom_guardrail import (
|
|||
log_guardrail_information,
|
||||
)
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import resolve_structured_messages
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.guardrails import GuardrailEventHooks, Mode
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.straiker import (
|
||||
STRAIKER_WEBHOOK_SCHEMA_VERSION,
|
||||
StraikerGuardrailConfigModel,
|
||||
|
|
@ -42,7 +43,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.straiker import (
|
|||
StraikerWebhookStream,
|
||||
StraikerWebhookUsage,
|
||||
)
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs, Usage
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
|
@ -57,6 +58,7 @@ RETRY_STATUS = frozenset({408, 429, 500, 502, 503, 504})
|
|||
UNREACHABLE_STATUS = frozenset({502, 503, 504})
|
||||
_APPLICATION_METADATA_KEYS = frozenset({"agent_id", "app_name"})
|
||||
_OPAQUE_METADATA_SCALAR_TYPES = (str, int, float, bool)
|
||||
_JSON_DICT_ADAPTER = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -137,6 +139,44 @@ def _resolve_destination(request_data: dict) -> str | None:
|
|||
return None
|
||||
|
||||
|
||||
def _route_has_translation(request_data: dict) -> bool:
|
||||
from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route
|
||||
from litellm.llms import load_guardrail_translation_mappings
|
||||
|
||||
route = _as_dict(request_data.get("litellm_metadata")).get("user_api_key_request_route")
|
||||
if not isinstance(route, str) or not route:
|
||||
return False
|
||||
mappings = load_guardrail_translation_mappings()
|
||||
return any(call_type in mappings for call_type in get_call_types_for_route(route) or ())
|
||||
|
||||
|
||||
def _request_structured_messages(request_data: dict) -> list[dict[str, Any]] | None:
|
||||
messages = request_data.get("messages")
|
||||
if messages:
|
||||
return messages if isinstance(messages, list) else None
|
||||
if not _route_has_translation(request_data):
|
||||
return None
|
||||
return resolve_structured_messages(messages=None, request_kwargs=request_data)
|
||||
|
||||
|
||||
def _hook_name(value: object) -> str:
|
||||
return value.value if isinstance(value, GuardrailEventHooks) else str(value)
|
||||
|
||||
|
||||
def _configured_modes(event_hook: object) -> list[str] | None:
|
||||
if isinstance(event_hook, list):
|
||||
names = [_hook_name(v) for v in event_hook]
|
||||
elif isinstance(event_hook, (str, GuardrailEventHooks)):
|
||||
names = [_hook_name(event_hook)]
|
||||
elif isinstance(event_hook, Mode):
|
||||
default = event_hook.default if isinstance(event_hook.default, list) else [event_hook.default]
|
||||
tags = [v for value in event_hook.tags.values() for v in (value if isinstance(value, list) else [value])]
|
||||
names = [_hook_name(v) for v in (*default, *tags) if v is not None]
|
||||
else:
|
||||
return None
|
||||
return list(dict.fromkeys(names)) or None
|
||||
|
||||
|
||||
def _resolve_call_surface(logging_obj: LiteLLMLoggingObj | None, request_data: dict) -> str:
|
||||
call_type = (
|
||||
(getattr(logging_obj, "call_type", None) if logging_obj is not None else None)
|
||||
|
|
@ -146,23 +186,76 @@ def _resolve_call_surface(logging_obj: LiteLLMLoggingObj | None, request_data: d
|
|||
return call_type if isinstance(call_type, str) and call_type else "unknown"
|
||||
|
||||
|
||||
def _jsonable_dict(value: object) -> dict[str, object] | None:
|
||||
if isinstance(value, BaseModel):
|
||||
return _JSON_DICT_ADAPTER.validate_python(value.model_dump(mode="json", exclude_none=True))
|
||||
if isinstance(value, dict):
|
||||
return _JSON_DICT_ADAPTER.validate_python(value)
|
||||
return None
|
||||
|
||||
|
||||
def _opaque_dict_list(value: object) -> list[dict[str, object]] | None:
|
||||
if not isinstance(value, list):
|
||||
return None
|
||||
items = tuple(plain for item in value if (plain := _jsonable_dict(item)) is not None)
|
||||
return list(items) if items else None
|
||||
|
||||
|
||||
def _choice_terminal_reason(choice: object) -> str | None:
|
||||
if isinstance(choice, dict):
|
||||
return _as_optional_str(choice.get("finish_reason")) or _as_optional_str(choice.get("stop_reason"))
|
||||
return _as_optional_str(getattr(choice, "finish_reason", None)) or _as_optional_str(
|
||||
getattr(choice, "stop_reason", None)
|
||||
)
|
||||
|
||||
|
||||
def _response_finish_reason(response: Any) -> str | None:
|
||||
if response is None:
|
||||
return None
|
||||
if isinstance(response, dict):
|
||||
top = _as_optional_str(response.get("finish_reason")) or _as_optional_str(response.get("stop_reason"))
|
||||
if top:
|
||||
return top
|
||||
choices = response.get("choices")
|
||||
if not isinstance(choices, list):
|
||||
return None
|
||||
for choice in choices:
|
||||
reason = _choice_terminal_reason(choice)
|
||||
if reason:
|
||||
return reason
|
||||
return None
|
||||
|
||||
top = _as_optional_str(getattr(response, "finish_reason", None)) or _as_optional_str(
|
||||
getattr(response, "stop_reason", None)
|
||||
)
|
||||
if top:
|
||||
return top
|
||||
choices = getattr(response, "choices", None)
|
||||
if not isinstance(choices, list):
|
||||
return None
|
||||
for choice in choices:
|
||||
reason = getattr(choice, "finish_reason", None)
|
||||
if isinstance(reason, str) and reason:
|
||||
reason = _choice_terminal_reason(choice)
|
||||
if reason:
|
||||
return reason
|
||||
return None
|
||||
|
||||
|
||||
def _as_optional_int(value: object) -> int | None:
|
||||
return value if isinstance(value, int) and not isinstance(value, bool) else None
|
||||
|
||||
|
||||
def _usage_token_count(usage: object, openai_key: str, anthropic_key: str) -> int | None:
|
||||
get = usage.get if isinstance(usage, dict) else lambda key: getattr(usage, key, None)
|
||||
openai_count = _as_optional_int(get(openai_key))
|
||||
return openai_count if openai_count is not None else _as_optional_int(get(anthropic_key))
|
||||
|
||||
|
||||
def _build_usage(response: object) -> StraikerWebhookUsage | None:
|
||||
usage = getattr(response, "usage", None)
|
||||
if not isinstance(usage, Usage):
|
||||
usage = response.get("usage") if isinstance(response, dict) else getattr(response, "usage", None)
|
||||
if usage is None:
|
||||
return None
|
||||
input_tokens = usage.prompt_tokens
|
||||
output_tokens = usage.completion_tokens
|
||||
input_tokens = _usage_token_count(usage, "prompt_tokens", "input_tokens")
|
||||
output_tokens = _usage_token_count(usage, "completion_tokens", "output_tokens")
|
||||
if input_tokens is None and output_tokens is None:
|
||||
return None
|
||||
return StraikerWebhookUsage(input_tokens=input_tokens, output_tokens=output_tokens)
|
||||
|
|
@ -234,6 +327,8 @@ class StraikerGuardrail(CustomGuardrail):
|
|||
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
|
||||
super().__init__(**kwargs)
|
||||
|
||||
self.configured_modes = _configured_modes(self.event_hook)
|
||||
|
||||
def _webhook_url(self) -> str:
|
||||
return f"{self.api_base}{WEBHOOK_PATH}"
|
||||
|
||||
|
|
@ -263,6 +358,7 @@ class StraikerGuardrail(CustomGuardrail):
|
|||
) -> StraikerWebhookContext:
|
||||
return StraikerWebhookContext(
|
||||
call_surface=_resolve_call_surface(logging_obj, request_data),
|
||||
mode=self.configured_modes,
|
||||
model=model,
|
||||
model_provider=_resolve_provider(request_data, model),
|
||||
destination=_resolve_destination(request_data),
|
||||
|
|
@ -287,9 +383,9 @@ class StraikerGuardrail(CustomGuardrail):
|
|||
content = StraikerWebhookContent(
|
||||
texts=list(inputs.get("texts") or []),
|
||||
images=list(inputs.get("images") or []),
|
||||
structured_messages=inputs.get("structured_messages"),
|
||||
tools=inputs.get("tools"),
|
||||
tool_calls=inputs.get("tool_calls"),
|
||||
structured_messages=_opaque_dict_list(inputs.get("structured_messages")),
|
||||
tools=_opaque_dict_list(inputs.get("tools")),
|
||||
tool_calls=_opaque_dict_list(inputs.get("tool_calls")),
|
||||
)
|
||||
|
||||
if input_type == "request":
|
||||
|
|
@ -305,9 +401,8 @@ class StraikerGuardrail(CustomGuardrail):
|
|||
|
||||
response_obj = request_data.get("response")
|
||||
content.finish_reason = _response_finish_reason(response_obj)
|
||||
original_messages = request_data.get("messages")
|
||||
request_content = StraikerWebhookContent(
|
||||
structured_messages=original_messages if isinstance(original_messages, list) else None,
|
||||
structured_messages=_opaque_dict_list(_request_structured_messages(request_data)),
|
||||
)
|
||||
phase: Literal["none", "assembled"] = "assembled" if _is_streamed_request(request_data) else "none"
|
||||
event = StraikerWebhookEvent(type="post_call", id=event_id, stream=StraikerWebhookStream(phase=phase))
|
||||
|
|
|
|||
|
|
@ -147,8 +147,10 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
litellm_logging_obj=data.get("litellm_logging_obj"),
|
||||
)
|
||||
|
||||
# Add guardrail to applied guardrails header
|
||||
add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=guardrail_to_apply.guardrail_name)
|
||||
if not guardrail_to_apply.records_own_guardrail_information:
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=guardrail_to_apply.guardrail_name
|
||||
)
|
||||
return data
|
||||
|
||||
async def async_moderation_hook(
|
||||
|
|
@ -274,8 +276,10 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
if e.original_response is None:
|
||||
e.original_response = response
|
||||
raise
|
||||
# Add guardrail to applied guardrails header
|
||||
add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=guardrail_to_apply.guardrail_name)
|
||||
if not guardrail_to_apply.records_own_guardrail_information:
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=guardrail_to_apply.guardrail_name
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
|
|
|
|||
|
|
@ -38,6 +38,10 @@ from litellm._uuid import uuid
|
|||
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
get_or_create_metadata_bucket,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
|
@ -668,16 +672,13 @@ def _carry_guardrail_logging_info(request_data: dict, guardrail_data: Optional[d
|
|||
"""
|
||||
if guardrail_data is None:
|
||||
return
|
||||
source_metadata = guardrail_data.get("metadata")
|
||||
if not isinstance(source_metadata, dict):
|
||||
return
|
||||
source_key = get_metadata_variable_name_from_kwargs(guardrail_data)
|
||||
source_metadata = guardrail_data.get(source_key) or {}
|
||||
entries = source_metadata.get("standard_logging_guardrail_information")
|
||||
if not entries:
|
||||
return
|
||||
|
||||
metadata = request_data.get("metadata")
|
||||
if not isinstance(metadata, dict):
|
||||
metadata = request_data["metadata"] = {}
|
||||
_, metadata = get_or_create_metadata_bucket(request_data)
|
||||
metadata.setdefault("standard_logging_guardrail_information", list(entries))
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import asyncio
|
||||
from typing import TYPE_CHECKING, Any, Literal, Optional
|
||||
from typing import TYPE_CHECKING, Any, Literal, Mapping, Optional
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException, status
|
||||
|
|
@ -145,6 +145,30 @@ class ProxyModelNotFoundError(HTTPException):
|
|||
super().__init__(status_code=status.HTTP_400_BAD_REQUEST, detail=detail)
|
||||
|
||||
|
||||
REQUIRED_BODY_PARAM_BY_ROUTE: Mapping[str, str] = {
|
||||
"acompletion": "messages",
|
||||
"aembedding": "input",
|
||||
}
|
||||
|
||||
|
||||
class ProxyMissingRequiredParamError(HTTPException):
|
||||
def __init__(self, route: str, param: str):
|
||||
detail = {"error": f"{route}: Missing required parameter: '{param}'."}
|
||||
super().__init__(status_code=status.HTTP_400_BAD_REQUEST, detail=detail)
|
||||
self.type = "invalid_request_error"
|
||||
self.param = param
|
||||
|
||||
|
||||
def raise_if_required_body_param_missing(route_type: str, data: Mapping[str, object]) -> None:
|
||||
required_param = REQUIRED_BODY_PARAM_BY_ROUTE.get(route_type)
|
||||
if required_param is None or data.get(required_param) is not None:
|
||||
return
|
||||
raise ProxyMissingRequiredParamError(
|
||||
route=ROUTE_ENDPOINT_MAPPING.get(route_type, route_type),
|
||||
param=required_param,
|
||||
)
|
||||
|
||||
|
||||
def get_team_id_from_data(data: dict) -> Optional[str]:
|
||||
"""
|
||||
Get the team id from the data's metadata or litellm_metadata params.
|
||||
|
|
@ -353,6 +377,8 @@ async def route_request(
|
|||
"""
|
||||
Common helper to route the request
|
||||
"""
|
||||
raise_if_required_body_param_missing(route_type=route_type, data=data)
|
||||
|
||||
await add_shared_session_to_data(data)
|
||||
|
||||
# Strip router-internal mock_testing_* flags. Combined with an
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ from typing import (
|
|||
Dict,
|
||||
Iterable,
|
||||
List,
|
||||
Mapping,
|
||||
Optional,
|
||||
Type,
|
||||
Union,
|
||||
|
|
@ -24,6 +25,7 @@ from litellm.types.llms.openai import (
|
|||
ResponseInputParam,
|
||||
ResponsesAPIOptionalRequestParams,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamOptions,
|
||||
ResponseText,
|
||||
)
|
||||
from litellm.types.responses.main import DecodedResponseId
|
||||
|
|
@ -35,6 +37,17 @@ from litellm.types.utils import (
|
|||
)
|
||||
|
||||
|
||||
def normalize_responses_api_stream_options(
|
||||
stream_options: object,
|
||||
) -> ResponsesAPIStreamOptions | None:
|
||||
if not isinstance(stream_options, Mapping):
|
||||
return None
|
||||
include_obfuscation = stream_options.get("include_obfuscation")
|
||||
if not isinstance(include_obfuscation, bool):
|
||||
return None
|
||||
return ResponsesAPIStreamOptions(include_obfuscation=include_obfuscation)
|
||||
|
||||
|
||||
class ResponsesAPIRequestUtils:
|
||||
"""Helper utils for constructing ResponseAPI requests"""
|
||||
|
||||
|
|
@ -156,15 +169,19 @@ class ResponsesAPIRequestUtils:
|
|||
drop_params=should_drop_params,
|
||||
)
|
||||
|
||||
stream_options = normalize_responses_api_stream_options(mapped_params.get("stream_options"))
|
||||
params_with_normalized_stream_options = {
|
||||
**{key: value for key, value in mapped_params.items() if key != "stream_options"},
|
||||
**({} if stream_options is None else {"stream_options": stream_options}),
|
||||
}
|
||||
|
||||
# add any allowed_openai_params to the mapped_params
|
||||
mapped_params = _apply_openai_param_overrides(
|
||||
optional_params=mapped_params,
|
||||
return _apply_openai_param_overrides(
|
||||
optional_params=params_with_normalized_stream_options,
|
||||
non_default_params=non_default_params,
|
||||
allowed_openai_params=allowed_openai_params or [],
|
||||
)
|
||||
|
||||
return mapped_params
|
||||
|
||||
@staticmethod
|
||||
def get_requested_response_api_optional_param(
|
||||
params: Dict[str, Any],
|
||||
|
|
|
|||
|
|
@ -1145,6 +1145,10 @@ class ContextManagementEntry(TypedDict, total=False):
|
|||
"""Token threshold at which compaction is triggered for this entry. Minimum 1000."""
|
||||
|
||||
|
||||
class ResponsesAPIStreamOptions(TypedDict, total=False):
|
||||
include_obfuscation: bool
|
||||
|
||||
|
||||
class ResponsesAPIOptionalRequestParams(TypedDict, total=False):
|
||||
"""TypedDict for Optional parameters supported by the responses API."""
|
||||
|
||||
|
|
@ -1171,7 +1175,7 @@ class ResponsesAPIOptionalRequestParams(TypedDict, total=False):
|
|||
max_tool_calls: Optional[int]
|
||||
prompt_cache_key: Optional[str]
|
||||
prompt_cache_retention: Optional[str]
|
||||
stream_options: Optional[dict]
|
||||
stream_options: Optional[ResponsesAPIStreamOptions]
|
||||
top_logprobs: Optional[int]
|
||||
partial_images: Optional[int] # Number of partial images to generate (1-3) for streaming image generation
|
||||
context_management: Optional[List[ContextManagementEntry]]
|
||||
|
|
|
|||
|
|
@ -4,9 +4,6 @@ from typing import Literal
|
|||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolCallChunk
|
||||
from litellm.types.utils import ChatCompletionMessageToolCall
|
||||
|
||||
from .base import GuardrailConfigModel
|
||||
|
||||
StraikerWebhookEventType = Literal["pre_call", "post_call"]
|
||||
|
|
@ -32,9 +29,9 @@ class StraikerWebhookContent(BaseModel):
|
|||
|
||||
texts: list[str] = Field(default_factory=list)
|
||||
images: list[str] = Field(default_factory=list)
|
||||
structured_messages: list[AllMessageValues] | None = None
|
||||
structured_messages: list[dict[str, object]] | None = None
|
||||
tools: list[dict[str, object]] | None = None
|
||||
tool_calls: list[ChatCompletionToolCallChunk] | list[ChatCompletionMessageToolCall] | None = None
|
||||
tool_calls: list[dict[str, object]] | None = None
|
||||
finish_reason: str | None = None
|
||||
|
||||
|
||||
|
|
@ -45,6 +42,7 @@ class StraikerWebhookUsage(BaseModel):
|
|||
|
||||
class StraikerWebhookContext(BaseModel):
|
||||
call_surface: str
|
||||
mode: list[str] | None = None
|
||||
model: str | None = None
|
||||
model_provider: str | None = None
|
||||
destination: str | None = None
|
||||
|
|
|
|||
|
|
@ -22,78 +22,6 @@ def redis_no_ping():
|
|||
yield
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace", [None, "test"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_redis_cache_async_increment(namespace, monkeypatch, redis_no_ping):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
# Create an AsyncMock for the Redis client
|
||||
mock_redis_instance = AsyncMock()
|
||||
|
||||
# Make sure the mock can be used as an async context manager
|
||||
mock_redis_instance.__aenter__.return_value = mock_redis_instance
|
||||
mock_redis_instance.__aexit__.return_value = None
|
||||
|
||||
assert redis_cache is not None
|
||||
|
||||
expected_key = "test:test" if namespace else "test"
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
# Call async_set_cache
|
||||
await redis_cache.async_increment(key=expected_key, value=1)
|
||||
|
||||
# Verify that the set method was called on the mock Redis instance
|
||||
mock_redis_instance.incrbyfloat.assert_called_once_with(
|
||||
name=expected_key, amount=1
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redis_cache_async_increment_refresh_ttl_true_bumps_existing_ttl(
|
||||
monkeypatch, redis_no_ping
|
||||
):
|
||||
"""With refresh_ttl=True, every increment should call expire() to bump
|
||||
the TTL, even when the key already has a TTL (counter-style use)."""
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache()
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_redis_instance.__aenter__.return_value = mock_redis_instance
|
||||
mock_redis_instance.__aexit__.return_value = None
|
||||
mock_redis_instance.ttl.return_value = 42 # key already has ~42s left
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
await redis_cache.async_increment(
|
||||
key="spend:team_member:u:t", value=0.05, refresh_ttl=True
|
||||
)
|
||||
|
||||
mock_redis_instance.expire.assert_awaited_once_with("spend:team_member:u:t", 60)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redis_cache_async_increment_default_does_not_bump_existing_ttl(
|
||||
monkeypatch, redis_no_ping
|
||||
):
|
||||
"""Default (refresh_ttl=False) preserves window-style semantics: TTL is
|
||||
set only on first creation, never refreshed (used by rate-limit windows)."""
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache()
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_redis_instance.__aenter__.return_value = mock_redis_instance
|
||||
mock_redis_instance.__aexit__.return_value = None
|
||||
mock_redis_instance.ttl.return_value = 42 # key already has ~42s left
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
await redis_cache.async_increment(key="rate_limit:window", value=1)
|
||||
|
||||
mock_redis_instance.expire.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace", [None, "litellm"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_delete_cache_applies_namespace(
|
||||
|
|
@ -140,42 +68,6 @@ async def test_redis_client_init_with_socket_timeout(monkeypatch, redis_no_ping)
|
|||
assert client.connection_pool.connection_kwargs["socket_timeout"] == 1.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redis_cache_async_batch_get_cache(monkeypatch, redis_no_ping):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache()
|
||||
|
||||
# Create an AsyncMock for the Redis client
|
||||
mock_redis_instance = AsyncMock()
|
||||
|
||||
# Make sure the mock can be used as an async context manager
|
||||
mock_redis_instance.__aenter__.return_value = mock_redis_instance
|
||||
mock_redis_instance.__aexit__.return_value = None
|
||||
|
||||
# Setup the return value for mget
|
||||
mock_redis_instance.mget.return_value = [
|
||||
b'{"key1": "value1"}',
|
||||
None,
|
||||
b'{"key3": "value3"}',
|
||||
]
|
||||
|
||||
test_keys = ["key1", "key2", "key3"]
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
# Call async_batch_get_cache
|
||||
result = await redis_cache.async_batch_get_cache(key_list=test_keys)
|
||||
|
||||
# Verify mget was called with the correct keys
|
||||
mock_redis_instance.mget.assert_called_once()
|
||||
|
||||
# Check that results were properly decoded
|
||||
assert result["key1"] == {"key1": "value1"}
|
||||
assert result["key2"] is None
|
||||
assert result["key3"] == {"key3": "value3"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_lpop_count_for_older_redis_versions(monkeypatch):
|
||||
"""Test the helper method that handles LPOP with count for Redis versions < 7.0"""
|
||||
|
|
@ -202,41 +94,6 @@ async def test_handle_lpop_count_for_older_redis_versions(monkeypatch):
|
|||
assert mock_pipeline.execute.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_rpush_pipeline_executes_all_operations(monkeypatch, redis_no_ping):
|
||||
"""Verify that multiple rpush ops are batched into a single pipeline execute"""
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache()
|
||||
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_pipeline = MagicMock()
|
||||
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
|
||||
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
|
||||
mock_pipeline.rpush = MagicMock()
|
||||
mock_pipeline.execute = AsyncMock(return_value=[3, 5, 1])
|
||||
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
|
||||
|
||||
from litellm.types.caching import RedisPipelineRpushOperation
|
||||
|
||||
rpush_list = [
|
||||
RedisPipelineRpushOperation(key="key1", values=["a", "b"]),
|
||||
RedisPipelineRpushOperation(key="key2", values=["c"]),
|
||||
RedisPipelineRpushOperation(key="key3", values=["d", "e", "f"]),
|
||||
]
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
result = await redis_cache.async_rpush_pipeline(rpush_list=rpush_list)
|
||||
|
||||
assert result == [3, 5, 1]
|
||||
assert mock_pipeline.rpush.call_count == 3
|
||||
mock_pipeline.rpush.assert_any_call("key1", "a", "b")
|
||||
mock_pipeline.rpush.assert_any_call("key2", "c")
|
||||
mock_pipeline.rpush.assert_any_call("key3", "d", "e", "f")
|
||||
mock_pipeline.execute.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_rpush_pipeline_empty_list_returns_empty(
|
||||
monkeypatch, redis_no_ping
|
||||
|
|
@ -256,183 +113,6 @@ async def test_async_rpush_pipeline_empty_list_returns_empty(
|
|||
mock_redis_instance.pipeline.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_rpush_pipeline_raises_on_redis_error(monkeypatch, redis_no_ping):
|
||||
"""Pipeline errors should propagate"""
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache()
|
||||
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_pipeline = MagicMock()
|
||||
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
|
||||
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
|
||||
mock_pipeline.rpush = MagicMock()
|
||||
mock_pipeline.execute = AsyncMock(side_effect=ConnectionError("Redis down"))
|
||||
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
|
||||
|
||||
from litellm.types.caching import RedisPipelineRpushOperation
|
||||
|
||||
rpush_list = [RedisPipelineRpushOperation(key="key1", values=["a"])]
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with pytest.raises(ConnectionError, match="Redis down"):
|
||||
await redis_cache.async_rpush_pipeline(rpush_list=rpush_list)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_lpop_pipeline_single_round_trip(monkeypatch, redis_no_ping):
|
||||
"""Verify that multiple lpop ops are batched into a single pipeline execute"""
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache()
|
||||
redis_cache.redis_version = "7.0.0"
|
||||
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_pipeline = MagicMock()
|
||||
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
|
||||
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
|
||||
mock_pipeline.lpop = MagicMock()
|
||||
mock_pipeline.execute = AsyncMock(
|
||||
return_value=[
|
||||
[b"val1", b"val2"], # key1 results
|
||||
None, # key2 empty
|
||||
[b"val3"], # key3 results
|
||||
]
|
||||
)
|
||||
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
|
||||
|
||||
from litellm.types.caching import RedisPipelineLpopOperation
|
||||
|
||||
lpop_list = [
|
||||
RedisPipelineLpopOperation(key="key1", count=10),
|
||||
RedisPipelineLpopOperation(key="key2", count=10),
|
||||
RedisPipelineLpopOperation(key="key3", count=5),
|
||||
]
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
results = await redis_cache.async_lpop_pipeline(lpop_list=lpop_list)
|
||||
|
||||
assert len(results) == 3
|
||||
assert results[0] == ["val1", "val2"]
|
||||
assert results[1] is None
|
||||
assert results[2] == ["val3"]
|
||||
mock_pipeline.execute.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_lpop_pipeline_redis_lt7_regroups_flat_results(
|
||||
monkeypatch, redis_no_ping
|
||||
):
|
||||
"""Verify Redis < 7 fallback issues individual LPOPs and regroups correctly"""
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache()
|
||||
redis_cache.redis_version = "6.2.0"
|
||||
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_pipeline = MagicMock()
|
||||
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
|
||||
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
|
||||
mock_pipeline.lpop = MagicMock()
|
||||
|
||||
# With count=3 for key1 and count=2 for key2, we get 5 individual LPOP commands
|
||||
# Simulate: key1 has 2 values then None, key2 has 1 value then None
|
||||
mock_pipeline.execute = AsyncMock(
|
||||
return_value=[
|
||||
b"val1",
|
||||
b"val2",
|
||||
None, # 3 LPOPs for key1
|
||||
b"val3",
|
||||
None, # 2 LPOPs for key2
|
||||
]
|
||||
)
|
||||
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
|
||||
|
||||
from litellm.types.caching import RedisPipelineLpopOperation
|
||||
|
||||
lpop_list = [
|
||||
RedisPipelineLpopOperation(key="key1", count=3),
|
||||
RedisPipelineLpopOperation(key="key2", count=2),
|
||||
]
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
results = await redis_cache.async_lpop_pipeline(lpop_list=lpop_list)
|
||||
|
||||
assert len(results) == 2
|
||||
assert results[0] == ["val1", "val2"] # 2 values, None filtered out
|
||||
assert results[1] == ["val3"] # 1 value, None filtered out
|
||||
# All 5 individual LPOPs should be queued, but only 1 execute() call
|
||||
assert mock_pipeline.lpop.call_count == 5
|
||||
mock_pipeline.execute.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_rpush_pipeline_raises_on_per_command_error(
|
||||
monkeypatch, redis_no_ping
|
||||
):
|
||||
"""Verify that per-command errors in pipeline results are raised, not silently dropped"""
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache()
|
||||
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_pipeline = MagicMock()
|
||||
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
|
||||
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
|
||||
mock_pipeline.rpush = MagicMock()
|
||||
# Simulate: first RPUSH succeeds, second returns a per-command error
|
||||
mock_pipeline.execute = AsyncMock(return_value=[3, Exception("WRONGTYPE")])
|
||||
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
|
||||
|
||||
from litellm.types.caching import RedisPipelineRpushOperation
|
||||
|
||||
rpush_list = [
|
||||
RedisPipelineRpushOperation(key="key1", values=["a"]),
|
||||
RedisPipelineRpushOperation(key="key2", values=["b"]),
|
||||
]
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with pytest.raises(Exception, match="WRONGTYPE"):
|
||||
await redis_cache.async_rpush_pipeline(rpush_list=rpush_list)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_lpop_pipeline_raises_on_per_command_error(
|
||||
monkeypatch, redis_no_ping
|
||||
):
|
||||
"""Verify that per-command errors in LPOP pipeline results are raised, not silently dropped"""
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache()
|
||||
redis_cache.redis_version = "7.0.0"
|
||||
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_pipeline = MagicMock()
|
||||
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
|
||||
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
|
||||
mock_pipeline.lpop = MagicMock()
|
||||
# Simulate: first LPOP succeeds, second returns a per-command error
|
||||
mock_pipeline.execute = AsyncMock(return_value=[[b"val1"], Exception("WRONGTYPE")])
|
||||
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
|
||||
|
||||
from litellm.types.caching import RedisPipelineLpopOperation
|
||||
|
||||
lpop_list = [
|
||||
RedisPipelineLpopOperation(key="key1", count=10),
|
||||
RedisPipelineLpopOperation(key="key2", count=10),
|
||||
]
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with pytest.raises(Exception, match="WRONGTYPE"):
|
||||
await redis_cache.async_lpop_pipeline(lpop_list=lpop_list)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_lpop_pipeline_empty_list(monkeypatch, redis_no_ping):
|
||||
"""Empty lpop_list should return empty list without touching Redis"""
|
||||
|
|
@ -450,111 +130,6 @@ async def test_async_lpop_pipeline_empty_list(monkeypatch, redis_no_ping):
|
|||
mock_redis_instance.pipeline.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_lpop_pipeline_propagates_redis_exception(
|
||||
monkeypatch, redis_no_ping
|
||||
):
|
||||
"""Pipeline errors should propagate"""
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache()
|
||||
redis_cache.redis_version = "7.0.0"
|
||||
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_pipeline = MagicMock()
|
||||
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
|
||||
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
|
||||
mock_pipeline.lpop = MagicMock()
|
||||
mock_pipeline.execute = AsyncMock(side_effect=ConnectionError("Redis down"))
|
||||
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
|
||||
|
||||
from litellm.types.caching import RedisPipelineLpopOperation
|
||||
|
||||
lpop_list = [RedisPipelineLpopOperation(key="key1", count=10)]
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with pytest.raises(ConnectionError, match="Redis down"):
|
||||
await redis_cache.async_lpop_pipeline(lpop_list=lpop_list)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"redis_version",
|
||||
[
|
||||
# Standard cases
|
||||
"7.0.0", # Standard Redis string version
|
||||
7.0, # Valkey/ElastiCache float version (THE BUG this fix addresses)
|
||||
7, # Integer version (e.g., from some Redis forks)
|
||||
# Version < 7
|
||||
"6", # String without dots, version < 7
|
||||
# Malformed versions (fallback to 7)
|
||||
"latest", # Non-numeric version
|
||||
"", # Empty string
|
||||
-7.0, # Negative float
|
||||
# Format variations
|
||||
" 7.0.0 ", # Whitespace (should be stripped)
|
||||
"7.0.0-rc1", # Version with suffix
|
||||
"10.0.0", # Double digit major version
|
||||
],
|
||||
)
|
||||
async def test_async_lpop_with_float_redis_version(
|
||||
monkeypatch, redis_no_ping, redis_version
|
||||
):
|
||||
"""
|
||||
Test async_lpop with various Redis version formats (especially float).
|
||||
|
||||
This test specifically addresses the issue where AWS ElastiCache Valkey
|
||||
returns redis_version as a float (e.g., 7.0) instead of a string (e.g., "7.0.0"),
|
||||
which caused a 'float' object has no attribute 'split' error when trying to
|
||||
use the Redis transaction buffer feature.
|
||||
|
||||
The fix converts the version to a string and handles edge cases like:
|
||||
- Floats (7.0) and integers (7)
|
||||
- Strings with/without dots ("7" vs "7.0.0")
|
||||
- Malformed versions ("v7.0.0", "latest") - fallback to version 7
|
||||
- Whitespace (" 7.0.0 ")
|
||||
- Negative versions (fallback to version 7)
|
||||
|
||||
Related: Database deadlock issues when use_redis_transaction_buffer is enabled.
|
||||
"""
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
|
||||
# Create RedisCache instance
|
||||
redis_cache = RedisCache()
|
||||
redis_cache.redis_version = redis_version # Set the version to test
|
||||
|
||||
# Create an AsyncMock for the Redis client
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_redis_instance.__aenter__.return_value = mock_redis_instance
|
||||
mock_redis_instance.__aexit__.return_value = None
|
||||
|
||||
# Mock lpop to return a test value (Redis >= 7.0 behavior)
|
||||
mock_redis_instance.lpop.return_value = [b"value1", b"value2"]
|
||||
|
||||
# Mock pipeline for Redis < 7.0 (used when major_version < 7)
|
||||
mock_pipeline = MagicMock()
|
||||
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
|
||||
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
|
||||
# Make pipeline() a regular method (not async) that returns the mock
|
||||
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
|
||||
|
||||
# Mock handle_lpop_count_for_older_redis_versions for Redis < 7
|
||||
with patch.object(
|
||||
redis_cache,
|
||||
"handle_lpop_count_for_older_redis_versions",
|
||||
return_value=[b"value1", b"value2"],
|
||||
):
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
# Call async_lpop with count - this should not raise AttributeError
|
||||
result = await redis_cache.async_lpop(key="test_key", count=2)
|
||||
|
||||
# Verify the method completed without error
|
||||
assert result is not None
|
||||
|
||||
|
||||
# LIT-3374: the namespace must be applied uniformly across every key-taking
|
||||
# Redis operation, not just get/set/increment. Before the fix these paths wrote
|
||||
# or read raw keys, so with a namespace configured the prefixed keys other
|
||||
|
|
|
|||
|
|
@ -2853,3 +2853,77 @@ def test_streaming_function_call_tool_id_for_degenerate_call_id():
|
|||
|
||||
assert stream_tool_id("fc_unique_abc123", "call_0") == "fc_unique_abc123"
|
||||
assert stream_tool_id("fc_2", "call_tokyo") == "call_tokyo"
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"stream_options,expected_wire_stream_options",
|
||||
[
|
||||
({"include_usage": True, "include_obfuscation": False}, {"include_obfuscation": False}),
|
||||
({"include_usage": True}, None),
|
||||
],
|
||||
)
|
||||
async def test_acompletion_bridge_normalizes_stream_options_on_the_wire(
|
||||
stream_options, expected_wire_stream_options
|
||||
):
|
||||
"""include_usage must be stripped from the /v1/responses body; include_obfuscation must survive as a dict."""
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
responses_payload = {
|
||||
"id": "resp_bridge_stream_options",
|
||||
"object": "response",
|
||||
"created_at": 1734366691,
|
||||
"status": "completed",
|
||||
"model": "gpt-5.5",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_1",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "hi", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"parallel_tool_calls": True,
|
||||
"usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2},
|
||||
"error": None,
|
||||
"incomplete_details": None,
|
||||
"instructions": None,
|
||||
"metadata": None,
|
||||
"temperature": None,
|
||||
"tool_choice": "auto",
|
||||
"tools": [],
|
||||
"top_p": None,
|
||||
"max_output_tokens": None,
|
||||
"previous_response_id": None,
|
||||
"reasoning": None,
|
||||
"truncation": None,
|
||||
"user": None,
|
||||
}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.text = json.dumps(responses_payload)
|
||||
mock_response.headers = httpx.Headers({})
|
||||
mock_response.json.return_value = responses_payload
|
||||
|
||||
with patch.object(AsyncHTTPHandler, "post", new_callable=AsyncMock) as mock_post:
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
await litellm.acompletion(
|
||||
model="openai/responses/gpt-5.5",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
api_key="fake-api-key",
|
||||
stream_options=stream_options,
|
||||
)
|
||||
|
||||
mock_post.assert_called_once()
|
||||
post_kwargs = mock_post.call_args.kwargs
|
||||
request_body = post_kwargs["json"] if "json" in post_kwargs else json.loads(post_kwargs["data"])
|
||||
if expected_wire_stream_options is None:
|
||||
assert "stream_options" not in request_body
|
||||
else:
|
||||
assert request_body["stream_options"] == expected_wire_stream_options
|
||||
|
|
|
|||
|
|
@ -3,9 +3,12 @@ from unittest.mock import AsyncMock
|
|||
|
||||
import pytest
|
||||
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
)
|
||||
from litellm.proxy._types import CallTypes, UserAPIKeyAuth
|
||||
from litellm.types.utils import GuardrailTracingDetail
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailTracingDetail
|
||||
|
||||
|
||||
class TestCustomGuardrailDeploymentHook:
|
||||
|
|
@ -654,6 +657,53 @@ class TestGuardrailLoggingAggregation:
|
|||
assert len(info) == 2
|
||||
assert info[1]["guardrail_name"] == "test_guardrail"
|
||||
|
||||
def test_caller_metadata_does_not_divert_the_entry_from_the_reader(self):
|
||||
"""A caller-supplied `metadata` field must not send the entry to a bucket the
|
||||
spend log never reads. Routes in LITELLM_METADATA_ROUTES (/v1/messages,
|
||||
/v1/responses, batches, files) seed `litellm_metadata`, and Claude Code sends
|
||||
`metadata.user_id`, so both keys are present on the same request."""
|
||||
request_data = {
|
||||
"metadata": {"user_id": "device-account-session"},
|
||||
"litellm_metadata": {"user_api_key_hash": "abc"},
|
||||
}
|
||||
|
||||
self._invoke_add_log(request_data)
|
||||
|
||||
assert (
|
||||
"standard_logging_guardrail_information" not in request_data["metadata"]
|
||||
), "entry landed in the caller's metadata, where the spend log does not read it"
|
||||
info = request_data["litellm_metadata"][
|
||||
"standard_logging_guardrail_information"
|
||||
]
|
||||
assert len(info) == 1
|
||||
assert info[0]["guardrail_name"] == "test_guardrail"
|
||||
|
||||
def test_entry_and_applied_guardrails_header_share_one_bucket(self):
|
||||
"""The x-litellm-applied-guardrails writer and the guardrail-info writer must
|
||||
resolve the same bucket, otherwise the response header and the spend log
|
||||
disagree about whether the guardrail ran."""
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_to_applied_guardrails_header,
|
||||
)
|
||||
|
||||
request_data = {
|
||||
"metadata": {"user_id": "device-account-session"},
|
||||
"litellm_metadata": {},
|
||||
}
|
||||
|
||||
self._invoke_add_log(request_data)
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=request_data, guardrail_name="test_guardrail"
|
||||
)
|
||||
|
||||
buckets = {
|
||||
key
|
||||
for key in ("metadata", "litellm_metadata")
|
||||
for field in ("standard_logging_guardrail_information", "applied_guardrails")
|
||||
if field in request_data[key]
|
||||
}
|
||||
assert buckets == {"litellm_metadata"}
|
||||
|
||||
|
||||
class TestGuardrailOtelSpanEmission:
|
||||
"""Recording a guardrail emits its otel span inline, so every guardrail
|
||||
|
|
@ -1947,3 +1997,55 @@ class TestOnlyScanNewMessages:
|
|||
cache.async_set_cache = AsyncMock(side_effect=RuntimeError("redis down"))
|
||||
|
||||
await guardrail.mark_texts_scanned(texts=["a"], request_data={"litellm_session_id": "s1"}, cache=cache)
|
||||
|
||||
|
||||
def _guardrail_entries(request_data: dict) -> list:
|
||||
container = request_data.get("metadata") or request_data.get("litellm_metadata") or {}
|
||||
entries = container.get("standard_logging_guardrail_information")
|
||||
return entries if isinstance(entries, list) else []
|
||||
|
||||
|
||||
class _NoopGuardrail(CustomGuardrail):
|
||||
"""apply_guardrail that returns the inputs untouched and records nothing."""
|
||||
|
||||
@log_guardrail_information
|
||||
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
|
||||
return inputs
|
||||
|
||||
|
||||
class _NoopSelfLoggingGuardrail(_NoopGuardrail):
|
||||
records_own_guardrail_information = True
|
||||
|
||||
|
||||
class TestRecordsOwnGuardrailInformation:
|
||||
"""The @log_guardrail_information decorator must not synthesize an "allow"/"success"
|
||||
entry for a no-op apply_guardrail when the guardrail sets
|
||||
records_own_guardrail_information (LIT-4650)."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_default_noop_apply_guardrail_is_auto_logged(self):
|
||||
guardrail = _NoopGuardrail(guardrail_name="g1")
|
||||
request_data: dict = {"model": "gpt-4o"}
|
||||
|
||||
await guardrail.apply_guardrail(
|
||||
inputs=GenericGuardrailAPIInputs(texts=["x"]),
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
entries = _guardrail_entries(request_data)
|
||||
assert len(entries) == 1
|
||||
assert entries[0]["guardrail_status"] == "success"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_self_logging_noop_apply_guardrail_is_not_logged(self):
|
||||
guardrail = _NoopSelfLoggingGuardrail(guardrail_name="g2")
|
||||
request_data: dict = {"model": "gpt-4o"}
|
||||
|
||||
await guardrail.apply_guardrail(
|
||||
inputs=GenericGuardrailAPIInputs(texts=["x"]),
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert _guardrail_entries(request_data) == []
|
||||
|
|
|
|||
|
|
@ -59,8 +59,11 @@ def test_syncs_from_metadata_key():
|
|||
assert result == [entry]
|
||||
|
||||
|
||||
def test_metadata_wins_over_litellm_metadata():
|
||||
"""metadata key takes precedence over litellm_metadata when both are present."""
|
||||
def test_litellm_metadata_wins_over_caller_metadata():
|
||||
"""When both keys are present the helper must read the bucket the writer used,
|
||||
which get_or_create_metadata_bucket resolves to litellm_metadata. Reading the
|
||||
caller's metadata instead is how a guardrail entry went missing from spend logs
|
||||
on the routes that seed litellm_metadata."""
|
||||
entry_meta = _make_slg_entry("from-metadata")
|
||||
entry_lm = _make_slg_entry("from-litellm_metadata")
|
||||
request_data = {
|
||||
|
|
@ -74,7 +77,25 @@ def test_metadata_wins_over_litellm_metadata():
|
|||
result = logging_obj.litellm_params["metadata"].get(
|
||||
"standard_logging_guardrail_information"
|
||||
)
|
||||
assert result == [entry_meta]
|
||||
assert result == [entry_lm]
|
||||
|
||||
|
||||
def test_syncs_when_caller_sends_its_own_metadata():
|
||||
"""The Claude Code shape: caller metadata present, guardrail entry in the seeded
|
||||
litellm_metadata bucket. The entry must still reach the spend-log payload."""
|
||||
entry = _make_slg_entry()
|
||||
request_data = {
|
||||
"metadata": {"user_id": "device-account-session"},
|
||||
"litellm_metadata": {"standard_logging_guardrail_information": [entry]},
|
||||
}
|
||||
logging_obj = _FakeLogging()
|
||||
|
||||
_sync_guardrail_info_to_logging_obj(request_data, logging_obj)
|
||||
|
||||
result = logging_obj.litellm_params["metadata"].get(
|
||||
"standard_logging_guardrail_information"
|
||||
)
|
||||
assert result == [entry]
|
||||
|
||||
|
||||
def test_noop_when_no_guardrail_info():
|
||||
|
|
|
|||
|
|
@ -279,6 +279,48 @@ class TestGuardrailSpanOnViolation(unittest.TestCase):
|
|||
parent_span.context.span_id,
|
||||
)
|
||||
|
||||
def test_post_call_failure_hook_emits_span_when_caller_sends_metadata(self):
|
||||
"""On routes that seed ``litellm_metadata`` the guardrail entry lives there,
|
||||
not in the caller's own ``metadata`` field. Reading a hard-coded ``metadata``
|
||||
key drops the span for exactly the requests that carry both."""
|
||||
otel, provider, exporter = _make_otel()
|
||||
parent_span = provider.get_tracer(__name__).start_span(PROXY_SPAN_NAME)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="sk-test",
|
||||
parent_otel_span=parent_span,
|
||||
request_route="/v1/messages",
|
||||
)
|
||||
|
||||
request_data = {
|
||||
"model": "claude-haiku",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"metadata": {"user_id": "device-account-session"},
|
||||
"litellm_metadata": {
|
||||
"standard_logging_guardrail_information": [
|
||||
_slg_entry("guardrail_intervened", _bedrock_block_response())
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
_run(
|
||||
otel.async_post_call_failure_hook(
|
||||
request_data=request_data,
|
||||
original_exception=Exception("guardrail blocked"),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
)
|
||||
|
||||
guardrail_spans = [
|
||||
s for s in exporter.get_finished_spans() if s.name == GUARDRAIL_SPAN_NAME
|
||||
]
|
||||
self.assertEqual(
|
||||
len(guardrail_spans),
|
||||
1,
|
||||
"the guardrail span must be emitted from the resolved metadata bucket, "
|
||||
"not from a hard-coded 'metadata' key",
|
||||
)
|
||||
|
||||
def test_handle_failure_and_post_call_failure_hook_dedupe(self):
|
||||
"""When _handle_failure and async_post_call_failure_hook BOTH fire
|
||||
for the same request (the production flow on a guardrail block),
|
||||
|
|
|
|||
|
|
@ -4,12 +4,53 @@ import pytest
|
|||
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
_FINISH_REASON_MAP,
|
||||
get_or_create_metadata_bucket,
|
||||
map_finish_reason,
|
||||
reconstruct_model_name,
|
||||
redact_nested_match_and_regex_keys,
|
||||
)
|
||||
|
||||
|
||||
class TestGetOrCreateMetadataBucket:
|
||||
"""The single owner every guardrail writer and reader shares, so the response
|
||||
header and the spend log can never disagree about which dict a record lives in."""
|
||||
|
||||
def test_prefers_litellm_metadata_when_both_present(self):
|
||||
request_data = {"metadata": {"user_id": "caller"}, "litellm_metadata": {}}
|
||||
|
||||
key, bucket = get_or_create_metadata_bucket(request_data)
|
||||
|
||||
assert key == "litellm_metadata"
|
||||
assert bucket is request_data["litellm_metadata"]
|
||||
|
||||
def test_uses_metadata_when_litellm_metadata_absent(self):
|
||||
request_data = {"metadata": {"user_id": "caller"}}
|
||||
|
||||
key, bucket = get_or_create_metadata_bucket(request_data)
|
||||
|
||||
assert key == "metadata"
|
||||
assert bucket is request_data["metadata"]
|
||||
|
||||
def test_creates_the_bucket_in_place_when_missing(self):
|
||||
request_data: dict = {}
|
||||
|
||||
key, bucket = get_or_create_metadata_bucket(request_data)
|
||||
|
||||
assert key == "metadata"
|
||||
assert request_data["metadata"] is bucket
|
||||
bucket["k"] = "v"
|
||||
assert request_data["metadata"]["k"] == "v"
|
||||
|
||||
def test_replaces_a_non_dict_bucket(self):
|
||||
request_data = {"litellm_metadata": None}
|
||||
|
||||
key, bucket = get_or_create_metadata_bucket(request_data)
|
||||
|
||||
assert key == "litellm_metadata"
|
||||
assert isinstance(request_data["litellm_metadata"], dict)
|
||||
assert bucket is request_data["litellm_metadata"]
|
||||
|
||||
|
||||
def test_reconstruct_model_name_prefers_deployment_value():
|
||||
"""Ensure deployment metadata wins when reconstructing the model name."""
|
||||
|
||||
|
|
|
|||
|
|
@ -57,6 +57,98 @@ class MockDynamicGuardrail(CustomGuardrail):
|
|||
return inputs
|
||||
|
||||
|
||||
class MockRecordingGuardrail(CustomGuardrail):
|
||||
"""Mock guardrail that records the request_data it was handed."""
|
||||
|
||||
def __init__(self, guardrail_name: str):
|
||||
super().__init__(guardrail_name=guardrail_name)
|
||||
self.request_data: Optional[dict] = None
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
self.request_data = request_data
|
||||
return inputs
|
||||
|
||||
|
||||
class TestAnthropicMessagesHandlerStreamingRequestData:
|
||||
"""Post-call guardrails on streaming /v1/messages receive the response and identity metadata"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_terminal_chunk_passes_assembled_response_and_metadata(self):
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockRecordingGuardrail(guardrail_name="test")
|
||||
mock_response = ModelResponse(
|
||||
id="msg_123",
|
||||
created=1234567890,
|
||||
model="claude-sonnet-4-5",
|
||||
object="chat.completion",
|
||||
choices=[
|
||||
Choices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
message=Message(content="Hello world", role="assistant"),
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(handler, "_check_streaming_has_ended", return_value=True),
|
||||
patch(
|
||||
"litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler._build_complete_streaming_response",
|
||||
return_value=mock_response,
|
||||
),
|
||||
):
|
||||
await handler.process_output_streaming_response(
|
||||
responses_so_far=[b"data: some chunk"],
|
||||
guardrail_to_apply=guardrail,
|
||||
litellm_logging_obj=MagicMock(),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_id="u-1", team_id="t-1"),
|
||||
request_data={"model": "claude-sonnet-4-5"},
|
||||
)
|
||||
|
||||
assert guardrail.request_data is not None
|
||||
assert guardrail.request_data["response"] is mock_response
|
||||
assert (
|
||||
guardrail.request_data["litellm_metadata"]["user_api_key_user_id"] == "u-1"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mid_stream_chunk_passes_responses_so_far_and_metadata(self):
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockRecordingGuardrail(guardrail_name="test")
|
||||
responses_so_far = [b"data: some chunk"]
|
||||
|
||||
with (
|
||||
patch.object(handler, "_check_streaming_has_ended", return_value=False),
|
||||
patch.object(
|
||||
handler, "get_streaming_string_so_far", return_value="partial text"
|
||||
),
|
||||
):
|
||||
await handler.process_output_streaming_response(
|
||||
responses_so_far=responses_so_far,
|
||||
guardrail_to_apply=guardrail,
|
||||
litellm_logging_obj=MagicMock(),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_id="u-1", team_id="t-1"),
|
||||
request_data={"model": "claude-sonnet-4-5"},
|
||||
)
|
||||
|
||||
assert guardrail.request_data is not None
|
||||
assert guardrail.request_data["responses"] is responses_so_far
|
||||
assert (
|
||||
guardrail.request_data["litellm_metadata"]["user_api_key_user_id"] == "u-1"
|
||||
)
|
||||
|
||||
|
||||
class TestAnthropicMessagesHandlerStreamingOutputProcessing:
|
||||
"""Test streaming output processing functionality"""
|
||||
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import json
|
|||
import os
|
||||
import sys
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
|
|
@ -9,6 +10,7 @@ sys.path.insert(0, os.path.abspath("../../../../.."))
|
|||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import litellm
|
||||
from litellm.anthropic_interface import messages
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.types.utils import Delta, ModelResponse, StreamingChoices
|
||||
|
|
@ -37,6 +39,68 @@ def test_anthropic_experimental_pass_through_messages_handler():
|
|||
assert mock_responses.call_args.kwargs["api_key"] == "test-api-key"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_model_does_not_forward_stream_options_to_responses_api():
|
||||
"""
|
||||
Regression test for LIT-4779. `always_include_stream_usage` injects
|
||||
stream_options={'include_usage': True} into every streaming request, but OpenAI
|
||||
models on /v1/messages go to the Responses API, which 400s on that param.
|
||||
"""
|
||||
responses_payload = {
|
||||
"id": "resp_stream_options",
|
||||
"object": "response",
|
||||
"created_at": 1734366691,
|
||||
"status": "completed",
|
||||
"model": "gpt-5.5",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_1",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "hi", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"parallel_tool_calls": True,
|
||||
"usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2},
|
||||
"error": None,
|
||||
"incomplete_details": None,
|
||||
"instructions": None,
|
||||
"metadata": None,
|
||||
"temperature": None,
|
||||
"tool_choice": "auto",
|
||||
"tools": [],
|
||||
"top_p": None,
|
||||
"max_output_tokens": None,
|
||||
"previous_response_id": None,
|
||||
"reasoning": None,
|
||||
"truncation": None,
|
||||
"user": None,
|
||||
}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.text = json.dumps(responses_payload)
|
||||
mock_response.headers = httpx.Headers({})
|
||||
mock_response.json.return_value = responses_payload
|
||||
|
||||
with patch.object(AsyncHTTPHandler, "post", new_callable=AsyncMock) as mock_post:
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
await litellm.anthropic.messages.acreate(
|
||||
max_tokens=100,
|
||||
messages=[{"role": "user", "content": "Hello, how are you?"}],
|
||||
model="openai/gpt-5.5",
|
||||
api_key="test-api-key",
|
||||
stream_options={"include_usage": True},
|
||||
)
|
||||
|
||||
mock_post.assert_called_once()
|
||||
post_kwargs = mock_post.call_args.kwargs
|
||||
request_body = post_kwargs["json"] if "json" in post_kwargs else json.loads(post_kwargs["data"])
|
||||
assert "stream_options" not in request_body
|
||||
|
||||
|
||||
def test_anthropic_experimental_pass_through_messages_handler_dynamic_api_key_and_api_base_and_custom_values():
|
||||
"""
|
||||
Test that api key, api base, and extra kwargs are forwarded to litellm.completion for Azure models.
|
||||
|
|
|
|||
|
|
@ -181,28 +181,6 @@ async def test_force_ipv4_transport():
|
|||
litellm.disable_aiohttp_transport = original_disable
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ssl_context_transport():
|
||||
"""Test transport creation with SSL context"""
|
||||
# Create a test SSL context
|
||||
ssl_context = ssl.create_default_context()
|
||||
|
||||
transport = AsyncHTTPHandler._create_async_transport(ssl_context=ssl_context)
|
||||
assert transport is not None
|
||||
|
||||
try:
|
||||
if isinstance(transport, LiteLLMAiohttpTransport):
|
||||
# Get the client session and verify SSL context is passed through
|
||||
client_session = transport._get_valid_client_session()
|
||||
assert isinstance(client_session, ClientSession)
|
||||
assert isinstance(client_session.connector, TCPConnector)
|
||||
# Verify the connector has SSL context set by checking if it's using SSL
|
||||
assert client_session.connector._ssl is not None
|
||||
finally:
|
||||
if isinstance(transport, LiteLLMAiohttpTransport):
|
||||
await transport.aclose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aiohttp_disabled_transport():
|
||||
"""Test transport creation with aiohttp disabled"""
|
||||
|
|
@ -339,44 +317,6 @@ async def test_ssl_context_with_shared_session():
|
|||
litellm.disable_aiohttp_transport = original_disable
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aiohttp_transport_trust_env_setting(monkeypatch):
|
||||
"""Test that trust_env setting is properly configured in aiohttp transport"""
|
||||
transports = []
|
||||
try:
|
||||
# Test 1: Default trust_env behavior
|
||||
transport = AsyncHTTPHandler._create_aiohttp_transport()
|
||||
transports.append(transport)
|
||||
client_session = transport._get_valid_client_session()
|
||||
|
||||
# Default should be False (litellm.aiohttp_trust_env default)
|
||||
default_trust_env = getattr(litellm, "aiohttp_trust_env", False)
|
||||
assert client_session._trust_env == default_trust_env
|
||||
|
||||
# Test 2: Environment variable override
|
||||
monkeypatch.setenv("AIOHTTP_TRUST_ENV", "True")
|
||||
transport_with_env = AsyncHTTPHandler._create_aiohttp_transport()
|
||||
transports.append(transport_with_env)
|
||||
client_session_with_env = transport_with_env._get_valid_client_session()
|
||||
|
||||
# Should be True when environment variable is set
|
||||
assert client_session_with_env._trust_env is True
|
||||
|
||||
# Test 3: Verify environment variable with False value
|
||||
monkeypatch.setenv("AIOHTTP_TRUST_ENV", "False")
|
||||
transport_with_false_env = AsyncHTTPHandler._create_aiohttp_transport()
|
||||
transports.append(transport_with_false_env)
|
||||
client_session_with_false_env = (
|
||||
transport_with_false_env._get_valid_client_session()
|
||||
)
|
||||
|
||||
# Should respect the litellm.aiohttp_trust_env setting when env var is False
|
||||
assert client_session_with_false_env._trust_env == default_trust_env
|
||||
finally:
|
||||
for t in transports:
|
||||
await t.aclose()
|
||||
|
||||
|
||||
def test_get_ssl_configuration():
|
||||
"""Test that get_ssl_configuration() returns a proper SSL context with certifi CA bundle
|
||||
when no environment variables are set."""
|
||||
|
|
@ -443,36 +383,6 @@ async def test_create_aiohttp_transport_with_shared_session():
|
|||
assert not callable(transport.client) # Should not be callable
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_aiohttp_transport_without_shared_session():
|
||||
"""Test that _create_aiohttp_transport creates new session when none provided"""
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
# Test without shared session
|
||||
transport = AsyncHTTPHandler._create_aiohttp_transport(shared_session=None)
|
||||
|
||||
# Verify the transport uses a lambda function (for backward compatibility)
|
||||
assert callable(transport.client) # Should be a lambda function
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_aiohttp_transport_with_closed_session():
|
||||
"""Test that _create_aiohttp_transport creates new session when shared session is closed"""
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
# Create a mock closed session
|
||||
mock_session = MockClientSession()
|
||||
mock_session.closed = True
|
||||
|
||||
# Test with closed session
|
||||
transport = AsyncHTTPHandler._create_aiohttp_transport(
|
||||
shared_session=mock_session # type: ignore
|
||||
)
|
||||
|
||||
# Verify the transport creates a new session (lambda function)
|
||||
assert callable(transport.client) # Should be a lambda function
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_handler_with_shared_session():
|
||||
"""Test AsyncHTTPHandler initialization with shared session"""
|
||||
|
|
@ -622,27 +532,6 @@ async def test_session_reuse_integration():
|
|||
await client2.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_validation():
|
||||
"""Test that session validation works correctly"""
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
# Test with None session
|
||||
transport1 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=None)
|
||||
assert callable(transport1.client) # Should create lambda
|
||||
|
||||
# Test with closed session
|
||||
mock_closed_session = MockClientSession()
|
||||
mock_closed_session.closed = True
|
||||
transport2 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=mock_closed_session) # type: ignore
|
||||
assert callable(transport2.client) # Should create lambda
|
||||
|
||||
# Test with valid session
|
||||
mock_valid_session = MockClientSession()
|
||||
transport3 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=mock_valid_session) # type: ignore
|
||||
assert transport3.client is mock_valid_session # Should reuse session
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"env_curve,litellm_curve,expected_curve,should_call",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -1,15 +1,10 @@
|
|||
"""
|
||||
Tests for LlmPassthroughRouteHandler and the guardrail_translation_mappings registry.
|
||||
Tests for the guardrail_translation_mappings registry.
|
||||
|
||||
Validates:
|
||||
- allm_passthrough_route is registered in the mappings (regression: this was the bug)
|
||||
- Bedrock provider is dispatched to BedrockPassthroughGuardrailHandler
|
||||
- Unknown provider skips apply_guardrail
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from litellm.llms.pass_through.guardrail_translation import (
|
||||
guardrail_translation_mappings,
|
||||
)
|
||||
|
|
@ -40,185 +35,3 @@ class TestRegistry:
|
|||
is PassThroughEndpointHandler
|
||||
)
|
||||
|
||||
|
||||
def _make_guardrail() -> MagicMock:
|
||||
g = MagicMock()
|
||||
g.guardrail_name = "test-guard"
|
||||
g.apply_guardrail = AsyncMock(return_value={"texts": []})
|
||||
g.skip_system_message_in_guardrail = False
|
||||
g.skip_tool_message_in_guardrail = False
|
||||
return g
|
||||
|
||||
|
||||
class TestLlmPassthroughRouteHandlerInput:
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_provider_delegates_to_bedrock_handler(self):
|
||||
handler = LlmPassthroughRouteHandler()
|
||||
data = {
|
||||
"custom_llm_provider": "bedrock",
|
||||
"endpoint": "model/anthropic.claude-3-sonnet/converse",
|
||||
"data": {"messages": [{"role": "user", "content": [{"text": "hi"}]}]},
|
||||
}
|
||||
guardrail = _make_guardrail()
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
guardrail.apply_guardrail.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unknown_provider_skips_apply_guardrail(self):
|
||||
handler = LlmPassthroughRouteHandler()
|
||||
data = {
|
||||
"custom_llm_provider": "some_unknown_provider",
|
||||
"endpoint": "v1/chat/completions",
|
||||
"data": {"messages": [{"role": "user", "content": "hi"}]},
|
||||
}
|
||||
guardrail = _make_guardrail()
|
||||
|
||||
result = await handler.process_input_messages(
|
||||
data=data, guardrail_to_apply=guardrail
|
||||
)
|
||||
|
||||
guardrail.apply_guardrail.assert_not_called()
|
||||
assert result is data
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_provider_skips(self):
|
||||
handler = LlmPassthroughRouteHandler()
|
||||
data = {"endpoint": "foo/bar", "data": {}}
|
||||
guardrail = _make_guardrail()
|
||||
|
||||
result = await handler.process_input_messages(
|
||||
data=data, guardrail_to_apply=guardrail
|
||||
)
|
||||
|
||||
guardrail.apply_guardrail.assert_not_called()
|
||||
assert result is data
|
||||
|
||||
|
||||
class TestLlmPassthroughRouteHandlerOutput:
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_provider_delegates_output_to_bedrock_handler(self):
|
||||
handler = LlmPassthroughRouteHandler()
|
||||
response = {
|
||||
"output": {
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": [{"text": "hello"}],
|
||||
}
|
||||
}
|
||||
}
|
||||
request_data = {
|
||||
"custom_llm_provider": "bedrock",
|
||||
"endpoint": "model/anthropic.claude-3-sonnet/converse",
|
||||
}
|
||||
guardrail = _make_guardrail()
|
||||
|
||||
await handler.process_output_response(
|
||||
response=response,
|
||||
guardrail_to_apply=guardrail,
|
||||
request_data=request_data,
|
||||
)
|
||||
|
||||
guardrail.apply_guardrail.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unknown_provider_skips_output(self):
|
||||
handler = LlmPassthroughRouteHandler()
|
||||
response = {"some": "response"}
|
||||
request_data = {"custom_llm_provider": "unknown"}
|
||||
guardrail = _make_guardrail()
|
||||
|
||||
result = await handler.process_output_response(
|
||||
response=response,
|
||||
guardrail_to_apply=guardrail,
|
||||
request_data=request_data,
|
||||
)
|
||||
|
||||
guardrail.apply_guardrail.assert_not_called()
|
||||
assert result is response
|
||||
|
||||
|
||||
class TestDeAnonymizeEventStream:
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_provider_dispatches_to_handler(self):
|
||||
body = b"original-stream-bytes"
|
||||
expected = b"de-anonymized-bytes"
|
||||
proxy_logging_obj = MagicMock()
|
||||
user_api_key_dict = MagicMock()
|
||||
|
||||
with patch(
|
||||
"litellm.llms.bedrock.passthrough.guardrail_translation.handler."
|
||||
"BedrockPassthroughGuardrailHandler.de_anonymize_event_stream",
|
||||
new=AsyncMock(return_value=expected),
|
||||
) as mock_handler:
|
||||
result = await LlmPassthroughRouteHandler.de_anonymize_event_stream(
|
||||
body_bytes=body,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
data={"custom_llm_provider": "bedrock"},
|
||||
)
|
||||
|
||||
mock_handler.assert_awaited_once()
|
||||
assert result == expected
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unknown_provider_returns_original_bytes(self):
|
||||
body = b"original-stream-bytes"
|
||||
|
||||
result = await LlmPassthroughRouteHandler.de_anonymize_event_stream(
|
||||
body_bytes=body,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
user_api_key_dict=MagicMock(),
|
||||
data={"custom_llm_provider": "anthropic"},
|
||||
)
|
||||
|
||||
assert result is body
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_provider_returns_original_bytes(self):
|
||||
body = b"original-stream-bytes"
|
||||
|
||||
result = await LlmPassthroughRouteHandler.de_anonymize_event_stream(
|
||||
body_bytes=body,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
user_api_key_dict=MagicMock(),
|
||||
data={},
|
||||
)
|
||||
|
||||
assert result is body
|
||||
|
||||
|
||||
class TestSupportsEventStreamDeAnonymization:
|
||||
def test_bedrock_converse_stream_is_supported(self):
|
||||
assert (
|
||||
LlmPassthroughRouteHandler.supports_event_stream_de_anonymization(
|
||||
"bedrock", "model/us.amazon.nova-lite-v1:0/converse-stream"
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
def test_bedrock_invoke_stream_is_not_supported(self):
|
||||
assert (
|
||||
LlmPassthroughRouteHandler.supports_event_stream_de_anonymization(
|
||||
"bedrock",
|
||||
"model/us.amazon.nova-lite-v1:0/invoke-with-response-stream",
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
def test_unknown_provider_is_not_supported(self):
|
||||
assert (
|
||||
LlmPassthroughRouteHandler.supports_event_stream_de_anonymization(
|
||||
"anthropic", "model/foo/converse-stream"
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
def test_missing_provider_is_not_supported(self):
|
||||
assert (
|
||||
LlmPassthroughRouteHandler.supports_event_stream_de_anonymization(
|
||||
None, "model/foo/converse-stream"
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
|
|
|||
|
|
@ -729,10 +729,15 @@ async def test_openai_moderation_post_call_request_data_passthrough():
|
|||
|
||||
mock_make_request.assert_called_once()
|
||||
|
||||
# Guardrail info in the REAL request_data (not a throwaway)
|
||||
guardrail_info_list = request_data["metadata"].get(
|
||||
"standard_logging_guardrail_information"
|
||||
# Guardrail info in the REAL request_data (not a throwaway). The unified hook
|
||||
# seeds litellm_metadata, so read the bucket the resolver names rather than
|
||||
# assuming "metadata"; the spend log reads it the same way.
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
)
|
||||
|
||||
bucket = request_data[get_metadata_variable_name_from_kwargs(request_data)]
|
||||
guardrail_info_list = bucket.get("standard_logging_guardrail_information")
|
||||
assert guardrail_info_list is not None
|
||||
assert isinstance(guardrail_info_list[0]["guardrail_response"], dict)
|
||||
assert "results" in guardrail_info_list[0]["guardrail_response"]
|
||||
|
|
|
|||
|
|
@ -259,10 +259,15 @@ async def test_openai_moderation_streaming_end_of_stream_request_data_passthroug
|
|||
):
|
||||
pass
|
||||
|
||||
# Verify guardrail info reached the REAL request_data (not a throwaway)
|
||||
guardrail_info_list = request_data["metadata"].get(
|
||||
"standard_logging_guardrail_information"
|
||||
# Verify guardrail info reached the REAL request_data (not a throwaway). The
|
||||
# unified hook seeds litellm_metadata, so read the bucket the resolver names
|
||||
# rather than assuming "metadata"; the spend log reads it the same way.
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
)
|
||||
|
||||
bucket = request_data[get_metadata_variable_name_from_kwargs(request_data)]
|
||||
guardrail_info_list = bucket.get("standard_logging_guardrail_information")
|
||||
assert (
|
||||
guardrail_info_list is not None
|
||||
), "Guardrail info should be in request_data after streaming"
|
||||
|
|
|
|||
|
|
@ -114,6 +114,24 @@ def guardrail() -> HeadroomGuardrail:
|
|||
return _make_guardrail()
|
||||
|
||||
|
||||
def _recorded_guardrail_entries(request_data: dict) -> list:
|
||||
for container_key in ("metadata", "litellm_metadata"):
|
||||
container = request_data.get(container_key)
|
||||
if isinstance(container, dict):
|
||||
entries = container.get("standard_logging_guardrail_information")
|
||||
if isinstance(entries, list):
|
||||
return entries
|
||||
return []
|
||||
|
||||
|
||||
def _applied_guardrails(request_data: dict) -> list:
|
||||
for container_key in ("metadata", "litellm_metadata"):
|
||||
container = request_data.get(container_key)
|
||||
if isinstance(container, dict) and isinstance(container.get("applied_guardrails"), list):
|
||||
return container["applied_guardrails"]
|
||||
return []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_compresses_and_returns_structured_messages(
|
||||
guardrail: HeadroomGuardrail,
|
||||
|
|
@ -123,6 +141,7 @@ async def test_apply_guardrail_compresses_and_returns_structured_messages(
|
|||
structured_messages=ORIGINAL_MESSAGES,
|
||||
)
|
||||
mock_response = _make_compress_response(COMPRESSED_MESSAGES)
|
||||
request_data = {"model": "gpt-4o"}
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler,
|
||||
|
|
@ -132,12 +151,19 @@ async def test_apply_guardrail_compresses_and_returns_structured_messages(
|
|||
):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data={"model": "gpt-4o"},
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert result.get("structured_messages") == COMPRESSED_MESSAGES
|
||||
|
||||
entries = _recorded_guardrail_entries(request_data)
|
||||
assert len(entries) == 1
|
||||
assert entries[0]["guardrail_name"] == "headroom"
|
||||
assert entries[0]["guardrail_status"] == "success"
|
||||
assert entries[0]["guardrail_provider"] == "headroom"
|
||||
assert "headroom" in _applied_guardrails(request_data)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_injects_retrieve_tool_when_hashes_present(
|
||||
|
|
@ -719,6 +745,7 @@ async def test_apply_guardrail_bypass_header_skips_compression(
|
|||
mock_post.assert_not_called()
|
||||
|
||||
assert result.get("structured_messages") == ORIGINAL_MESSAGES
|
||||
assert _recorded_guardrail_entries(request_data) == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -729,16 +756,18 @@ async def test_apply_guardrail_response_type_passthrough(
|
|||
texts=["some response text"],
|
||||
structured_messages=ORIGINAL_MESSAGES,
|
||||
)
|
||||
request_data: dict = {"model": "gpt-4o"}
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data={},
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
)
|
||||
mock_post.assert_not_called()
|
||||
|
||||
assert result is inputs
|
||||
assert _recorded_guardrail_entries(request_data) == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -746,16 +775,48 @@ async def test_apply_guardrail_empty_structured_messages_passthrough(
|
|||
guardrail: HeadroomGuardrail,
|
||||
):
|
||||
inputs = GenericGuardrailAPIInputs(texts=["hello"])
|
||||
request_data: dict = {"model": "gpt-4o"}
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data={},
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
mock_post.assert_not_called()
|
||||
|
||||
assert result is inputs
|
||||
assert _recorded_guardrail_entries(request_data) == []
|
||||
assert "headroom" not in _applied_guardrails(request_data)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_passthrough_handler_does_not_log_headroom_as_run(
|
||||
guardrail: HeadroomGuardrail,
|
||||
):
|
||||
"""Regression for LIT-4650.
|
||||
|
||||
A passthrough request drives headroom through PassThroughEndpointHandler, which
|
||||
only supplies `texts` (no `structured_messages`). Headroom cannot compress that
|
||||
shape and no-ops, so it must not appear in the spend log's
|
||||
standard_logging_guardrail_information as a successful run.
|
||||
"""
|
||||
from litellm.llms.pass_through.guardrail_translation.handler import (
|
||||
PassThroughEndpointHandler,
|
||||
)
|
||||
|
||||
data = {"model": "gpt-4o", "messages": [{"role": "user", "content": "hello"}]}
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
|
||||
await PassThroughEndpointHandler().process_input_messages(
|
||||
data=data,
|
||||
guardrail_to_apply=guardrail,
|
||||
litellm_logging_obj=None,
|
||||
)
|
||||
mock_post.assert_not_called()
|
||||
|
||||
assert _recorded_guardrail_entries(data) == []
|
||||
assert "headroom" not in _applied_guardrails(data)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -884,6 +945,7 @@ async def test_apply_guardrail_transport_error_fail_open_forwards_uncompressed()
|
|||
texts=["hello"],
|
||||
structured_messages=ORIGINAL_MESSAGES,
|
||||
)
|
||||
request_data = {"model": "gpt-4o"}
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler,
|
||||
|
|
@ -893,12 +955,18 @@ async def test_apply_guardrail_transport_error_fail_open_forwards_uncompressed()
|
|||
):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data={},
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert result["structured_messages"] == ORIGINAL_MESSAGES
|
||||
|
||||
entries = _recorded_guardrail_entries(request_data)
|
||||
assert len(entries) == 1
|
||||
assert entries[0]["guardrail_name"] == "headroom"
|
||||
assert entries[0]["guardrail_status"] == "guardrail_failed_to_respond"
|
||||
assert "headroom" in _applied_guardrails(request_data)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_http_error_fail_open_forwards_uncompressed():
|
||||
|
|
|
|||
|
|
@ -3502,6 +3502,44 @@ async def test_single_scan_response_stays_a_dict():
|
|||
assert isinstance(request_data["metadata"]["_model_armor_response"], dict)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scan_result_reaches_the_logger_on_a_seeded_route():
|
||||
"""On routes that seed `litellm_metadata` the scan result must land in that bucket
|
||||
and be found by `_process_response`. Writing the file-scan result through the shared
|
||||
resolver while the text-scan writers and the reader used a hard-coded `metadata` key
|
||||
split the record in two, so the logged guardrail payload came back empty."""
|
||||
guardrail = _make_guardrail()
|
||||
pdf_b64 = base64.b64encode(PDF_BYTES).decode("utf-8")
|
||||
request_data = {
|
||||
"model": "claude-haiku",
|
||||
"messages": [_file_message(pdf_b64)],
|
||||
"metadata": {"user_id": "device-account-session"},
|
||||
"litellm_metadata": {"guardrails": ["model-armor-test"]},
|
||||
}
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler,
|
||||
"post",
|
||||
AsyncMock(return_value=_armor_response(blocked=False)),
|
||||
):
|
||||
await guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=MagicMock(spec=DualCache),
|
||||
data=request_data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
assert "_model_armor_response" not in request_data["metadata"]
|
||||
assert "_model_armor_response" in request_data["litellm_metadata"]
|
||||
|
||||
before = len(request_data["litellm_metadata"].get("standard_logging_guardrail_information", []))
|
||||
guardrail._process_response(response=None, request_data=request_data)
|
||||
|
||||
logged = request_data["litellm_metadata"]["standard_logging_guardrail_information"]
|
||||
assert len(logged) == before + 1
|
||||
assert logged[-1]["guardrail_response"], "the logger recorded an empty Model Armor payload"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_blocks_supported_document_with_undecodable_base64():
|
||||
"""A supported document whose inline base64 will not decode cannot be scanned, so it fails closed."""
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import json
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import httpx
|
||||
|
|
@ -8,6 +9,9 @@ from litellm.exceptions import GuardrailRaisedException, ModifyResponseException
|
|||
from litellm.proxy.guardrails.guardrail_hooks.straiker import initialize_guardrail
|
||||
from litellm.proxy.guardrails.guardrail_hooks.straiker.straiker import (
|
||||
StraikerGuardrail,
|
||||
_build_usage,
|
||||
_request_structured_messages,
|
||||
_response_finish_reason,
|
||||
)
|
||||
from litellm.proxy.guardrails.guardrail_registry import (
|
||||
guardrail_class_registry,
|
||||
|
|
@ -17,7 +21,14 @@ from litellm.types.proxy.guardrails.guardrail_hooks.straiker import (
|
|||
StraikerGuardrailConfigModel,
|
||||
StraikerGuardrailConfigModelOptionalParams,
|
||||
)
|
||||
from litellm.types.utils import Choices, Message, ModelResponse, Usage
|
||||
from litellm.types.utils import (
|
||||
ChatCompletionMessageToolCall,
|
||||
Choices,
|
||||
Function,
|
||||
Message,
|
||||
ModelResponse,
|
||||
Usage,
|
||||
)
|
||||
|
||||
|
||||
def _mock_response(action: str, turn_id: str = "turn-1", schema_version: str = "1", **extra) -> MagicMock:
|
||||
|
|
@ -208,7 +219,9 @@ async def test_request_envelope_transport_and_shape():
|
|||
"metadata": {"user_api_key_alias": "team-key", "agent_id": "chatbot-app", "app_name": "Chatbot"},
|
||||
}
|
||||
|
||||
out = await g.apply_guardrail(inputs=inputs, request_data=request_data, input_type="request", logging_obj=_logging_obj())
|
||||
out = await g.apply_guardrail(
|
||||
inputs=inputs, request_data=request_data, input_type="request", logging_obj=_logging_obj()
|
||||
)
|
||||
|
||||
assert out is inputs
|
||||
url = g.async_handler.post.call_args.args[0]
|
||||
|
|
@ -232,6 +245,35 @@ async def test_request_envelope_transport_and_shape():
|
|||
assert "metadata" not in payload
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_request_envelope_ignores_unsupported_opaque_items():
|
||||
g = _make_guardrail()
|
||||
g.async_handler.post.return_value = _mock_response("NONE")
|
||||
|
||||
await g.apply_guardrail(
|
||||
inputs={
|
||||
"texts": ["hello"],
|
||||
"tools": [
|
||||
object(),
|
||||
{
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "parameters": {"type": "object"}},
|
||||
},
|
||||
],
|
||||
},
|
||||
request_data={"model": "m", "messages": [{"role": "user", "content": "hello"}]},
|
||||
input_type="request",
|
||||
logging_obj=_logging_obj(),
|
||||
)
|
||||
|
||||
assert _posted_payload(g)["request"]["tools"] == [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "parameters": {"type": "object"}},
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_webhook_metadata_session_id_and_opaque_passthrough():
|
||||
g = _make_guardrail()
|
||||
|
|
@ -331,6 +373,64 @@ async def test_context_session_id_from_request_metadata():
|
|||
assert "metadata" not in payload
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_context_mode_from_string_event_hook():
|
||||
g = _make_guardrail(event_hook="pre_call")
|
||||
g.async_handler.post.return_value = _mock_response("NONE")
|
||||
await g.apply_guardrail(
|
||||
inputs={"texts": ["x"]},
|
||||
request_data={"model": "m"},
|
||||
input_type="request",
|
||||
logging_obj=_logging_obj(),
|
||||
)
|
||||
assert _posted_payload(g)["context"]["mode"] == ["pre_call"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_context_mode_from_list_event_hook():
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
g = _make_guardrail(event_hook=[GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call])
|
||||
g.async_handler.post.return_value = _mock_response("NONE")
|
||||
await g.apply_guardrail(
|
||||
inputs={"texts": ["x"]},
|
||||
request_data={"model": "m"},
|
||||
input_type="request",
|
||||
logging_obj=_logging_obj(),
|
||||
)
|
||||
assert _posted_payload(g)["context"]["mode"] == ["pre_call", "post_call"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_context_mode_from_tagged_mode_is_flattened_and_deduped():
|
||||
from litellm.types.guardrails import Mode
|
||||
|
||||
g = _make_guardrail(
|
||||
event_hook=Mode(tags={"team-a": "pre_call", "team-b": ["post_call", "pre_call"]}, default="post_call")
|
||||
)
|
||||
g.async_handler.post.return_value = _mock_response("NONE")
|
||||
await g.apply_guardrail(
|
||||
inputs={"texts": ["x"]},
|
||||
request_data={"model": "m"},
|
||||
input_type="request",
|
||||
logging_obj=_logging_obj(),
|
||||
)
|
||||
assert _posted_payload(g)["context"]["mode"] == ["post_call", "pre_call"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_context_mode_omitted_when_event_hook_absent():
|
||||
g = _make_guardrail(event_hook=None)
|
||||
g.async_handler.post.return_value = _mock_response("NONE")
|
||||
await g.apply_guardrail(
|
||||
inputs={"texts": ["x"]},
|
||||
request_data={"model": "m"},
|
||||
input_type="request",
|
||||
logging_obj=_logging_obj(),
|
||||
)
|
||||
assert "mode" not in _posted_payload(g)["context"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_identity_key_and_team_coalesce_alias_over_id():
|
||||
g = _make_guardrail()
|
||||
|
|
@ -426,6 +526,7 @@ async def test_application_source_from_agent_id():
|
|||
)
|
||||
assert _posted_payload(g)["application"] == {"source": "analytics-app", "name": "Analytics"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_request_block_raises_guardrail_exception_with_reason():
|
||||
g = _make_guardrail()
|
||||
|
|
@ -560,6 +661,34 @@ async def test_response_envelope_and_block_replaces_response():
|
|||
assert payload["request"]["structured_messages"] == [{"role": "user", "content": "original prompt"}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_resolves_request_from_responses_input_when_messages_absent():
|
||||
g = _make_guardrail()
|
||||
g.async_handler.post.return_value = _mock_response("NONE")
|
||||
response = ModelResponse(
|
||||
choices=[Choices(finish_reason="stop", index=0, message=Message(content="answer", role="assistant"))],
|
||||
model="gpt-4o-mini",
|
||||
)
|
||||
request_data = {
|
||||
"model": "gpt-4o-mini",
|
||||
"input": "responses-surface prompt",
|
||||
"response": response,
|
||||
"litellm_metadata": {"user_api_key_request_route": "/v1/responses"},
|
||||
}
|
||||
|
||||
await g.apply_guardrail(
|
||||
inputs={"texts": ["answer"], "model": "gpt-4o-mini"},
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
logging_obj=_logging_obj(),
|
||||
)
|
||||
|
||||
payload = _posted_payload(g)
|
||||
assert payload["event"]["type"] == "post_call"
|
||||
messages = payload["request"]["structured_messages"]
|
||||
assert any(m.get("content") == "responses-surface prompt" for m in messages)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_fail_closed_raises_modify_response_exception():
|
||||
g = _make_guardrail(unreachable_fallback="fail_closed")
|
||||
|
|
@ -731,3 +860,210 @@ async def test_unreachable_http_status_fail_closed_blocks():
|
|||
await g.apply_guardrail(
|
||||
inputs={"texts": ["x"]}, request_data={"model": "m"}, input_type="request", logging_obj=_logging_obj()
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_preserves_anthropic_tool_blocks_in_request_messages():
|
||||
g = _make_guardrail()
|
||||
g.async_handler.post.return_value = _mock_response("NONE")
|
||||
anthropic_messages = [
|
||||
{"role": "user", "content": "What's the weather in Paris?"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": "toolu_1",
|
||||
"name": "get_weather",
|
||||
"input": {"city": "Paris"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_result",
|
||||
"tool_use_id": "toolu_1",
|
||||
"content": "18C, cloudy",
|
||||
}
|
||||
],
|
||||
},
|
||||
]
|
||||
response = {
|
||||
"id": "msg_1",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "Mild and cloudy."}],
|
||||
"stop_reason": "end_turn",
|
||||
"model": "claude-sonnet-5",
|
||||
}
|
||||
await g.apply_guardrail(
|
||||
inputs={"texts": ["Mild and cloudy."], "model": "claude-sonnet-5"},
|
||||
request_data={
|
||||
"model": "claude-sonnet-5",
|
||||
"messages": anthropic_messages,
|
||||
"response": response,
|
||||
},
|
||||
input_type="response",
|
||||
logging_obj=_logging_obj(),
|
||||
)
|
||||
payload = _posted_payload(g)
|
||||
assert payload["request"]["structured_messages"] == anthropic_messages
|
||||
assert payload["response"]["finish_reason"] == "end_turn"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_preserves_anthropic_tool_blocks_in_structured_messages():
|
||||
g = _make_guardrail()
|
||||
g.async_handler.post.return_value = _mock_response("NONE")
|
||||
anthropic_messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": "toolu_1",
|
||||
"name": "get_weather",
|
||||
"input": {"city": "Paris"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_result",
|
||||
"tool_use_id": "toolu_1",
|
||||
"content": "18C",
|
||||
}
|
||||
],
|
||||
},
|
||||
]
|
||||
await g.apply_guardrail(
|
||||
inputs={"structured_messages": anthropic_messages, "model": "claude-sonnet-5"},
|
||||
request_data={"model": "claude-sonnet-5", "messages": anthropic_messages},
|
||||
input_type="request",
|
||||
logging_obj=_logging_obj(),
|
||||
)
|
||||
assert _posted_payload(g)["request"]["structured_messages"] == anthropic_messages
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_response_finish_reason_from_openai_choices_still_works():
|
||||
g = _make_guardrail()
|
||||
g.async_handler.post.return_value = _mock_response("NONE")
|
||||
response = ModelResponse(
|
||||
choices=[Choices(finish_reason="tool_calls", index=0, message=Message(content=None, role="assistant"))],
|
||||
model="gpt-4o-mini",
|
||||
)
|
||||
await g.apply_guardrail(
|
||||
inputs={
|
||||
"texts": [],
|
||||
"tool_calls": [
|
||||
ChatCompletionMessageToolCall(
|
||||
id="c1",
|
||||
type="function",
|
||||
function=Function(name="f", arguments="{}"),
|
||||
)
|
||||
],
|
||||
},
|
||||
request_data={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}], "response": response},
|
||||
input_type="response",
|
||||
logging_obj=_logging_obj(),
|
||||
)
|
||||
payload = _posted_payload(g)
|
||||
assert payload["response"]["finish_reason"] == "tool_calls"
|
||||
assert payload["response"]["tool_calls"] == [
|
||||
{"id": "c1", "type": "function", "function": {"name": "f", "arguments": "{}"}}
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("response", "expected"),
|
||||
[
|
||||
(None, None),
|
||||
({"choices": "invalid"}, None),
|
||||
({"choices": [{"finish_reason": "length"}]}, "length"),
|
||||
({"choices": [{"stop_reason": "end_turn"}]}, "end_turn"),
|
||||
({"choices": [{}]}, None),
|
||||
(SimpleNamespace(stop_reason="end_turn"), "end_turn"),
|
||||
],
|
||||
)
|
||||
def test_response_finish_reason_handles_supported_shapes(response, expected):
|
||||
assert _response_finish_reason(response) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"request_data",
|
||||
[
|
||||
{"input": ["ssn 123-45-6789"], "litellm_metadata": {"user_api_key_request_route": "/vllm/v1/embeddings"}},
|
||||
{"input": [[1, 2, 3]], "litellm_metadata": {}},
|
||||
{"input": "confidential memo", "litellm_metadata": {}},
|
||||
{"input": "confidential memo"},
|
||||
],
|
||||
)
|
||||
def test_request_messages_not_resolved_for_unmapped_surfaces(request_data):
|
||||
"""Bodies from surfaces without a translation handler yield no messages, and never raise."""
|
||||
assert _request_structured_messages(request_data) is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("request_data", "expected"),
|
||||
[
|
||||
(
|
||||
{"messages": [{"role": "user", "content": "hi"}], "litellm_metadata": {}},
|
||||
[{"role": "user", "content": "hi"}],
|
||||
),
|
||||
(
|
||||
{
|
||||
"input": [{"role": "user", "content": "weather in Paris?"}],
|
||||
"litellm_metadata": {"user_api_key_request_route": "/v1/responses"},
|
||||
},
|
||||
[{"role": "user", "content": "weather in Paris?"}],
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_request_messages_resolved_for_mapped_surfaces(request_data, expected):
|
||||
assert _request_structured_messages(request_data) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("response", "expected"),
|
||||
[
|
||||
({"usage": {"input_tokens": 10, "output_tokens": 5}}, (10, 5)),
|
||||
({"usage": {"prompt_tokens": 7, "completion_tokens": 3}}, (7, 3)),
|
||||
(SimpleNamespace(usage=Usage(prompt_tokens=7, completion_tokens=3)), (7, 3)),
|
||||
({"usage": {"prompt_tokens": 0, "input_tokens": 99}}, (0, None)),
|
||||
({"usage": {}}, None),
|
||||
({}, None),
|
||||
],
|
||||
)
|
||||
def test_build_usage_handles_openai_and_anthropic_shapes(response, expected):
|
||||
usage = _build_usage(response)
|
||||
if expected is None:
|
||||
assert usage is None
|
||||
else:
|
||||
assert (usage.input_tokens, usage.output_tokens) == expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_non_streaming_response_reports_usage():
|
||||
g = _make_guardrail()
|
||||
g.async_handler.post.return_value = _mock_response("NONE")
|
||||
await g.apply_guardrail(
|
||||
inputs={"texts": ["hello"]},
|
||||
request_data={
|
||||
"model": "claude-sonnet-4-5",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"response": {
|
||||
"stop_reason": "end_turn",
|
||||
"usage": {"input_tokens": 10, "output_tokens": 5},
|
||||
},
|
||||
},
|
||||
input_type="response",
|
||||
logging_obj=_logging_obj(),
|
||||
)
|
||||
payload = _posted_payload(g)
|
||||
assert payload["usage"] == {"input_tokens": 10, "output_tokens": 5}
|
||||
assert payload["response"]["finish_reason"] == "end_turn"
|
||||
|
|
|
|||
|
|
@ -4,7 +4,10 @@ import pytest
|
|||
|
||||
import litellm
|
||||
from litellm.caching import DualCache
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
effective_skip_system_message_for_guardrail,
|
||||
|
|
@ -1490,3 +1493,109 @@ class TestStreamingTransform:
|
|||
|
||||
# None holdback treated as 0: full text emitted, no crash.
|
||||
assert "".join(_delta_text(i) for i in out) == "ABCDEF"
|
||||
|
||||
|
||||
def _applied_guardrails(data: dict) -> list:
|
||||
for key in ("metadata", "litellm_metadata"):
|
||||
meta = data.get(key)
|
||||
if isinstance(meta, dict) and isinstance(meta.get("applied_guardrails"), list):
|
||||
return meta["applied_guardrails"]
|
||||
return []
|
||||
|
||||
|
||||
class _TextsOnlyTranslation(BaseTranslation):
|
||||
"""Mimics a passthrough handler: hands the guardrail only `texts`, never
|
||||
structured_messages, so a structured_messages-based guardrail no-ops."""
|
||||
|
||||
async def process_input_messages(self, data, guardrail_to_apply, litellm_logging_obj=None): # type: ignore[override]
|
||||
await guardrail_to_apply.apply_guardrail(
|
||||
inputs={"texts": ["payload"]},
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
return data
|
||||
|
||||
async def process_output_response( # type: ignore[override]
|
||||
self,
|
||||
response,
|
||||
guardrail_to_apply,
|
||||
litellm_logging_obj=None,
|
||||
user_api_key_dict=None,
|
||||
request_data=None,
|
||||
):
|
||||
return response
|
||||
|
||||
|
||||
class _SelfLoggingGuardrail(CustomGuardrail):
|
||||
records_own_guardrail_information = True
|
||||
|
||||
def __init__(self, *, self_add: bool):
|
||||
super().__init__(guardrail_name="self-logging")
|
||||
self._self_add = self_add
|
||||
|
||||
def should_run_guardrail(self, data, event_type): # type: ignore[override]
|
||||
return True
|
||||
|
||||
@log_guardrail_information
|
||||
async def apply_guardrail(self, inputs, request_data, input_type, **kwargs):
|
||||
if self._self_add:
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_to_applied_guardrails_header,
|
||||
)
|
||||
|
||||
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
|
||||
return inputs
|
||||
|
||||
|
||||
class _AutoLoggingGuardrail(CustomGuardrail):
|
||||
def __init__(self):
|
||||
super().__init__(guardrail_name="auto-logging")
|
||||
|
||||
def should_run_guardrail(self, data, event_type): # type: ignore[override]
|
||||
return True
|
||||
|
||||
@log_guardrail_information
|
||||
async def apply_guardrail(self, inputs, request_data, input_type, **kwargs):
|
||||
return inputs
|
||||
|
||||
|
||||
class TestAppliedGuardrailsReflectsExecution:
|
||||
"""The unified hook must not auto-mark a self-logging guardrail
|
||||
(records_own_guardrail_information) as applied; such a guardrail owns that
|
||||
decision and marks itself only when it actually ran (LIT-4650). Ordinary
|
||||
guardrails are still auto-marked by the hook after dispatch."""
|
||||
|
||||
@staticmethod
|
||||
def _data(guardrail):
|
||||
return {
|
||||
"guardrail_to_apply": guardrail,
|
||||
"model": "gpt-4o",
|
||||
"messages": [{"role": "user", "content": "hello world"}],
|
||||
}
|
||||
|
||||
async def _run(self, guardrail):
|
||||
unified_module.endpoint_guardrail_translation_mappings = {CallTypes.pass_through: _TextsOnlyTranslation}
|
||||
data = self._data(guardrail)
|
||||
await UnifiedLLMGuardrails().async_pre_call_hook(
|
||||
user_api_key_dict=None,
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type=CallTypes.pass_through.value,
|
||||
)
|
||||
return data
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_self_logging_guardrail_is_not_auto_marked_applied(self):
|
||||
data = await self._run(_SelfLoggingGuardrail(self_add=False))
|
||||
assert "self-logging" not in _applied_guardrails(data)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_self_logging_guardrail_that_self_marks_is_applied(self):
|
||||
data = await self._run(_SelfLoggingGuardrail(self_add=True))
|
||||
assert "self-logging" in _applied_guardrails(data)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ordinary_guardrail_is_auto_marked_applied(self):
|
||||
data = await self._run(_AutoLoggingGuardrail())
|
||||
assert "auto-logging" in _applied_guardrails(data)
|
||||
|
|
|
|||
|
|
@ -1504,50 +1504,6 @@ def test_team_info_masking():
|
|||
assert "public-test-key" not in str(exc_info.value)
|
||||
|
||||
|
||||
def test_embedding_input_array_of_tokens(client_no_auth):
|
||||
"""
|
||||
Test to bypass decoding input as array of tokens for selected providers
|
||||
|
||||
Ref: https://github.com/BerriAI/litellm/issues/10113
|
||||
"""
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
# The client_no_auth fixture should initialize the router
|
||||
# Assert this to catch any router initialization regressions
|
||||
assert proxy_server.llm_router is not None, (
|
||||
"llm_router is None after client_no_auth fixture initialized. "
|
||||
"This indicates a router initialization issue that should be investigated."
|
||||
)
|
||||
|
||||
try:
|
||||
with mock.patch.object(
|
||||
proxy_server.llm_router,
|
||||
"aembedding",
|
||||
return_value=example_embedding_result,
|
||||
) as mock_aembedding:
|
||||
test_data = {
|
||||
"model": "vllm_embed_model",
|
||||
"input": [[2046, 13269, 158208]],
|
||||
}
|
||||
|
||||
response = client_no_auth.post("/v1/embeddings", json=test_data)
|
||||
|
||||
# Assert that aembedding was called, and that input was not modified
|
||||
mock_aembedding.assert_called_once()
|
||||
call_args, call_kwargs = mock_aembedding.call_args
|
||||
assert call_kwargs["model"] == "vllm_embed_model"
|
||||
assert call_kwargs["input"] == [[2046, 13269, 158208]]
|
||||
|
||||
assert response.status_code == 200
|
||||
result = response.json()
|
||||
print(len(result["data"][0]["embedding"]))
|
||||
assert (
|
||||
len(result["data"][0]["embedding"]) > 10
|
||||
) # this usually has len==1536 so
|
||||
except Exception as e:
|
||||
pytest.fail(f"LiteLLM Proxy test failed. Exception - {str(e)}")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_all_team_models():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -12,24 +12,25 @@ from litellm.proxy.route_llm_request import ProxyModelNotFoundError, route_reque
|
|||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route_type",
|
||||
"route_type, required_body_params",
|
||||
[
|
||||
"atext_completion",
|
||||
"acompletion",
|
||||
"aembedding",
|
||||
"aimage_generation",
|
||||
"aspeech",
|
||||
"atranscription",
|
||||
"amoderation",
|
||||
"arerank",
|
||||
("atext_completion", {}),
|
||||
("acompletion", {"messages": [{"role": "user", "content": "Hello"}]}),
|
||||
("aembedding", {"input": "Hello"}),
|
||||
("aimage_generation", {}),
|
||||
("aspeech", {}),
|
||||
("atranscription", {}),
|
||||
("amoderation", {}),
|
||||
("arerank", {}),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_dynamic_credentials(route_type):
|
||||
async def test_route_request_dynamic_credentials(route_type, required_body_params):
|
||||
data = {
|
||||
"model": "openai/gpt-4o-mini-2024-07-18",
|
||||
"api_key": "my-bad-key",
|
||||
"api_base": "https://api.openai.com/v1 ",
|
||||
**required_body_params,
|
||||
}
|
||||
llm_router = MagicMock()
|
||||
# Ensure that the dynamic method exists on the llm_router mock.
|
||||
|
|
@ -887,3 +888,59 @@ async def test_route_request_override_enable_tag_filtering_beats_body_value():
|
|||
|
||||
call_kwargs = llm_router.acompletion.call_args[1]
|
||||
assert call_kwargs["enable_tag_filtering"] is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route_type, param, route",
|
||||
[
|
||||
("acompletion", "messages", "/chat/completions"),
|
||||
("aembedding", "input", "/embeddings"),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize("data_extra", [{}, {"messages": None, "input": None}])
|
||||
def test_raise_if_required_body_param_missing_rejects_missing_param(route_type, param, route, data_extra):
|
||||
from litellm.proxy.route_llm_request import (
|
||||
ProxyMissingRequiredParamError,
|
||||
raise_if_required_body_param_missing,
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
|
||||
raise_if_required_body_param_missing(route_type=route_type, data={"model": "gpt-4o", **data_extra})
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert exc_info.value.param == param
|
||||
assert exc_info.value.type == "invalid_request_error"
|
||||
assert exc_info.value.detail == {"error": f"{route}: Missing required parameter: '{param}'."}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route_type, data",
|
||||
[
|
||||
("acompletion", {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]}),
|
||||
("acompletion", {"model": "gpt-4o", "messages": []}),
|
||||
("atext_completion", {"model": "gpt-4o"}),
|
||||
("aembedding", {"model": "text-embedding-3-small", "input": "hi"}),
|
||||
("arerank", {"model": "rerank-model"}),
|
||||
("aimage_generation", {"model": "dall-e-3"}),
|
||||
],
|
||||
)
|
||||
def test_raise_if_required_body_param_missing_allows_valid_requests(route_type, data):
|
||||
from litellm.proxy.route_llm_request import raise_if_required_body_param_missing
|
||||
|
||||
raise_if_required_body_param_missing(route_type=route_type, data=data)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_rejects_chat_completion_without_messages():
|
||||
"""A /chat/completions body without `messages` used to splat into
|
||||
Router.acompletion() and surface the resulting TypeError as a 500."""
|
||||
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
|
||||
|
||||
llm_router = MagicMock()
|
||||
|
||||
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
|
||||
await route_request({"model": "gpt-4o"}, llm_router, None, "acompletion")
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert exc_info.value.param == "messages"
|
||||
llm_router.acompletion.assert_not_called()
|
||||
|
|
|
|||
|
|
@ -198,6 +198,54 @@ async def test_aresponses_azure_shell_tool_400_maps_to_bad_request_error():
|
|||
assert "not supported" in str(excinfo.value).lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aresponses_drops_stream_options():
|
||||
"""The Responses API rejects include_usage, so include_usage-only stream_options must never reach the wire."""
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_post:
|
||||
mock_post.return_value = MockResponse(
|
||||
_minimal_responses_api_payload("resp_stream_options_test", "gpt-5.5"), 200
|
||||
)
|
||||
|
||||
await litellm.aresponses(
|
||||
model="openai/gpt-5.5",
|
||||
api_key="fake-api-key",
|
||||
input="hi",
|
||||
stream_options={"include_usage": True},
|
||||
)
|
||||
|
||||
mock_post.assert_called_once()
|
||||
post_kwargs = mock_post.call_args.kwargs
|
||||
request_body = post_kwargs["json"] if "json" in post_kwargs else json.loads(post_kwargs["data"])
|
||||
assert "stream_options" not in request_body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aresponses_keeps_include_obfuscation_in_stream_options():
|
||||
"""include_obfuscation is a valid Responses API stream option and must survive the include_usage strip."""
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_post:
|
||||
mock_post.return_value = MockResponse(
|
||||
_minimal_responses_api_payload("resp_stream_options_obfuscation", "gpt-5.5"), 200
|
||||
)
|
||||
|
||||
await litellm.aresponses(
|
||||
model="openai/gpt-5.5",
|
||||
api_key="fake-api-key",
|
||||
input="hi",
|
||||
stream_options={"include_usage": True, "include_obfuscation": False},
|
||||
)
|
||||
|
||||
mock_post.assert_called_once()
|
||||
post_kwargs = mock_post.call_args.kwargs
|
||||
request_body = post_kwargs["json"] if "json" in post_kwargs else json.loads(post_kwargs["data"])
|
||||
assert request_body["stream_options"] == {"include_obfuscation": False}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aresponses_request_level_drop_params_drops_bedrock_mantle_service_tier(
|
||||
monkeypatch,
|
||||
|
|
|
|||
|
|
@ -25,8 +25,17 @@ vi.mock("@/components/networking", () => ({
|
|||
|
||||
// Mock the child components to simplify testing
|
||||
vi.mock("@/components/activity_metrics", () => ({
|
||||
ActivityMetrics: () => <div>Activity Metrics</div>,
|
||||
processActivityData: () => ({ data: [], metadata: {} }),
|
||||
ActivityMetrics: ({ modelMetrics }: { modelMetrics?: { __source?: string } }) => (
|
||||
<div>
|
||||
<span>Activity Metrics</span>
|
||||
<span>{`metrics-source:${modelMetrics?.__source ?? "none"}`}</span>
|
||||
</div>
|
||||
),
|
||||
processActivityData: (_data: unknown, key: string) => ({ __source: key }),
|
||||
}));
|
||||
|
||||
vi.mock("../EndpointUsage/EndpointUsage", () => ({
|
||||
default: () => <div>Endpoint Usage Panel</div>,
|
||||
}));
|
||||
|
||||
vi.mock("@/components/UsagePage/components/EntityUsage/TopKeyView", () => ({
|
||||
|
|
@ -481,6 +490,54 @@ describe("EntityUsage", () => {
|
|||
expect(screen.getAllByText("Activity Metrics")[1]).toBeInTheDocument();
|
||||
});
|
||||
|
||||
const selectedPanels = (container: HTMLElement) =>
|
||||
Array.from(container.querySelectorAll("div.tremor-TabPanel-root")).filter(
|
||||
(panel) => panel.getAttribute("aria-selected") === "true",
|
||||
);
|
||||
|
||||
it.each([
|
||||
["Cost", "Tag Spend Overview"],
|
||||
["Model Activity", "metrics-source:models"],
|
||||
["Key Activity", "metrics-source:api_keys"],
|
||||
["Endpoint Activity", "Endpoint Usage Panel"],
|
||||
])("shows only the %s panel for a non-team entity type", async (tabLabel, marker) => {
|
||||
const { container } = render(<EntityUsage {...defaultProps} />);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(mockTagDailyActivityCall).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
act(() => {
|
||||
fireEvent.click(screen.getByText(tabLabel));
|
||||
});
|
||||
|
||||
const selected = selectedPanels(container);
|
||||
expect(selected).toHaveLength(1);
|
||||
expect(selected[0].textContent).toContain(marker);
|
||||
});
|
||||
|
||||
it.each([
|
||||
["Cost", "Team Spend Overview"],
|
||||
["Model Activity", "metrics-source:models"],
|
||||
["Agent Activity", "metrics-source:entities"],
|
||||
["Key Activity", "metrics-source:api_keys"],
|
||||
["Endpoint Activity", "Endpoint Usage Panel"],
|
||||
])("shows only the %s panel for the team entity type", async (tabLabel, marker) => {
|
||||
const { container } = render(<EntityUsage {...defaultProps} entityType="team" />);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(mockTeamDailyActivityCall).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
act(() => {
|
||||
fireEvent.click(screen.getByText(tabLabel));
|
||||
});
|
||||
|
||||
const selected = selectedPanels(container);
|
||||
expect(selected).toHaveLength(1);
|
||||
expect(selected[0].textContent).toContain(marker);
|
||||
});
|
||||
|
||||
it("should handle empty data gracefully", async () => {
|
||||
const emptyData = {
|
||||
results: [],
|
||||
|
|
|
|||
|
|
@ -25,7 +25,7 @@ import {
|
|||
} from "@tremor/react";
|
||||
import { ExportOutlined, LoadingOutlined } from "@ant-design/icons";
|
||||
import { Alert, Button } from "antd";
|
||||
import React, { useMemo, useState } from "react";
|
||||
import React, { type ReactNode, useMemo, useState } from "react";
|
||||
import TeamMultiSelect from "@/components/common_components/team_multi_select";
|
||||
import { ActivityMetrics, processActivityData } from "@/components/activity_metrics";
|
||||
import { UsageExportHeader } from "@/components/EntityUsageExport";
|
||||
|
|
@ -406,6 +406,304 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
|
|||
|
||||
const capitalizedEntityLabel = entityType.charAt(0).toUpperCase() + entityType.slice(1);
|
||||
|
||||
const costPanel = (
|
||||
<Grid numItems={2} className="gap-2 w-full">
|
||||
{/* Total Spend Card */}
|
||||
<Col numColSpan={2}>
|
||||
<Card>
|
||||
<Title>{capitalizedEntityLabel} Spend Overview</Title>
|
||||
<Grid numItems={5} className="gap-4 mt-4">
|
||||
<Card>
|
||||
<Title>Total Spend</Title>
|
||||
<Text className="text-2xl font-bold mt-2">
|
||||
${formatNumberWithCommas(spendData.metadata.total_spend, 2)}
|
||||
</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Total Requests</Title>
|
||||
<Text className="text-2xl font-bold mt-2">{spendData.metadata.total_api_requests.toLocaleString()}</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Successful Requests</Title>
|
||||
<Text className="text-2xl font-bold mt-2 text-green-600">
|
||||
{spendData.metadata.total_successful_requests.toLocaleString()}
|
||||
</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Failed Requests</Title>
|
||||
<Text className="text-2xl font-bold mt-2 text-red-600">
|
||||
{spendData.metadata.total_failed_requests.toLocaleString()}
|
||||
</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Total Tokens</Title>
|
||||
<Text className="text-2xl font-bold mt-2">{spendData.metadata.total_tokens.toLocaleString()}</Text>
|
||||
</Card>
|
||||
</Grid>
|
||||
</Card>
|
||||
</Col>
|
||||
|
||||
{/* Daily Spend Chart */}
|
||||
<Col numColSpan={2}>
|
||||
<ShadcnCard>
|
||||
<CardHeader>
|
||||
<CardTitle className="text-base font-semibold">Daily Spend</CardTitle>
|
||||
</CardHeader>
|
||||
<CardContent>
|
||||
<BarChart
|
||||
data={[...spendData.results].sort((a, b) => new Date(a.date).getTime() - new Date(b.date).getTime())}
|
||||
index="date"
|
||||
categories={["metrics.spend"]}
|
||||
colors={["cyan"]}
|
||||
valueFormatter={valueFormatterSpend}
|
||||
yAxisWidth={100}
|
||||
showLegend={false}
|
||||
customTooltip={({ payload, active }) => {
|
||||
if (!active || !payload?.[0]) return null;
|
||||
const data = payload[0].payload;
|
||||
const entityCount = Object.keys(data.breakdown.entities || {}).length;
|
||||
return (
|
||||
<div className="bg-white p-4 shadow-lg rounded-lg border">
|
||||
<p className="font-bold">{data.date}</p>
|
||||
<p className="text-cyan-500">Total Spend: ${formatNumberWithCommas(data.metrics.spend, 2)}</p>
|
||||
<p className="text-gray-600">Total Requests: {data.metrics.api_requests}</p>
|
||||
<p className="text-gray-600">Successful: {data.metrics.successful_requests}</p>
|
||||
<p className="text-gray-600">Failed: {data.metrics.failed_requests}</p>
|
||||
<p className="text-gray-600">Total Tokens: {data.metrics.total_tokens}</p>
|
||||
<p className="text-gray-600">
|
||||
Total {capitalizedEntityLabel}s: {entityCount}
|
||||
</p>
|
||||
<div className="mt-2 border-t pt-2">
|
||||
<p className="font-semibold">Spend by {capitalizedEntityLabel}:</p>
|
||||
{Object.entries(data.breakdown.entities || {})
|
||||
.sort(([, a], [, b]) => {
|
||||
const spendA = (a as EntityMetrics).metrics.spend;
|
||||
const spendB = (b as EntityMetrics).metrics.spend;
|
||||
return spendB - spendA;
|
||||
})
|
||||
.slice(0, 5)
|
||||
.map(([entity, entityData]) => {
|
||||
const metrics = entityData as EntityMetrics;
|
||||
return (
|
||||
<p key={entity} className="text-sm text-gray-600">
|
||||
{getEntityLabel(entity, metrics.metadata)}: $
|
||||
{formatNumberWithCommas(metrics.metrics.spend, 2)}
|
||||
</p>
|
||||
);
|
||||
})}
|
||||
{entityCount > 5 && <p className="text-sm text-gray-500 italic">...and {entityCount - 5} more</p>}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}}
|
||||
/>
|
||||
</CardContent>
|
||||
</ShadcnCard>
|
||||
</Col>
|
||||
|
||||
{/* Entity Breakdown Section */}
|
||||
<Col numColSpan={2}>
|
||||
<Card>
|
||||
<div className="flex flex-col space-y-4">
|
||||
<div className="flex flex-col space-y-2">
|
||||
<Title>Spend Per {capitalizedEntityLabel}</Title>
|
||||
<Subtitle className="text-xs">Showing Top 5 by Spend</Subtitle>
|
||||
<div className="flex items-center text-sm text-gray-500">
|
||||
<span>Get Started by Tracking cost per {capitalizedEntityLabel} </span>
|
||||
<a
|
||||
href="https://docs.litellm.ai/docs/proxy/enterprise#spend-tracking"
|
||||
className="text-blue-500 hover:text-blue-700 ml-1"
|
||||
>
|
||||
here
|
||||
</a>
|
||||
</div>
|
||||
</div>
|
||||
<Grid numItems={2} className="gap-6">
|
||||
<Col numColSpan={1}>
|
||||
<BarChart
|
||||
className="mt-4 h-52"
|
||||
data={getProcessedEntityBreakdownForChart()}
|
||||
index="metadata.alias_display"
|
||||
categories={["metrics.spend"]}
|
||||
colors={["cyan"]}
|
||||
valueFormatter={valueFormatterSpend}
|
||||
layout="vertical"
|
||||
showLegend={false}
|
||||
yAxisWidth={150}
|
||||
customTooltip={({ payload, active }) => {
|
||||
if (!active || !payload?.[0]) return null;
|
||||
const data = payload[0].payload;
|
||||
return (
|
||||
<div className="bg-white p-4 shadow-lg rounded-lg border">
|
||||
<p className="font-bold">{data.metadata.alias}</p>
|
||||
<p className="text-cyan-500">Spend: ${formatNumberWithCommas(data.metrics.spend, 4)}</p>
|
||||
<p className="text-gray-600">Requests: {data.metrics.api_requests.toLocaleString()}</p>
|
||||
<p className="text-green-600">
|
||||
Successful: {data.metrics.successful_requests.toLocaleString()}
|
||||
</p>
|
||||
<p className="text-red-600">Failed: {data.metrics.failed_requests.toLocaleString()}</p>
|
||||
<p className="text-gray-600">Tokens: {data.metrics.total_tokens.toLocaleString()}</p>
|
||||
</div>
|
||||
);
|
||||
}}
|
||||
/>
|
||||
</Col>
|
||||
<Col numColSpan={1}>
|
||||
<div className="h-52 overflow-y-auto">
|
||||
<Table>
|
||||
<TableHead>
|
||||
<TableRow>
|
||||
<TableHeaderCell>{capitalizedEntityLabel}</TableHeaderCell>
|
||||
<TableHeaderCell>Spend</TableHeaderCell>
|
||||
<TableHeaderCell className="text-green-600">Successful</TableHeaderCell>
|
||||
<TableHeaderCell className="text-red-600">Failed</TableHeaderCell>
|
||||
<TableHeaderCell>Tokens</TableHeaderCell>
|
||||
</TableRow>
|
||||
</TableHead>
|
||||
<TableBody>
|
||||
{getEntityBreakdown()
|
||||
.filter((entity) => entity.metrics.spend > 0)
|
||||
.map((entity) => (
|
||||
<TableRow key={entity.metadata.id}>
|
||||
<TableCell>{entity.metadata.alias}</TableCell>
|
||||
<TableCell>
|
||||
<MoneyCell value={entity.metrics.spend} decimals={4} />
|
||||
</TableCell>
|
||||
<TableCell className="text-green-600">
|
||||
{entity.metrics.successful_requests.toLocaleString()}
|
||||
</TableCell>
|
||||
<TableCell className="text-red-600">
|
||||
{entity.metrics.failed_requests.toLocaleString()}
|
||||
</TableCell>
|
||||
<TableCell>{entity.metrics.total_tokens.toLocaleString()}</TableCell>
|
||||
</TableRow>
|
||||
))}
|
||||
</TableBody>
|
||||
</Table>
|
||||
</div>
|
||||
</Col>
|
||||
</Grid>
|
||||
</div>
|
||||
</Card>
|
||||
</Col>
|
||||
|
||||
{/* Top API Keys */}
|
||||
<Col numColSpan={1}>
|
||||
<Card>
|
||||
<Title>Top Virtual Keys</Title>
|
||||
<TopKeyView
|
||||
topKeys={getTopAPIKeys()}
|
||||
teams={null}
|
||||
showTags={entityType === "tag"}
|
||||
topKeysLimit={topKeysLimit}
|
||||
setTopKeysLimit={setTopKeysLimit}
|
||||
/>
|
||||
</Card>
|
||||
</Col>
|
||||
|
||||
{/* Top Models */}
|
||||
<Col numColSpan={1}>
|
||||
<Card>
|
||||
<Title>{entityType === "agent" ? "Top Agents" : "Top Models"}</Title>
|
||||
<TopModelView
|
||||
topModels={getTopModels()}
|
||||
topModelsLimit={topModelsLimit}
|
||||
setTopModelsLimit={setTopModelsLimit}
|
||||
/>
|
||||
</Card>
|
||||
</Col>
|
||||
|
||||
{/* Top Agents - only for team entity type */}
|
||||
{entityType === "team" && (
|
||||
<Col numColSpan={2}>
|
||||
<Card>
|
||||
<Title>Top Agents Driving Spend</Title>
|
||||
<TopModelView
|
||||
topModels={getTopAgents()}
|
||||
topModelsLimit={topAgentsLimit}
|
||||
setTopModelsLimit={setTopAgentsLimit}
|
||||
/>
|
||||
</Card>
|
||||
</Col>
|
||||
)}
|
||||
|
||||
{/* Spend by Provider */}
|
||||
<Col numColSpan={2}>
|
||||
<Card>
|
||||
<div className="flex flex-col space-y-4">
|
||||
<Title>Provider Usage</Title>
|
||||
<Grid numItems={2}>
|
||||
<Col numColSpan={1}>
|
||||
<DonutChart
|
||||
className="mt-4 h-40"
|
||||
data={getProviderSpend()}
|
||||
index="provider"
|
||||
category="spend"
|
||||
valueFormatter={(value) => `$${formatNumberWithCommas(value, 2)}`}
|
||||
colors={["cyan", "blue", "indigo", "violet", "purple"]}
|
||||
showLabel
|
||||
startAngle={90}
|
||||
endAngle={-270}
|
||||
/>
|
||||
</Col>
|
||||
<Col numColSpan={1}>
|
||||
<Table>
|
||||
<TableHead>
|
||||
<TableRow>
|
||||
<TableHeaderCell>Provider</TableHeaderCell>
|
||||
<TableHeaderCell>Spend</TableHeaderCell>
|
||||
<TableHeaderCell className="text-green-600">Successful</TableHeaderCell>
|
||||
<TableHeaderCell className="text-red-600">Failed</TableHeaderCell>
|
||||
<TableHeaderCell>Tokens</TableHeaderCell>
|
||||
</TableRow>
|
||||
</TableHead>
|
||||
<TableBody>
|
||||
{getProviderSpend().map((provider) => (
|
||||
<TableRow key={provider.provider}>
|
||||
<TableCell>
|
||||
<div className="flex items-center space-x-2">
|
||||
{provider.provider && <Logo provider={provider.provider} className="w-4 h-4" />}
|
||||
<span>{provider.provider}</span>
|
||||
</div>
|
||||
</TableCell>
|
||||
<TableCell>
|
||||
<MoneyCell value={provider.spend} decimals={2} />
|
||||
</TableCell>
|
||||
<TableCell className="text-green-600">
|
||||
{provider.successful_requests.toLocaleString()}
|
||||
</TableCell>
|
||||
<TableCell className="text-red-600">{provider.failed_requests.toLocaleString()}</TableCell>
|
||||
<TableCell>{provider.tokens.toLocaleString()}</TableCell>
|
||||
</TableRow>
|
||||
))}
|
||||
</TableBody>
|
||||
</Table>
|
||||
</Col>
|
||||
</Grid>
|
||||
</div>
|
||||
</Card>
|
||||
</Col>
|
||||
</Grid>
|
||||
);
|
||||
|
||||
const tabs: readonly { key: string; label: string; content: ReactNode }[] = [
|
||||
{ key: "cost", label: "Cost", content: costPanel },
|
||||
{
|
||||
key: "models",
|
||||
label: entityType === "agent" ? "Request / Token Consumption" : "Model Activity",
|
||||
content: <ActivityMetrics modelMetrics={modelMetrics} hidePromptCachingMetrics={entityType === "agent"} />,
|
||||
},
|
||||
...(entityType === "team"
|
||||
? [{ key: "agents", label: "Agent Activity", content: <ActivityMetrics modelMetrics={agentMetrics} /> }]
|
||||
: []),
|
||||
{
|
||||
key: "keys",
|
||||
label: "Key Activity",
|
||||
content: <ActivityMetrics modelMetrics={keyMetrics} hidePromptCachingMetrics={entityType === "agent"} />,
|
||||
},
|
||||
{ key: "endpoints", label: "Endpoint Activity", content: <EndpointUsage userSpendData={spendData} /> },
|
||||
];
|
||||
|
||||
return (
|
||||
<div style={{ width: "100%" }} className="relative">
|
||||
{isFetchingMore && (
|
||||
|
|
@ -501,320 +799,14 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
|
|||
/>
|
||||
<TabGroup>
|
||||
<TabList variant="solid" className="mt-1">
|
||||
<Tab>Cost</Tab>
|
||||
<Tab>{entityType === "agent" ? "Request / Token Consumption" : "Model Activity"}</Tab>
|
||||
{entityType === "team" ? <Tab>Agent Activity</Tab> : <></>}
|
||||
<Tab>Key Activity</Tab>
|
||||
<Tab>Endpoint Activity</Tab>
|
||||
{tabs.map(({ key, label }) => (
|
||||
<Tab key={key}>{label}</Tab>
|
||||
))}
|
||||
</TabList>
|
||||
<TabPanels>
|
||||
<TabPanel>
|
||||
<Grid numItems={2} className="gap-2 w-full">
|
||||
{/* Total Spend Card */}
|
||||
<Col numColSpan={2}>
|
||||
<Card>
|
||||
<Title>{capitalizedEntityLabel} Spend Overview</Title>
|
||||
<Grid numItems={5} className="gap-4 mt-4">
|
||||
<Card>
|
||||
<Title>Total Spend</Title>
|
||||
<Text className="text-2xl font-bold mt-2">
|
||||
${formatNumberWithCommas(spendData.metadata.total_spend, 2)}
|
||||
</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Total Requests</Title>
|
||||
<Text className="text-2xl font-bold mt-2">
|
||||
{spendData.metadata.total_api_requests.toLocaleString()}
|
||||
</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Successful Requests</Title>
|
||||
<Text className="text-2xl font-bold mt-2 text-green-600">
|
||||
{spendData.metadata.total_successful_requests.toLocaleString()}
|
||||
</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Failed Requests</Title>
|
||||
<Text className="text-2xl font-bold mt-2 text-red-600">
|
||||
{spendData.metadata.total_failed_requests.toLocaleString()}
|
||||
</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Total Tokens</Title>
|
||||
<Text className="text-2xl font-bold mt-2">
|
||||
{spendData.metadata.total_tokens.toLocaleString()}
|
||||
</Text>
|
||||
</Card>
|
||||
</Grid>
|
||||
</Card>
|
||||
</Col>
|
||||
|
||||
{/* Daily Spend Chart */}
|
||||
<Col numColSpan={2}>
|
||||
<ShadcnCard>
|
||||
<CardHeader>
|
||||
<CardTitle className="text-base font-semibold">Daily Spend</CardTitle>
|
||||
</CardHeader>
|
||||
<CardContent>
|
||||
<BarChart
|
||||
data={[...spendData.results].sort(
|
||||
(a, b) => new Date(a.date).getTime() - new Date(b.date).getTime(),
|
||||
)}
|
||||
index="date"
|
||||
categories={["metrics.spend"]}
|
||||
colors={["cyan"]}
|
||||
valueFormatter={valueFormatterSpend}
|
||||
yAxisWidth={100}
|
||||
showLegend={false}
|
||||
customTooltip={({ payload, active }) => {
|
||||
if (!active || !payload?.[0]) return null;
|
||||
const data = payload[0].payload;
|
||||
const entityCount = Object.keys(data.breakdown.entities || {}).length;
|
||||
return (
|
||||
<div className="bg-white p-4 shadow-lg rounded-lg border">
|
||||
<p className="font-bold">{data.date}</p>
|
||||
<p className="text-cyan-500">
|
||||
Total Spend: ${formatNumberWithCommas(data.metrics.spend, 2)}
|
||||
</p>
|
||||
<p className="text-gray-600">Total Requests: {data.metrics.api_requests}</p>
|
||||
<p className="text-gray-600">Successful: {data.metrics.successful_requests}</p>
|
||||
<p className="text-gray-600">Failed: {data.metrics.failed_requests}</p>
|
||||
<p className="text-gray-600">Total Tokens: {data.metrics.total_tokens}</p>
|
||||
<p className="text-gray-600">
|
||||
Total {capitalizedEntityLabel}s: {entityCount}
|
||||
</p>
|
||||
<div className="mt-2 border-t pt-2">
|
||||
<p className="font-semibold">Spend by {capitalizedEntityLabel}:</p>
|
||||
{Object.entries(data.breakdown.entities || {})
|
||||
.sort(([, a], [, b]) => {
|
||||
const spendA = (a as EntityMetrics).metrics.spend;
|
||||
const spendB = (b as EntityMetrics).metrics.spend;
|
||||
return spendB - spendA;
|
||||
})
|
||||
.slice(0, 5)
|
||||
.map(([entity, entityData]) => {
|
||||
const metrics = entityData as EntityMetrics;
|
||||
return (
|
||||
<p key={entity} className="text-sm text-gray-600">
|
||||
{getEntityLabel(entity, metrics.metadata)}: $
|
||||
{formatNumberWithCommas(metrics.metrics.spend, 2)}
|
||||
</p>
|
||||
);
|
||||
})}
|
||||
{entityCount > 5 && (
|
||||
<p className="text-sm text-gray-500 italic">...and {entityCount - 5} more</p>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}}
|
||||
/>
|
||||
</CardContent>
|
||||
</ShadcnCard>
|
||||
</Col>
|
||||
|
||||
{/* Entity Breakdown Section */}
|
||||
<Col numColSpan={2}>
|
||||
<Card>
|
||||
<div className="flex flex-col space-y-4">
|
||||
<div className="flex flex-col space-y-2">
|
||||
<Title>Spend Per {capitalizedEntityLabel}</Title>
|
||||
<Subtitle className="text-xs">Showing Top 5 by Spend</Subtitle>
|
||||
<div className="flex items-center text-sm text-gray-500">
|
||||
<span>Get Started by Tracking cost per {capitalizedEntityLabel} </span>
|
||||
<a
|
||||
href="https://docs.litellm.ai/docs/proxy/enterprise#spend-tracking"
|
||||
className="text-blue-500 hover:text-blue-700 ml-1"
|
||||
>
|
||||
here
|
||||
</a>
|
||||
</div>
|
||||
</div>
|
||||
<Grid numItems={2} className="gap-6">
|
||||
<Col numColSpan={1}>
|
||||
<BarChart
|
||||
className="mt-4 h-52"
|
||||
data={getProcessedEntityBreakdownForChart()}
|
||||
index="metadata.alias_display"
|
||||
categories={["metrics.spend"]}
|
||||
colors={["cyan"]}
|
||||
valueFormatter={valueFormatterSpend}
|
||||
layout="vertical"
|
||||
showLegend={false}
|
||||
yAxisWidth={150}
|
||||
customTooltip={({ payload, active }) => {
|
||||
if (!active || !payload?.[0]) return null;
|
||||
const data = payload[0].payload;
|
||||
return (
|
||||
<div className="bg-white p-4 shadow-lg rounded-lg border">
|
||||
<p className="font-bold">{data.metadata.alias}</p>
|
||||
<p className="text-cyan-500">Spend: ${formatNumberWithCommas(data.metrics.spend, 4)}</p>
|
||||
<p className="text-gray-600">Requests: {data.metrics.api_requests.toLocaleString()}</p>
|
||||
<p className="text-green-600">
|
||||
Successful: {data.metrics.successful_requests.toLocaleString()}
|
||||
</p>
|
||||
<p className="text-red-600">Failed: {data.metrics.failed_requests.toLocaleString()}</p>
|
||||
<p className="text-gray-600">Tokens: {data.metrics.total_tokens.toLocaleString()}</p>
|
||||
</div>
|
||||
);
|
||||
}}
|
||||
/>
|
||||
</Col>
|
||||
<Col numColSpan={1}>
|
||||
<div className="h-52 overflow-y-auto">
|
||||
<Table>
|
||||
<TableHead>
|
||||
<TableRow>
|
||||
<TableHeaderCell>{capitalizedEntityLabel}</TableHeaderCell>
|
||||
<TableHeaderCell>Spend</TableHeaderCell>
|
||||
<TableHeaderCell className="text-green-600">Successful</TableHeaderCell>
|
||||
<TableHeaderCell className="text-red-600">Failed</TableHeaderCell>
|
||||
<TableHeaderCell>Tokens</TableHeaderCell>
|
||||
</TableRow>
|
||||
</TableHead>
|
||||
<TableBody>
|
||||
{getEntityBreakdown()
|
||||
.filter((entity) => entity.metrics.spend > 0)
|
||||
.map((entity) => (
|
||||
<TableRow key={entity.metadata.id}>
|
||||
<TableCell>{entity.metadata.alias}</TableCell>
|
||||
<TableCell>
|
||||
<MoneyCell value={entity.metrics.spend} decimals={4} />
|
||||
</TableCell>
|
||||
<TableCell className="text-green-600">
|
||||
{entity.metrics.successful_requests.toLocaleString()}
|
||||
</TableCell>
|
||||
<TableCell className="text-red-600">
|
||||
{entity.metrics.failed_requests.toLocaleString()}
|
||||
</TableCell>
|
||||
<TableCell>{entity.metrics.total_tokens.toLocaleString()}</TableCell>
|
||||
</TableRow>
|
||||
))}
|
||||
</TableBody>
|
||||
</Table>
|
||||
</div>
|
||||
</Col>
|
||||
</Grid>
|
||||
</div>
|
||||
</Card>
|
||||
</Col>
|
||||
|
||||
{/* Top API Keys */}
|
||||
<Col numColSpan={1}>
|
||||
<Card>
|
||||
<Title>Top Virtual Keys</Title>
|
||||
<TopKeyView
|
||||
topKeys={getTopAPIKeys()}
|
||||
teams={null}
|
||||
showTags={entityType === "tag"}
|
||||
topKeysLimit={topKeysLimit}
|
||||
setTopKeysLimit={setTopKeysLimit}
|
||||
/>
|
||||
</Card>
|
||||
</Col>
|
||||
|
||||
{/* Top Models */}
|
||||
<Col numColSpan={1}>
|
||||
<Card>
|
||||
<Title>{entityType === "agent" ? "Top Agents" : "Top Models"}</Title>
|
||||
<TopModelView
|
||||
topModels={getTopModels()}
|
||||
topModelsLimit={topModelsLimit}
|
||||
setTopModelsLimit={setTopModelsLimit}
|
||||
/>
|
||||
</Card>
|
||||
</Col>
|
||||
|
||||
{/* Top Agents - only for team entity type */}
|
||||
{entityType === "team" && (
|
||||
<Col numColSpan={2}>
|
||||
<Card>
|
||||
<Title>Top Agents Driving Spend</Title>
|
||||
<TopModelView
|
||||
topModels={getTopAgents()}
|
||||
topModelsLimit={topAgentsLimit}
|
||||
setTopModelsLimit={setTopAgentsLimit}
|
||||
/>
|
||||
</Card>
|
||||
</Col>
|
||||
)}
|
||||
|
||||
{/* Spend by Provider */}
|
||||
<Col numColSpan={2}>
|
||||
<Card>
|
||||
<div className="flex flex-col space-y-4">
|
||||
<Title>Provider Usage</Title>
|
||||
<Grid numItems={2}>
|
||||
<Col numColSpan={1}>
|
||||
<DonutChart
|
||||
className="mt-4 h-40"
|
||||
data={getProviderSpend()}
|
||||
index="provider"
|
||||
category="spend"
|
||||
valueFormatter={(value) => `$${formatNumberWithCommas(value, 2)}`}
|
||||
colors={["cyan", "blue", "indigo", "violet", "purple"]}
|
||||
showLabel
|
||||
startAngle={90}
|
||||
endAngle={-270}
|
||||
/>
|
||||
</Col>
|
||||
<Col numColSpan={1}>
|
||||
<Table>
|
||||
<TableHead>
|
||||
<TableRow>
|
||||
<TableHeaderCell>Provider</TableHeaderCell>
|
||||
<TableHeaderCell>Spend</TableHeaderCell>
|
||||
<TableHeaderCell className="text-green-600">Successful</TableHeaderCell>
|
||||
<TableHeaderCell className="text-red-600">Failed</TableHeaderCell>
|
||||
<TableHeaderCell>Tokens</TableHeaderCell>
|
||||
</TableRow>
|
||||
</TableHead>
|
||||
<TableBody>
|
||||
{getProviderSpend().map((provider) => (
|
||||
<TableRow key={provider.provider}>
|
||||
<TableCell>
|
||||
<div className="flex items-center space-x-2">
|
||||
{provider.provider && <Logo provider={provider.provider} className="w-4 h-4" />}
|
||||
<span>{provider.provider}</span>
|
||||
</div>
|
||||
</TableCell>
|
||||
<TableCell>
|
||||
<MoneyCell value={provider.spend} decimals={2} />
|
||||
</TableCell>
|
||||
<TableCell className="text-green-600">
|
||||
{provider.successful_requests.toLocaleString()}
|
||||
</TableCell>
|
||||
<TableCell className="text-red-600">
|
||||
{provider.failed_requests.toLocaleString()}
|
||||
</TableCell>
|
||||
<TableCell>{provider.tokens.toLocaleString()}</TableCell>
|
||||
</TableRow>
|
||||
))}
|
||||
</TableBody>
|
||||
</Table>
|
||||
</Col>
|
||||
</Grid>
|
||||
</div>
|
||||
</Card>
|
||||
</Col>
|
||||
</Grid>
|
||||
</TabPanel>
|
||||
<TabPanel>
|
||||
<ActivityMetrics modelMetrics={modelMetrics} hidePromptCachingMetrics={entityType === "agent"} />
|
||||
</TabPanel>
|
||||
{entityType === "team" ? (
|
||||
<TabPanel>
|
||||
<ActivityMetrics modelMetrics={agentMetrics} />
|
||||
</TabPanel>
|
||||
) : (
|
||||
<></>
|
||||
)}
|
||||
<TabPanel>
|
||||
<ActivityMetrics modelMetrics={keyMetrics} hidePromptCachingMetrics={entityType === "agent"} />
|
||||
</TabPanel>
|
||||
<TabPanel>
|
||||
<EndpointUsage userSpendData={spendData} />
|
||||
</TabPanel>
|
||||
{tabs.map(({ key, content }) => (
|
||||
<TabPanel key={key}>{content}</TabPanel>
|
||||
))}
|
||||
</TabPanels>
|
||||
</TabGroup>
|
||||
</div>
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue