diff --git a/ci_cd/generate_model_prices_schema.py b/ci_cd/generate_model_prices_schema.py index 1b60f986ca4..153fbc0fdc2 100644 --- a/ci_cd/generate_model_prices_schema.py +++ b/ci_cd/generate_model_prices_schema.py @@ -41,6 +41,11 @@ OBJECT_KEYS: dict[str, JsonSchema] = { }, "additionalProperties": False, }, + "guardrail_cost_per_unit": { + "type": "object", + "description": "USD cost per billable guardrail unit, keyed by the provider's usage counter name (e.g. Bedrock's contentPolicyUnits).", + "additionalProperties": NONNEG_NUMBER, + }, "metadata": { "type": "object", "description": "Free-form notes about the entry (e.g. pricing derivation).", diff --git a/litellm/__init__.py b/litellm/__init__.py index ae0fee11aeb..1ecb04b6e54 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -792,6 +792,8 @@ def _populate_provider_model_sets(model_cost_map: Dict) -> None: nlp_cloud_models.add(key) elif value.get("litellm_provider") == "aleph_alpha": aleph_alpha_models.add(key) + elif value.get("litellm_provider") == "bedrock" and value.get("mode") == "guardrail": + pass elif value.get("litellm_provider") == "bedrock" and not is_bedrock_pricing_only_model(key): bedrock_models.add(key) elif value.get("litellm_provider") == "bedrock_converse": diff --git a/litellm/caching/_embedding_router.py b/litellm/caching/_embedding_router.py index 1073b34ef25..8dfcddf158a 100644 --- a/litellm/caching/_embedding_router.py +++ b/litellm/caching/_embedding_router.py @@ -12,8 +12,11 @@ This module is dependency-injected: callers pass the proxy ``llm_router`` and from __future__ import annotations +from collections.abc import Sequence from typing import TYPE_CHECKING, Any, Final +import litellm + if TYPE_CHECKING: from litellm.router import Router @@ -41,3 +44,28 @@ def build_router_embedding_metadata( metadata: Final[dict[str, Any]] = dict(request_metadata or {}) metadata["semantic-cache-embedding"] = True return metadata + + +def resolve_embedding_max_input_tokens( + configured_max_input_tokens: int | None, + embedding_model: str, + router: Router | None, +) -> int | None: + """Explicit cache setting first, else the Router deployment's configured ``max_input_tokens``.""" + if configured_max_input_tokens is not None: + return configured_max_input_tokens + if router is None: + return None + deployment_max_input_tokens, _ = router.get_configured_token_limits(embedding_model) + return deployment_max_input_tokens + + +def truncate_embedding_input(prompt: str, embedding_model: str, max_input_tokens: int | None) -> str: + """Keep only the first ``max_input_tokens`` tokens of ``prompt`` for the embedding call.""" + if max_input_tokens is None: + return prompt + tokens: Final[Sequence[int]] = litellm.encode(model=embedding_model, text=prompt) + if len(tokens) <= max_input_tokens: + return prompt + truncated: Final[str] = litellm.decode(model=embedding_model, tokens=tokens[:max_input_tokens]) + return truncated diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index f0fb91b987f..6b68ae98111 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -97,6 +97,7 @@ class Cache: qdrant_quantization_config: str | None = None, qdrant_semantic_cache_embedding_model: str = "text-embedding-ada-002", qdrant_semantic_cache_vector_size: int | None = None, + semantic_cache_embedding_max_input_tokens: int | None = None, # GCP IAM authentication parameters gcp_service_account: str | None = None, gcp_ssl_ca_certs: str | None = None, @@ -122,6 +123,7 @@ class Cache: qdrant_api_key (str, optional): The api_key for the local or cloud qdrant cluster. qdrant_collection_name (str, optional): The name for your qdrant collection. Required if type is "qdrant-semantic". similarity_threshold (float, optional): The similarity threshold for semantic-caching, Required if type is "redis-semantic" or "qdrant-semantic". + semantic_cache_embedding_max_input_tokens (int, optional): Truncate prompts to this many tokens before embedding them for semantic caching. Defaults to the embedding deployment's configured max_input_tokens. # Disk Cache Args disk_cache_dir (str, optional): The directory for the disk cache. Defaults to None. @@ -192,6 +194,7 @@ class Cache: similarity_threshold=similarity_threshold, embedding_model=redis_semantic_cache_embedding_model, index_name=redis_semantic_cache_index_name, + embedding_max_input_tokens=semantic_cache_embedding_max_input_tokens, **kwargs, ) elif type == LiteLLMCacheType.VALKEY_SEMANTIC: @@ -207,6 +210,7 @@ class Cache: embedding_model=valkey_semantic_cache_embedding_model, index_name=valkey_semantic_cache_index_name, startup_nodes=redis_startup_nodes, + embedding_max_input_tokens=semantic_cache_embedding_max_input_tokens, **kwargs, ) elif type == LiteLLMCacheType.QDRANT_SEMANTIC: @@ -218,6 +222,7 @@ class Cache: quantization_config=qdrant_quantization_config, embedding_model=qdrant_semantic_cache_embedding_model, vector_size=qdrant_semantic_cache_vector_size, + embedding_max_input_tokens=semantic_cache_embedding_max_input_tokens, ) elif type == LiteLLMCacheType.LOCAL: self.cache = InMemoryCache() diff --git a/litellm/caching/qdrant_semantic_cache.py b/litellm/caching/qdrant_semantic_cache.py index 8f8323550f3..8270c655d82 100644 --- a/litellm/caching/qdrant_semantic_cache.py +++ b/litellm/caching/qdrant_semantic_cache.py @@ -12,7 +12,7 @@ import ast import asyncio import json import os -from typing import Any, Final, cast +from typing import TYPE_CHECKING, Any, Final, cast import litellm from litellm._logging import print_verbose @@ -22,12 +22,21 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( ) from litellm.types.utils import EmbeddingResponse -from ._embedding_router import build_router_embedding_metadata, resolve_embedding_router +from ._embedding_router import ( + build_router_embedding_metadata, + resolve_embedding_max_input_tokens, + resolve_embedding_router, + truncate_embedding_input, +) from .base_cache import BaseCache +if TYPE_CHECKING: + from litellm.router import Router + class QdrantSemanticCache(BaseCache): CACHE_KEY_FIELD_NAME = "litellm_cache_key" + embedding_max_input_tokens: int | None = None def __init__( self, @@ -39,6 +48,7 @@ class QdrantSemanticCache(BaseCache): embedding_model="text-embedding-ada-002", host_type=None, vector_size=None, + embedding_max_input_tokens: int | None = None, ): from litellm.llms.custom_httpx.http_handler import ( _get_httpx_client, @@ -57,6 +67,7 @@ class QdrantSemanticCache(BaseCache): raise Exception("similarity_threshold must be provided, passed None") self.similarity_threshold = similarity_threshold self.embedding_model = embedding_model + self.embedding_max_input_tokens = embedding_max_input_tokens self.vector_size = vector_size if vector_size is not None else QDRANT_VECTOR_SIZE headers = {} @@ -188,6 +199,13 @@ class QdrantSemanticCache(BaseCache): cached_key: Final = payload.get(self.CACHE_KEY_FIELD_NAME) return cached_key is not None and str(cached_key) == str(key) + def _embedding_input(self, prompt: str, router: "Router | None") -> str: + return truncate_embedding_input( + prompt, + self.embedding_model, + resolve_embedding_max_input_tokens(self.embedding_max_input_tokens, self.embedding_model, router), + ) + def _get_embedding(self, prompt: str, metadata: dict[str, Any] | None = None) -> EmbeddingResponse: """Embed via the proxy Router when it serves the model, else direct.""" try: @@ -197,16 +215,17 @@ class QdrantSemanticCache(BaseCache): llm_router = None router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list) + embedding_input: Final = self._embedding_input(prompt, router) if router is not None: return router.embedding( model=self.embedding_model, - input=prompt, + input=embedding_input, cache={"no-store": True, "no-cache": True}, metadata=build_router_embedding_metadata(metadata), ) return litellm.embedding( model=self.embedding_model, - input=prompt, + input=embedding_input, cache={"no-store": True, "no-cache": True}, ) @@ -218,17 +237,18 @@ class QdrantSemanticCache(BaseCache): llm_router = None router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list) + embedding_input: Final = self._embedding_input(prompt, router) if router is not None: return await router.aembedding( model=self.embedding_model, - input=prompt, + input=embedding_input, cache={"no-store": True, "no-cache": True}, metadata=build_router_embedding_metadata(metadata), ) return await litellm.aembedding( model=self.embedding_model, - input=prompt, + input=embedding_input, cache={"no-store": True, "no-cache": True}, ) diff --git a/litellm/caching/redis_semantic_cache.py b/litellm/caching/redis_semantic_cache.py index 604d6395ea1..d91260f4d9c 100644 --- a/litellm/caching/redis_semantic_cache.py +++ b/litellm/caching/redis_semantic_cache.py @@ -14,7 +14,7 @@ import asyncio import json import os from collections.abc import Callable, Mapping -from typing import Any, Final, cast +from typing import TYPE_CHECKING, Any, Final, cast import litellm from litellm._logging import print_verbose, verbose_logger @@ -23,9 +23,17 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( ) from litellm.types.utils import EmbeddingResponse -from ._embedding_router import build_router_embedding_metadata, resolve_embedding_router +from ._embedding_router import ( + build_router_embedding_metadata, + resolve_embedding_max_input_tokens, + resolve_embedding_router, + truncate_embedding_input, +) from .base_cache import BaseCache +if TYPE_CHECKING: + from litellm.router import Router + class RedisSemanticCache(BaseCache): """ @@ -38,6 +46,7 @@ class RedisSemanticCache(BaseCache): DEFAULT_REDIS_INDEX_NAME: str = "litellm_semantic_cache_index" CACHE_KEY_FIELD_NAME: str = "litellm_cache_key" + embedding_max_input_tokens: int | None = None def __init__( self, @@ -48,6 +57,7 @@ class RedisSemanticCache(BaseCache): similarity_threshold: float | None = None, embedding_model: str = "text-embedding-ada-002", index_name: str | None = None, + embedding_max_input_tokens: int | None = None, **kwargs: object, ): """ @@ -62,6 +72,8 @@ class RedisSemanticCache(BaseCache): where 1.0 requires exact matches and 0.0 accepts any match embedding_model: Model to use for generating embeddings index_name: Name for the Redis index + embedding_max_input_tokens: Truncate prompts to this many tokens before + embedding; defaults to the Router deployment's configured max_input_tokens ttl: Default time-to-live for cache entries in seconds **kwargs: Additional arguments passed to the Redis client @@ -86,6 +98,7 @@ class RedisSemanticCache(BaseCache): # While similarity: 1 = most similar, 0 = least similar self.distance_threshold = 1 - similarity_threshold self.embedding_model = embedding_model + self.embedding_max_input_tokens = embedding_max_input_tokens # Set up Redis connection if redis_url is None: @@ -307,6 +320,13 @@ class RedisSemanticCache(BaseCache): return dict_method() return value + def _embedding_input(self, prompt: str, router: "Router | None") -> str: + return truncate_embedding_input( + prompt, + self.embedding_model, + resolve_embedding_max_input_tokens(self.embedding_max_input_tokens, self.embedding_model, router), + ) + def _get_embedding(self, prompt: str, metadata: dict[str, Any] | None = None) -> list[float]: """ Routes through the proxy Router when the embedding model is a Router @@ -320,12 +340,13 @@ class RedisSemanticCache(BaseCache): llm_router = None router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list) + embedding_input: Final = self._embedding_input(prompt, router) if router is not None: embedding_response = cast( EmbeddingResponse, router.embedding( model=self.embedding_model, - input=prompt, + input=embedding_input, cache={"no-store": True, "no-cache": True}, metadata=build_router_embedding_metadata(metadata), ), @@ -335,7 +356,7 @@ class RedisSemanticCache(BaseCache): EmbeddingResponse, litellm.embedding( model=self.embedding_model, - input=prompt, + input=embedding_input, cache={"no-store": True, "no-cache": True}, ), ) @@ -490,18 +511,19 @@ class RedisSemanticCache(BaseCache): llm_router = None router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list) + embedding_input: Final = self._embedding_input(prompt, router) try: if router is not None: embedding_response = await router.aembedding( model=self.embedding_model, - input=prompt, + input=embedding_input, cache={"no-store": True, "no-cache": True}, metadata=build_router_embedding_metadata(metadata), ) else: embedding_response = await litellm.aembedding( model=self.embedding_model, - input=prompt, + input=embedding_input, cache={"no-store": True, "no-cache": True}, ) return embedding_response["data"][0]["embedding"] diff --git a/litellm/caching/valkey_semantic_cache.py b/litellm/caching/valkey_semantic_cache.py index cf6fab780fc..737d212a89d 100644 --- a/litellm/caching/valkey_semantic_cache.py +++ b/litellm/caching/valkey_semantic_cache.py @@ -61,6 +61,7 @@ class ValkeySemanticCache(RedisSemanticCache): startup_nodes: list | None = None, sync_client: Redis | None = None, async_client: AsyncRedis | None = None, + embedding_max_input_tokens: int | None = None, **kwargs: Any, ): if similarity_threshold is None: @@ -78,6 +79,7 @@ class ValkeySemanticCache(RedisSemanticCache): self.similarity_threshold = similarity_threshold self.embedding_model = embedding_model + self.embedding_max_input_tokens = embedding_max_input_tokens self.index_name = index_name or self.DEFAULT_VALKEY_INDEX_NAME self.key_prefix = f"{self.index_name}:" self._index_dim: int | None = None diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index edb4d56a5b7..10e681b816e 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -64,6 +64,10 @@ from litellm.integrations.mlflow import MlflowLogger from litellm.integrations.sqs import SQSLogger from litellm.litellm_core_utils.core_helpers import reconstruct_model_name from litellm.litellm_core_utils.get_litellm_params import get_litellm_params +from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import ( + cost_breakdown_with_guardrail, + guardrail_information_cost, +) from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import ( StandardBuiltInToolCostTracking, ) @@ -5650,12 +5654,14 @@ def get_standard_logging_object_payload( base_model = metadata.get("deployment") custom_pricing: Final = use_custom_pricing_for_model(litellm_params=litellm_params) raw_response_cost: Final = kwargs.get("response_cost") - response_cost: Final[float] = raw_response_cost or 0.0 + llm_response_cost: Final[float] = raw_response_cost or 0.0 + guardrail_cost: Final = guardrail_information_cost(metadata.get("standard_logging_guardrail_information")) + response_cost: Final[float] = llm_response_cost + guardrail_cost # clean up litellm hidden params clean_hidden_params: Final = StandardLoggingPayloadSetup.get_hidden_params(hidden_params) if clean_hidden_params["response_cost"] is None and raw_response_cost is not None: - clean_hidden_params["response_cost"] = response_cost + clean_hidden_params["response_cost"] = llm_response_cost model_cost_information: Final = StandardLoggingPayloadSetup.get_model_cost_information( base_model=base_model, @@ -5735,7 +5741,7 @@ def get_standard_logging_object_payload( metadata=clean_metadata, cache_key=clean_hidden_params["cache_key"], response_cost=response_cost, - cost_breakdown=logging_obj.cost_breakdown, + cost_breakdown=cost_breakdown_with_guardrail(logging_obj.cost_breakdown, guardrail_cost), total_tokens=usage_dict.get("total_tokens", 0), prompt_tokens=usage_dict.get("prompt_tokens", 0), completion_tokens=usage_dict.get("completion_tokens", 0), diff --git a/litellm/litellm_core_utils/llm_cost_calc/guardrail_cost.py b/litellm/litellm_core_utils/llm_cost_calc/guardrail_cost.py new file mode 100644 index 00000000000..4645a8c3074 --- /dev/null +++ b/litellm/litellm_core_utils/llm_cost_calc/guardrail_cost.py @@ -0,0 +1,78 @@ +import math +from collections.abc import Mapping +from typing import Final + +from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError + +import litellm +from litellm._logging import verbose_logger +from litellm.types.utils import CostBreakdown + +BEDROCK_GUARDRAIL_PRICING_KEY: Final = "bedrock/guardrails" + + +class GuardrailPricing(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True) + + guardrail_cost_per_unit: Mapping[str, float] + + +class GuardrailCostEntry(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True) + + guardrail_cost: float | None = None + + +GuardrailInformationShape = tuple[GuardrailCostEntry, ...] | GuardrailCostEntry | None + +_GUARDRAIL_INFORMATION_ADAPTER: Final[TypeAdapter[GuardrailInformationShape]] = TypeAdapter(GuardrailInformationShape) + + +def _bedrock_guardrail_pricing(aws_region_name: str | None) -> GuardrailPricing | None: + regional_key: Final = f"bedrock/{aws_region_name}/guardrails" if aws_region_name else None + for key in (regional_key, BEDROCK_GUARDRAIL_PRICING_KEY): + if key is None or key not in litellm.model_cost: + continue + try: + return GuardrailPricing.model_validate(litellm.model_cost[key]) + except ValidationError as e: + verbose_logger.warning("Ignoring malformed guardrail pricing entry %s: %s", key, e) + return None + + +def bedrock_guardrail_cost(usage_units: Mapping[str, int], aws_region_name: str | None) -> float: + pricing: Final = _bedrock_guardrail_pricing(aws_region_name) + if pricing is None: + return 0.0 + return sum(units * pricing.guardrail_cost_per_unit.get(counter, 0.0) for counter, units in usage_units.items()) + + +def _billable_entry_cost(entry: GuardrailCostEntry) -> float: + cost: Final = entry.guardrail_cost + if cost is None or not math.isfinite(cost) or cost <= 0.0: + return 0.0 + return cost + + +def guardrail_information_cost(guardrail_information: object) -> float: + try: + parsed: Final = _GUARDRAIL_INFORMATION_ADAPTER.validate_python(guardrail_information) + except ValidationError: + return 0.0 + if parsed is None: + return 0.0 + if isinstance(parsed, GuardrailCostEntry): + return _billable_entry_cost(parsed) + return sum(_billable_entry_cost(entry) for entry in parsed) + + +def cost_breakdown_with_guardrail(cost_breakdown: CostBreakdown | None, guardrail_cost: float) -> CostBreakdown | None: + if guardrail_cost <= 0.0: + return cost_breakdown + existing: Final[CostBreakdown] = cost_breakdown if cost_breakdown is not None else CostBreakdown() + merged: Final[CostBreakdown] = { + **existing, + "guardrail_cost": guardrail_cost, + "total_cost": existing.get("total_cost", 0.0) + guardrail_cost, + } + return merged diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index d5b47466477..eccc783dd8b 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -896,6 +896,7 @@ class BaseLLMHTTPHandler: ) if provider_config is None: raise ValueError(f"Provider {custom_llm_provider} does not support embedding") + embedding_extra_body: Final[Mapping[str, object] | None] = optional_params.pop("extra_body", None) # get config from model, custom llm provider headers = provider_config.validate_environment( api_key=api_key, @@ -920,6 +921,8 @@ class BaseLLMHTTPHandler: optional_params=optional_params, headers=headers, ) + if embedding_extra_body: + data.update(embedding_extra_body) # Some providers (e.g. OCI) require request signing after the body is built. # The default BaseConfig.sign_request returns (headers, None) — a no-op for diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 78b53cefc53..409022016b0 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -10052,6 +10052,21 @@ "output_cost_per_second": 0.0066027, "supports_tool_choice": true }, + "bedrock/guardrails": { + "guardrail_cost_per_unit": { + "automatedReasoningPolicyUnits": 0.00017, + "contentPolicyImageUnits": 0.00075, + "contentPolicyUnits": 0.00015, + "contextualGroundingPolicyUnits": 0.0001, + "sensitiveInformationPolicyFreeUnits": 0.0, + "sensitiveInformationPolicyUnits": 0.0001, + "topicPolicyUnits": 0.00015, + "wordPolicyUnits": 0.0 + }, + "litellm_provider": "bedrock", + "mode": "guardrail", + "source": "https://aws.amazon.com/bedrock/pricing/" + }, "bedrock/ap-northeast-1/1-month-commitment/anthropic.claude-instant-v1": { "input_cost_per_second": 0.01475, "litellm_provider": "bedrock", diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 7c93155964d..a0b69ecb0bf 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -32,11 +32,13 @@ from litellm.constants import ( UNSAFE_PROXY_RESPONSE_HEADERS, ) from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket from litellm.litellm_core_utils.dd_tracing import NullTracer, tracer from litellm.litellm_core_utils.get_supported_openai_params import ( get_supported_openai_params, ) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import guardrail_information_cost from litellm.litellm_core_utils.llm_response_utils.get_headers import ( get_response_headers, ) @@ -2462,11 +2464,21 @@ class ProxyBaseLLMRequestProcessing: additional_headers = hidden_params.get("additional_headers", {}) or {} recover_response_cost: Final = not response_cost and hidden_params.get("response_cost") is None - response_cost_for_headers: Final = ( + llm_cost_for_headers: Final = ( self._response_cost_from_logging_obj(response=response, logging_obj=logging_obj) or "" if recover_response_cost else response_cost ) + _, request_metadata_bucket = get_or_create_metadata_bucket(self.data) + guardrail_cost_for_headers: Final = guardrail_information_cost( + request_metadata_bucket.get("standard_logging_guardrail_information") + ) + response_cost_for_headers: Final = ( + (llm_cost_for_headers if isinstance(llm_cost_for_headers, (int, float)) else 0.0) + + guardrail_cost_for_headers + if guardrail_cost_for_headers > 0 + else llm_cost_for_headers + ) fastapi_response.headers.update( ProxyBaseLLMRequestProcessing.get_custom_headers( diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 93cbb989e23..c70a2ee8a74 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -31,6 +31,7 @@ from litellm.constants import BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS from litellm.exceptions import ModifyResponseException from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.litellm_core_utils.core_helpers import redact_nested_match_and_regex_keys +from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import bedrock_guardrail_cost from litellm.llms.anthropic.chat.guardrail_translation.handler import AnthropicMessagesHandler from litellm.llms.base_llm.guardrail_translation.utils import ( effective_scan_only_tool_results_for_guardrail, @@ -872,6 +873,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): credentials, aws_region_name = self._load_credentials() allow_chunking: Final = not self._content_uses_contextual_grounding(content) + completed_chunk_usages: Final[list[BedrockGuardrailUsage]] = [] # mutable-ok: billed-chunk usage accumulator try: responses: Final = await self._apply_guardrail_content_with_chunking( content=content, @@ -883,6 +885,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): event_type=event_type, start_time=start_time, allow_chunking=allow_chunking, + completed_chunk_usages=completed_chunk_usages, ) except HTTPException as exc: if not isinstance(exc.detail, dict): @@ -891,6 +894,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): request_data=request_data, event_type=event_type, start_time=start_time, + aws_region_name=aws_region_name, + completed_chunk_usages=completed_chunk_usages, ) raise merged_response: Final = self._merge_bedrock_guardrail_responses(responses) @@ -899,6 +904,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): request_data=request_data, event_type=event_type, start_time=start_time, + aws_region_name=aws_region_name, ) return merged_response @@ -913,6 +919,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): event_type: GuardrailEventHooks, start_time: "datetime", allow_chunking: bool, + completed_chunk_usages: list[BedrockGuardrailUsage], # mutable-ok: billed-chunk usage accumulator ) -> tuple[BedrockContentChunkResult, ...]: """Post `content` to ApplyGuardrail, chunking only if AWS rejects it as too large. @@ -959,6 +966,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): request_data=request_data, event_type=event_type, start_time=start_time, + completed_chunk_usages=completed_chunk_usages, ) return ( BedrockContentChunkResult( @@ -989,6 +997,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): event_type=event_type, start_time=start_time, allow_chunking=allow_chunking, + completed_chunk_usages=completed_chunk_usages, ) for batch in batches ] @@ -1015,6 +1024,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): event_type=event_type, start_time=start_time, allow_chunking=allow_chunking, + completed_chunk_usages=completed_chunk_usages, ) second_results: Final = await self._apply_guardrail_content_with_chunking( content=second_half, @@ -1026,6 +1036,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): event_type=event_type, start_time=start_time, allow_chunking=allow_chunking, + completed_chunk_usages=completed_chunk_usages, ) combined_results: Final = tuple(first_results) + tuple(second_results) if is_single_item_text_split: @@ -1045,6 +1056,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): request_data: dict | None, # mutable-ok: proxy request body dict, mutated by the logging helper event_type: GuardrailEventHooks, start_time: "datetime", + completed_chunk_usages: list[BedrockGuardrailUsage], # mutable-ok: passed through to the single-call layer ) -> BedrockGuardrailResponse: """Post one ApplyGuardrail call for `content`, retrying with exponential backoff on AWS ThrottlingException (HTTP 429). @@ -1072,6 +1084,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): request_data=request_data, event_type=event_type, start_time=start_time, + completed_chunk_usages=completed_chunk_usages, ) except HTTPException as exc: if ( @@ -1093,6 +1106,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): request_data: dict | None, # mutable-ok: proxy request body dict, mutated by the logging helper event_type: GuardrailEventHooks, start_time: "datetime", + completed_chunk_usages: list[BedrockGuardrailUsage], # mutable-ok: billed-chunk usage accumulator ) -> BedrockGuardrailResponse: """Make exactly one signed ApplyGuardrail HTTP call for `content` and parse the result. Raises HTTPException on a guardrail block or any @@ -1108,7 +1122,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): A block is logged here rather than by the caller: it ends the whole chunking flow immediately, with no further chunks attempted, so there is no later - merged response for the caller to log instead. + merged response for the caller to log instead. The logged usage still spans + the whole logical request: chunks that passed before the block appended what + AWS billed them to ``completed_chunk_usages``, and the attempt log sums those + with the blocking call's own usage. """ bedrock_request_data: Final = { # mutable-ok: outbound JSON request body **base_request_data, @@ -1151,10 +1168,17 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): request_data=request_data, event_type=event_type, start_time=start_time, + aws_region_name=aws_region_name, + completed_chunk_usages=completed_chunk_usages, ) raise self._get_http_exception_for_blocked_guardrail( bedrock_guardrail_response, request_data=request_data ) + response_usage: Final = bedrock_guardrail_response.get("usage") + if isinstance(response_usage, dict): + completed_chunk_usages.append( + response_usage + ) # rebind-ok: accumulator threaded from make_bedrock_api_request, recording this billed call return bedrock_guardrail_response status_code, detail_message = self._parse_bedrock_guardrail_error_response(httpx_response) @@ -1172,14 +1196,31 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): request_data: dict | None, # mutable-ok: proxy request body dict, mutated by the logging helper event_type: GuardrailEventHooks, start_time: "datetime", + aws_region_name: str | None, + completed_chunk_usages: Sequence[BedrockGuardrailUsage], ) -> None: - """Log a single ApplyGuardrail HTTP attempt as-is (its own status, - derived from its own response). Used only for the blocked-content - case, which ends the whole chunking flow immediately.""" - tracing_detail: Final = self._build_tracing_detail(BedrockGuardrailResponse(**json_response)) + """Log the blocking ApplyGuardrail attempt, which ends the whole chunking + flow immediately. Its status derives from its own response, but its usage + (and so its cost) spans every billed call of the logical request: the + chunks that passed before the block plus the blocking call itself.""" + blocking_usage: Final = json_response.get("usage") + billed_usages: Final[tuple[BedrockGuardrailUsage, ...]] = tuple(completed_chunk_usages) + ( + (blocking_usage,) if isinstance(blocking_usage, dict) else () + ) + logged_json_response: Final = ( + { # mutable-ok: raw AWS JSON payload carrying the total billed usage + **json_response, + "usage": self._sum_usage_counters(billed_usages), + } + if completed_chunk_usages + else json_response + ) + tracing_detail: Final = self._build_tracing_detail( + BedrockGuardrailResponse(**logged_json_response), aws_region_name=aws_region_name + ) self.add_standard_logging_guardrail_information_to_request_data( guardrail_provider=self.guardrail_provider, - guardrail_json_response=json_response, + guardrail_json_response=logged_json_response, request_data=request_data or {}, # mutable-ok: logging helper requires a dict guardrail_status=self._get_bedrock_guardrail_response_status(response=httpx_response), start_time=start_time.timestamp(), @@ -1195,6 +1236,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): request_data: dict | None, # mutable-ok: proxy request body dict, mutated by the logging helper event_type: GuardrailEventHooks, start_time: "datetime", + aws_region_name: str | None, ) -> None: """Log one logical ApplyGuardrail call -- possibly several chunk calls under the hood -- using its final merged response, so a chunked @@ -1205,7 +1247,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ``Output.__type`` with an exception marker. That marker survives the merge, so the status is derived from the merged response rather than assumed to be a success, which is what the pre-chunking code reported for that shape.""" - tracing_detail: Final = self._build_tracing_detail(merged_response) + tracing_detail: Final = self._build_tracing_detail(merged_response, aws_region_name=aws_region_name) self.add_standard_logging_guardrail_information_to_request_data( guardrail_provider=self.guardrail_provider, guardrail_json_response=dict(merged_response), # mutable-ok: logging helper requires a dict @@ -1228,20 +1270,36 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): request_data: dict | None, # mutable-ok: proxy request body dict, mutated by the logging helper event_type: GuardrailEventHooks, start_time: "datetime", + aws_region_name: str | None, + completed_chunk_usages: Sequence[BedrockGuardrailUsage], ) -> None: """Log one logical ApplyGuardrail call that failed end-to-end (an unrecoverable too-large error, a non-size validation error, or exhausted throttle retries) as a single failure, rather than logging - every failed attempt chunking made along the way.""" + every failed attempt chunking made along the way. Chunk calls AWS + billed before the failure still carry their usage and cost.""" + billed_usage: Final = self._sum_usage_counters(completed_chunk_usages) if completed_chunk_usages else None + error_payload: Final = {"error": str(detail)} # mutable-ok: logging helper requires a dict + json_response: Final = ( + {**error_payload, "usage": billed_usage} # mutable-ok: logging helper requires a dict + if billed_usage is not None + else error_payload + ) + tracing_detail: Final = ( + self._build_tracing_detail(BedrockGuardrailResponse(usage=billed_usage), aws_region_name=aws_region_name) + if billed_usage is not None + else None + ) self.add_standard_logging_guardrail_information_to_request_data( guardrail_provider=self.guardrail_provider, - guardrail_json_response={"error": str(detail)}, # mutable-ok: logging helper requires a dict + guardrail_json_response=json_response, request_data=request_data or {}, # mutable-ok: logging helper requires a dict guardrail_status="guardrail_failed_to_respond", start_time=start_time.timestamp(), end_time=datetime.now(timezone.utc).timestamp(), duration=(datetime.now(timezone.utc) - start_time).total_seconds(), event_type=event_type, + tracing_detail=tracing_detail or None, ) @staticmethod @@ -1504,15 +1562,20 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): Keys are taken from the responses rather than from a fixed list, so a counter this code does not know about (AWS has added several) is still summed and reported instead of being silently dropped to zero.""" - chunk_usages: Final = tuple( - chunk_result.response.get("usage") or {} # mutable-ok: read-only empty fallback - for chunk_result in chunk_results + return BedrockGuardrail._sum_usage_counters( + tuple( + chunk_result.response.get("usage") or {} # mutable-ok: read-only empty fallback + for chunk_result in chunk_results + ) ) + + @staticmethod + def _sum_usage_counters(usages: Sequence[BedrockGuardrailUsage]) -> BedrockGuardrailUsage: return cast( # cast-ok: TypedDict assembled from a comprehension BedrockGuardrailUsage, { # mutable-ok: builds the TypedDict payload - key: sum(usage.get(key) or 0 for usage in chunk_usages) - for key in dict.fromkeys(key for usage in chunk_usages for key in usage) + key: sum(usage.get(key) or 0 for usage in usages) + for key in dict.fromkeys(key for usage in usages for key in usage) }, ) @@ -2036,7 +2099,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): return (status_code, err) return (status_code, message) - def _build_tracing_detail(self, response: BedrockGuardrailResponse) -> GuardrailTracingDetail: + def _build_tracing_detail( + self, response: BedrockGuardrailResponse, aws_region_name: str | None + ) -> GuardrailTracingDetail: """ Build the tracing detail from the raw Bedrock response, before redaction, so downstream loggers (OTEL, Langfuse, ...) get the @@ -2060,6 +2125,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): } if usage_units: tracing_detail["guardrail_usage"] = usage_units + tracing_detail["guardrail_cost"] = bedrock_guardrail_cost( + usage_units=usage_units, aws_region_name=aws_region_name + ) return tracing_detail def _extract_violation_category_names(self, response: BedrockGuardrailResponse) -> list[str]: diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 5dc92d82bda..99d0c94d11b 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -11,6 +11,7 @@ from litellm.litellm_core_utils.core_helpers import ( get_litellm_metadata_from_kwargs, ) from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup +from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import guardrail_information_cost from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_checks import ( get_key_object, @@ -184,9 +185,14 @@ class _ProxyDBLogger(CustomLogger): # recovered cost onto request_data (the usage rides along in # ``combined_usage_object`` for the token columns), so attribute the # real partial spend to this failure row instead of zero. - recovered_response_cost = 0.0 - if isinstance(request_data.get("combined_usage_object"), litellm.Usage): - recovered_response_cost = max(float(request_data.get("response_cost") or 0.0), 0.0) + recovered_stream_cost: Final = ( + max(float(request_data.get("response_cost") or 0.0), 0.0) + if isinstance(request_data.get("combined_usage_object"), litellm.Usage) + else 0.0 + ) + recovered_response_cost: Final = recovered_stream_cost + guardrail_information_cost( + existing_metadata.get("standard_logging_guardrail_information") + ) await proxy_logging_obj.db_spend_update_writer.update_database( token=user_api_key_dict.api_key, diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index ef6e590ba26..2ec5c34958c 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -296,7 +296,10 @@ _ALLOW_CLIENT_MESSAGE_REDACTION_OPT_OUT_METADATA_KEY: Final = "allow_client_mess _CLIENT_PRICING_CONTROL_FIELDS: Final = frozenset(CustomPricingLiteLLMParams.model_fields.keys()) # ``model_info`` carries the same pricing fields when read by # ``use_custom_pricing_for_model``; strip from metadata for the same reason. -_CLIENT_PRICING_METADATA_FIELDS: Final = frozenset({"model_info"}) +# ``standard_logging_guardrail_information`` is proxy-written telemetry summed +# into response_cost and spend; a client seeding it forges (even negative) +# guardrail cost. +_CLIENT_PRICING_METADATA_FIELDS: Final = frozenset({"model_info", "standard_logging_guardrail_information"}) _ALLOW_CLIENT_PRICING_OVERRIDE_METADATA_KEY: Final = "allow_client_pricing_override" # Request fields whose value, when URL-valued, becomes the outbound destination diff --git a/litellm/types/utils.py b/litellm/types/utils.py index c0e8f482429..d44d4cca6c4 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3027,6 +3027,11 @@ class StandardLoggingGuardrailInformation(TypedDict, total=False): provider's counter name (e.g. Bedrock's ``contentPolicyUnits``). Kept as a sibling of guardrail_response so spend-log prompt redaction never drops it.""" + guardrail_cost: ReadOnly[float | None] + """USD cost of this guardrail invocation, priced from ``guardrail_usage`` by the + provider hook. Summed into the request's ``response_cost`` so it counts against + spend and budgets like token cost.""" + class EvalVerdict(TypedDict, total=False): criterion_name: str @@ -3071,6 +3076,7 @@ class GuardrailTracingDetail(TypedDict, total=False): violation_categories: list[str] | None guardrail_action: str | None guardrail_usage: ReadOnly[Mapping[str, int] | None] + guardrail_cost: ReadOnly[float | None] StandardLoggingPayloadStatus = Literal["success", "failure"] @@ -3110,8 +3116,9 @@ class CostBreakdown(TypedDict, total=False): cache_creation_cost: float # Cost of cache-write tokens (premium rate) output_cost: float # Cost of output/completion tokens (includes reasoning if applicable) reasoning_cost: float # Cost of reasoning tokens (subset of output_cost) - total_cost: float # Total cost (input + output + tool usage) + total_cost: ReadOnly[float] # Total cost (input + output + tool usage + guardrail) tool_usage_cost: float # Cost of usage of built-in tools + guardrail_cost: ReadOnly[float] # Cost of guardrail invocations billed by the guardrail provider additional_costs: dict[str, float] # Free-form additional costs (e.g., {"azure_model_router_flat_cost": 0.00014}) original_cost: float # Cost before discount (optional) discount_percent: float # Discount percentage applied (e.g., 0.05 = 5%) (optional) diff --git a/litellm/utils.py b/litellm/utils.py index c4ad825e987..1c880ee9521 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -2110,7 +2110,7 @@ def encode(model="", text="", custom_tokenizer: dict | None = None): def decode( model="", - tokens: list[int] = [], + tokens: Sequence[int] = (), custom_tokenizer: dict | None = None, skip_special_tokens: bool = True, ): @@ -2132,7 +2132,7 @@ def decode( return dec -def _strip_huggingface_special_token_ids(tokenizer: Tokenizer, tokens: list[int]) -> list[int]: +def _strip_huggingface_special_token_ids(tokenizer: Tokenizer, tokens: Sequence[int]) -> Sequence[int]: try: added_tokens_decoder: Final = tokenizer.get_added_tokens_decoder() except Exception: diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 78b53cefc53..409022016b0 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -10052,6 +10052,21 @@ "output_cost_per_second": 0.0066027, "supports_tool_choice": true }, + "bedrock/guardrails": { + "guardrail_cost_per_unit": { + "automatedReasoningPolicyUnits": 0.00017, + "contentPolicyImageUnits": 0.00075, + "contentPolicyUnits": 0.00015, + "contextualGroundingPolicyUnits": 0.0001, + "sensitiveInformationPolicyFreeUnits": 0.0, + "sensitiveInformationPolicyUnits": 0.0001, + "topicPolicyUnits": 0.00015, + "wordPolicyUnits": 0.0 + }, + "litellm_provider": "bedrock", + "mode": "guardrail", + "source": "https://aws.amazon.com/bedrock/pricing/" + }, "bedrock/ap-northeast-1/1-month-commitment/anthropic.claude-instant-v1": { "input_cost_per_second": 0.01475, "litellm_provider": "bedrock", diff --git a/model_prices_and_context_window.schema.json b/model_prices_and_context_window.schema.json index 4c54822736c..cd02fde595f 100644 --- a/model_prices_and_context_window.schema.json +++ b/model_prices_and_context_window.schema.json @@ -186,6 +186,14 @@ "gemini_native_audio": { "type": "boolean" }, + "guardrail_cost_per_unit": { + "type": "object", + "description": "USD cost per billable guardrail unit, keyed by the provider's usage counter name (e.g. Bedrock's contentPolicyUnits).", + "additionalProperties": { + "type": "number", + "minimum": 0 + } + }, "input_cost_per_audio_per_second": { "type": "number", "minimum": 0 @@ -361,6 +369,7 @@ "chat", "completion", "embedding", + "guardrail", "image_edit", "image_generation", "moderation", diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 4a6ff6af902..06bee4a054c 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -33,7 +33,7 @@ "limit": 2 }, "B006": { - "limit": 178 + "limit": 177 }, "B008": { "limit": 503 diff --git a/tests/test_litellm/caching/test_caching.py b/tests/test_litellm/caching/test_caching.py index b65e8773c85..955b0e531bc 100644 --- a/tests/test_litellm/caching/test_caching.py +++ b/tests/test_litellm/caching/test_caching.py @@ -89,6 +89,20 @@ def _semantic_cache(): ) +@pytest.mark.parametrize( + "cache_type", + [LiteLLMCacheType.REDIS_SEMANTIC, LiteLLMCacheType.VALKEY_SEMANTIC], +) +def test_semantic_cache_embedding_max_input_tokens_reaches_backend(cache_type): + cache = Cache( + type=cache_type, + redis_url="redis://localhost:6379", + similarity_threshold=0.8, + semantic_cache_embedding_max_input_tokens=2048, + ) + assert cache.cache.embedding_max_input_tokens == 2048 + + def test_semantic_cache_key_excludes_prompt_so_paraphrases_share_a_bucket(): cache = _semantic_cache() tenant = {"user_api_key": "hash-abc"} diff --git a/tests/test_litellm/caching/test_embedding_router.py b/tests/test_litellm/caching/test_embedding_router.py index 550095a112a..9ebe669d32d 100644 --- a/tests/test_litellm/caching/test_embedding_router.py +++ b/tests/test_litellm/caching/test_embedding_router.py @@ -4,9 +4,12 @@ from unittest.mock import MagicMock sys.path.insert(0, os.path.abspath("../../..")) +import litellm from litellm.caching._embedding_router import ( build_router_embedding_metadata, + resolve_embedding_max_input_tokens, resolve_embedding_router, + truncate_embedding_input, ) @@ -65,3 +68,40 @@ def test_build_metadata_handles_none_and_does_not_mutate_input(): assert md == {"user_api_key": "sk-x", "semantic-cache-embedding": True} assert original == {"user_api_key": "sk-x"} assert build_router_embedding_metadata(None) == {"semantic-cache-embedding": True} + + +def test_resolve_max_input_tokens_prefers_configured_over_deployment(): + router = MagicMock() + router.get_configured_token_limits.return_value = (8191, None) + assert resolve_embedding_max_input_tokens(512, "sem-embed", router) == 512 + router.get_configured_token_limits.assert_not_called() + + +def test_resolve_max_input_tokens_falls_back_to_deployment_limit(): + router = MagicMock() + router.get_configured_token_limits.return_value = (8191, 4096) + assert resolve_embedding_max_input_tokens(None, "sem-embed", router) == 8191 + router.get_configured_token_limits.assert_called_once_with("sem-embed") + + +def test_resolve_max_input_tokens_is_none_without_router_or_deployment_limit(): + router = MagicMock() + router.get_configured_token_limits.return_value = (None, None) + assert resolve_embedding_max_input_tokens(None, "sem-embed", router) is None + assert resolve_embedding_max_input_tokens(None, "sem-embed", None) is None + + +def test_truncate_embedding_input_keeps_prompt_within_limit(): + prompt = "The quick brown fox jumps over the lazy dog" + assert truncate_embedding_input(prompt, "sem-embed", None) == prompt + assert truncate_embedding_input(prompt, "sem-embed", 100) == prompt + token_count = len(litellm.encode(model="sem-embed", text=prompt)) + assert truncate_embedding_input(prompt, "sem-embed", token_count) == prompt + + +def test_truncate_embedding_input_cuts_prompt_to_token_limit(): + prompt = " ".join(f"word{i}" for i in range(400)) + truncated = truncate_embedding_input(prompt, "sem-embed", 50) + assert prompt.startswith(truncated) + assert len(truncated) < len(prompt) + assert len(litellm.encode(model="sem-embed", text=truncated)) == 50 diff --git a/tests/test_litellm/caching/test_qdrant_semantic_cache.py b/tests/test_litellm/caching/test_qdrant_semantic_cache.py index 67d4e2d9892..852bed4a9df 100644 --- a/tests/test_litellm/caching/test_qdrant_semantic_cache.py +++ b/tests/test_litellm/caching/test_qdrant_semantic_cache.py @@ -43,6 +43,7 @@ def test_qdrant_semantic_cache_initialization(monkeypatch): qdrant_api_base="http://test.qdrant.local", qdrant_api_key="test_key", similarity_threshold=0.8, + embedding_max_input_tokens=512, ) # Verify the cache was initialized with correct parameters @@ -50,6 +51,7 @@ def test_qdrant_semantic_cache_initialization(monkeypatch): assert qdrant_cache.qdrant_api_base == "http://test.qdrant.local" assert qdrant_cache.qdrant_api_key == "test_key" assert qdrant_cache.similarity_threshold == 0.8 + assert qdrant_cache.embedding_max_input_tokens == 512 mock_sync_client_instance.put.assert_called_once_with( url="http://test.qdrant.local/collections/test_collection/index", headers={ @@ -832,6 +834,7 @@ def test_qdrant_sync_get_cache_routes_through_router(monkeypatch): cache.sync_client.post.return_value = search_response router = MagicMock() + router.get_configured_token_limits.return_value = (None, None) router.embedding = MagicMock( return_value={"data": [{"embedding": [0.3, 0.3, 0.3]}]} ) @@ -892,6 +895,7 @@ async def test_qdrant_async_embedding_forwards_full_metadata(monkeypatch): cache.embedding_model = "sem-embed" router = MagicMock() + router.get_configured_token_limits.return_value = (None, None) router.aembedding = AsyncMock(return_value={"data": [{"embedding": [0.1, 0.2]}]}) monkeypatch.setitem( sys.modules, @@ -908,3 +912,57 @@ async def test_qdrant_async_embedding_forwards_full_metadata(monkeypatch): assert md["user_api_key"] == "sk-x" assert md["user_api_key_team_id"] == "team-1" assert md["semantic-cache-embedding"] is True + + +LONG_PROMPT = " ".join(f"token{i}" for i in range(300)) + + +def _token_count(model, text): + import litellm + + return len(litellm.encode(model=model, text=text)) + + +def test_qdrant_get_embedding_truncates_to_deployment_max_input_tokens(monkeypatch): + from litellm.caching.qdrant_semantic_cache import QdrantSemanticCache + + cache = QdrantSemanticCache.__new__(QdrantSemanticCache) + cache.embedding_model = "sem-embed" + + router = MagicMock() + router.get_configured_token_limits.return_value = (5, None) + router.embedding = MagicMock(return_value={"data": [{"embedding": [0.5, 0.6]}]}) + monkeypatch.setitem( + sys.modules, + "litellm.proxy.proxy_server", + _router_proxy_module(router, "sem-embed"), + ) + + cache._get_embedding(LONG_PROMPT) + + sent_input = router.embedding.call_args.kwargs["input"] + assert LONG_PROMPT.startswith(sent_input) + assert _token_count("sem-embed", sent_input) == 5 + + +@pytest.mark.asyncio +async def test_qdrant_async_embedding_explicit_limit_beats_deployment_limit(monkeypatch): + from litellm.caching.qdrant_semantic_cache import QdrantSemanticCache + + cache = QdrantSemanticCache.__new__(QdrantSemanticCache) + cache.embedding_model = "sem-embed" + cache.embedding_max_input_tokens = 3 + + router = MagicMock() + router.get_configured_token_limits.return_value = (8191, None) + router.aembedding = AsyncMock(return_value={"data": [{"embedding": [0.1, 0.2]}]}) + monkeypatch.setitem( + sys.modules, + "litellm.proxy.proxy_server", + _router_proxy_module(router, "sem-embed"), + ) + + await cache._get_async_embedding(LONG_PROMPT) + + sent_input = router.aembedding.call_args.kwargs["input"] + assert _token_count("sem-embed", sent_input) == 3 diff --git a/tests/test_litellm/caching/test_redis_semantic_cache.py b/tests/test_litellm/caching/test_redis_semantic_cache.py index 1d3129d6467..9fd333cf87c 100644 --- a/tests/test_litellm/caching/test_redis_semantic_cache.py +++ b/tests/test_litellm/caching/test_redis_semantic_cache.py @@ -901,6 +901,7 @@ def test_redis_get_embedding_routes_through_router(monkeypatch): cache.embedding_model = "sem-embed" router = MagicMock() + router.get_configured_token_limits.return_value = (None, None) router.embedding = MagicMock(return_value={"data": [{"embedding": [0.5, 0.6]}]}) fake_proxy = types.ModuleType("litellm.proxy.proxy_server") fake_proxy.llm_router = router @@ -1145,6 +1146,7 @@ async def test_redis_async_embedding_forwards_full_metadata(monkeypatch): cache.embedding_model = "sem-embed" router = MagicMock() + router.get_configured_token_limits.return_value = (None, None) router.aembedding = AsyncMock(return_value={"data": [{"embedding": [0.1, 0.2]}]}) fake_proxy = types.ModuleType("litellm.proxy.proxy_server") fake_proxy.llm_router = router @@ -1162,6 +1164,100 @@ async def test_redis_async_embedding_forwards_full_metadata(monkeypatch): assert md["semantic-cache-embedding"] is True +LONG_PROMPT = " ".join(f"token{i}" for i in range(300)) + + +def _proxy_with_router(monkeypatch: pytest.MonkeyPatch, router: MagicMock, model_name: str) -> None: + import sys + import types + + fake_proxy = types.ModuleType("litellm.proxy.proxy_server") + fake_proxy.llm_router = router + fake_proxy.llm_model_list = [{"model_name": model_name}] + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_proxy) + + +def _token_count(model: str, text: str) -> int: + import litellm + + return len(litellm.encode(model=model, text=text)) + + +def test_redis_get_embedding_truncates_to_deployment_max_input_tokens(monkeypatch): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + cache = RedisSemanticCache.__new__(RedisSemanticCache) + cache.embedding_model = "sem-embed" + + router = MagicMock() + router.get_configured_token_limits.return_value = (5, None) + router.embedding = MagicMock(return_value={"data": [{"embedding": [0.5, 0.6]}]}) + _proxy_with_router(monkeypatch, router, "sem-embed") + + assert cache._get_embedding(LONG_PROMPT) == [0.5, 0.6] + + sent_input = router.embedding.call_args.kwargs["input"] + assert LONG_PROMPT.startswith(sent_input) + assert _token_count("sem-embed", sent_input) == 5 + assert _token_count("sem-embed", LONG_PROMPT) > 5 + + +@pytest.mark.asyncio +async def test_redis_async_embedding_explicit_limit_beats_deployment_limit(monkeypatch): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + cache = RedisSemanticCache.__new__(RedisSemanticCache) + cache.embedding_model = "sem-embed" + cache.embedding_max_input_tokens = 3 + + router = MagicMock() + router.get_configured_token_limits.return_value = (8191, None) + router.aembedding = AsyncMock(return_value={"data": [{"embedding": [0.1, 0.2]}]}) + _proxy_with_router(monkeypatch, router, "sem-embed") + + assert await cache._get_async_embedding(LONG_PROMPT) == [0.1, 0.2] + + sent_input = router.aembedding.call_args.kwargs["input"] + assert _token_count("sem-embed", sent_input) == 3 + + +def test_redis_get_embedding_truncates_direct_path_with_explicit_limit(monkeypatch): + import sys + import types + + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + cache = RedisSemanticCache.__new__(RedisSemanticCache) + cache.embedding_model = "text-embedding-3-small" + cache.embedding_max_input_tokens = 4 + + fake_proxy = types.ModuleType("litellm.proxy.proxy_server") + fake_proxy.llm_router = None + fake_proxy.llm_model_list = None + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_proxy) + + with patch( + "litellm.embedding", return_value={"data": [{"embedding": [0.1, 0.2]}]} + ) as direct_embed: + cache._get_embedding(LONG_PROMPT) + + sent_input = direct_embed.call_args.kwargs["input"] + assert _token_count("text-embedding-3-small", sent_input) == 4 + + +def test_redis_semantic_cache_init_stores_embedding_max_input_tokens(monkeypatch): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + cache = RedisSemanticCache( + redis_url="redis://localhost:6379", + similarity_threshold=0.8, + embedding_max_input_tokens=512, + ) + assert cache.embedding_max_input_tokens == 512 + default_cache = RedisSemanticCache(redis_url="redis://localhost:6379", similarity_threshold=0.8) + assert default_cache.embedding_max_input_tokens is None + + def test_redis_init_defers_redisvl_construction(monkeypatch): semantic_cache_mock = MagicMock() custom_vectorizer_mock = MagicMock() diff --git a/tests/test_litellm/caching/test_valkey_semantic_cache.py b/tests/test_litellm/caching/test_valkey_semantic_cache.py index d2df0a98e12..acf5a914e5c 100644 --- a/tests/test_litellm/caching/test_valkey_semantic_cache.py +++ b/tests/test_litellm/caching/test_valkey_semantic_cache.py @@ -105,6 +105,17 @@ def test_init_requires_similarity_threshold(): ValkeySemanticCache(sync_client=MagicMock(), async_client=AsyncMock()) +def test_init_stores_embedding_max_input_tokens(): + cache = ValkeySemanticCache( + similarity_threshold=0.8, + sync_client=MagicMock(), + async_client=AsyncMock(), + embedding_max_input_tokens=512, + ) + assert cache.embedding_max_input_tokens == 512 + assert _make_cache().embedding_max_input_tokens is None + + def test_init_rejects_cluster_startup_nodes(): with pytest.raises(ValueError, match="cluster-mode-enabled"): ValkeySemanticCache( diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_guardrail_cost.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_guardrail_cost.py new file mode 100644 index 00000000000..cf36a2b9b25 --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_guardrail_cost.py @@ -0,0 +1,113 @@ +import os + +import pytest + +import litellm +from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import ( + bedrock_guardrail_cost, + cost_breakdown_with_guardrail, + guardrail_information_cost, +) + + +@pytest.fixture +def synthetic_cost_map(monkeypatch): + monkeypatch.setattr( + litellm, + "model_cost", + { + "bedrock/guardrails": { + "guardrail_cost_per_unit": { + "contentPolicyUnits": 0.00015, + "topicPolicyUnits": 0.00015, + "wordPolicyUnits": 0.0, + } + }, + "bedrock/eu-west-1/guardrails": {"guardrail_cost_per_unit": {"contentPolicyUnits": 0.0002}}, + "bedrock/us-west-2/guardrails": {"guardrail_cost_per_unit": "malformed"}, + }, + ) + + +def test_bedrock_guardrail_cost_prices_each_counter(synthetic_cost_map): + cost = bedrock_guardrail_cost( + usage_units={"contentPolicyUnits": 2, "topicPolicyUnits": 1, "wordPolicyUnits": 5}, + aws_region_name="us-east-1", + ) + assert cost == pytest.approx(0.00045) + + +def test_bedrock_guardrail_cost_prefers_regional_entry(synthetic_cost_map): + cost = bedrock_guardrail_cost(usage_units={"contentPolicyUnits": 1}, aws_region_name="eu-west-1") + assert cost == pytest.approx(0.0002) + + +def test_bedrock_guardrail_cost_unknown_counter_is_free(synthetic_cost_map): + assert bedrock_guardrail_cost(usage_units={"someFutureCounter": 3}, aws_region_name="us-east-1") == 0.0 + + +def test_bedrock_guardrail_cost_malformed_regional_entry_falls_back(synthetic_cost_map): + cost = bedrock_guardrail_cost(usage_units={"contentPolicyUnits": 1}, aws_region_name="us-west-2") + assert cost == pytest.approx(0.00015) + + +def test_bedrock_guardrail_cost_no_pricing_entry(monkeypatch): + monkeypatch.setattr(litellm, "model_cost", {}) + assert bedrock_guardrail_cost(usage_units={"contentPolicyUnits": 1}, aws_region_name="us-east-1") == 0.0 + + +def test_shipped_bedrock_guardrail_prices_match_aws_pricing_page(): + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + assert litellm.model_cost["bedrock/guardrails"]["guardrail_cost_per_unit"] == { + "automatedReasoningPolicyUnits": 0.00017, + "contentPolicyImageUnits": 0.00075, + "contentPolicyUnits": 0.00015, + "contextualGroundingPolicyUnits": 0.0001, + "sensitiveInformationPolicyFreeUnits": 0.0, + "sensitiveInformationPolicyUnits": 0.0001, + "topicPolicyUnits": 0.00015, + "wordPolicyUnits": 0.0, + } + assert "bedrock/guardrails" not in litellm.bedrock_models + + +def test_guardrail_information_cost_sums_entries(): + entries = [ + {"guardrail_name": "a", "guardrail_cost": 0.0003}, + {"guardrail_name": "b", "guardrail_cost": None}, + {"guardrail_name": "c"}, + {"guardrail_name": "d", "guardrail_cost": 0.0001}, + ] + assert guardrail_information_cost(entries) == pytest.approx(0.0004) + + +def test_guardrail_information_cost_single_entry_and_garbage(): + assert guardrail_information_cost({"guardrail_cost": 0.0001}) == pytest.approx(0.0001) + assert guardrail_information_cost(None) == 0.0 + assert guardrail_information_cost("not-guardrail-info") == 0.0 + assert guardrail_information_cost([{"guardrail_cost": "bad"}]) == 0.0 + + +def test_guardrail_information_cost_ignores_negative_and_non_finite(): + entries = [ + {"guardrail_name": "forged-negative", "guardrail_cost": -0.005}, + {"guardrail_name": "forged-nan", "guardrail_cost": float("nan")}, + {"guardrail_name": "forged-inf", "guardrail_cost": float("inf")}, + {"guardrail_name": "real", "guardrail_cost": 0.0003}, + ] + assert guardrail_information_cost(entries) == pytest.approx(0.0003) + assert guardrail_information_cost({"guardrail_cost": -1.0}) == 0.0 + + +def test_cost_breakdown_with_guardrail_merges_and_creates(): + assert cost_breakdown_with_guardrail(None, 0.0) is None + untouched = {"input_cost": 0.1, "total_cost": 0.4} + assert cost_breakdown_with_guardrail(untouched, 0.0) is untouched + merged = cost_breakdown_with_guardrail({"input_cost": 0.1, "total_cost": 0.4}, 0.0003) + assert merged is not None + assert merged["guardrail_cost"] == pytest.approx(0.0003) + assert merged["total_cost"] == pytest.approx(0.4003) + assert merged["input_cost"] == pytest.approx(0.1) + created = cost_breakdown_with_guardrail(None, 0.0003) + assert created == {"guardrail_cost": 0.0003, "total_cost": 0.0003} diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 946f19b7658..0d54680fa81 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -4826,3 +4826,87 @@ async def test_restore_correlation_context_works_across_asyncio_task_boundary(): finally: trace_id_var.set("") session_id_var.set("") + + +def _build_success_payload(logging_obj, kwargs): + import datetime + + from litellm.litellm_core_utils.litellm_logging import ( + get_standard_logging_object_payload, + ) + + now = datetime.datetime.now() + return get_standard_logging_object_payload( + kwargs=kwargs, + init_response_obj={}, + start_time=now, + end_time=now, + logging_obj=logging_obj, + status="success", + ) + + +def _guardrail_kwargs(response_cost): + return { + "litellm_call_id": "guardrail-cost-call", + "model": "gpt-4o", + "messages": [], + "response_cost": response_cost, + "litellm_params": { + "metadata": { + "standard_logging_guardrail_information": [ + { + "guardrail_name": "bedrock-pre", + "guardrail_status": "success", + "guardrail_usage": {"topicPolicyUnits": 1, "contentPolicyUnits": 1}, + "guardrail_cost": 0.0003, + }, + {"guardrail_name": "no-usage-guardrail", "guardrail_status": "success"}, + ] + } + }, + } + + +def test_payload_response_cost_includes_guardrail_cost(logging_obj): + """LIT-5651: provider-billed guardrail cost must count in response_cost.""" + payload = _build_success_payload(logging_obj, _guardrail_kwargs(response_cost=0.0000429)) + + assert payload is not None + assert payload["response_cost"] == pytest.approx(0.0003429) + assert payload["cost_breakdown"] is not None + assert payload["cost_breakdown"]["guardrail_cost"] == pytest.approx(0.0003) + assert payload["cost_breakdown"]["total_cost"] == pytest.approx(0.0003) + assert payload["hidden_params"]["response_cost"] == pytest.approx(0.0000429) + + +def test_payload_guardrail_cost_merges_into_existing_cost_breakdown(logging_obj): + logging_obj.set_cost_breakdown( + input_cost=0.00003, + output_cost=0.0000129, + total_cost=0.0000429, + cost_for_built_in_tools_cost_usd_dollar=0.0, + ) + payload = _build_success_payload(logging_obj, _guardrail_kwargs(response_cost=0.0000429)) + + assert payload is not None + assert payload["response_cost"] == pytest.approx(0.0003429) + assert payload["cost_breakdown"]["guardrail_cost"] == pytest.approx(0.0003) + assert payload["cost_breakdown"]["total_cost"] == pytest.approx(0.0003429) + assert payload["cost_breakdown"]["input_cost"] == pytest.approx(0.00003) + assert logging_obj.cost_breakdown["total_cost"] == pytest.approx(0.0000429) + + +def test_payload_without_guardrail_cost_is_unchanged(logging_obj): + kwargs = { + "litellm_call_id": "no-guardrail-call", + "model": "gpt-4o", + "messages": [], + "response_cost": 0.0000429, + "litellm_params": {"metadata": {}}, + } + payload = _build_success_payload(logging_obj, kwargs) + + assert payload is not None + assert payload["response_cost"] == pytest.approx(0.0000429) + assert payload["cost_breakdown"] is None diff --git a/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_transformation.py b/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_transformation.py index 35c0a63573f..93c518599d6 100644 --- a/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_transformation.py +++ b/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_transformation.py @@ -8,7 +8,7 @@ especially ensuring that encoding_format is not included when not provided. import json import os import sys -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import Mock, patch import pytest @@ -289,6 +289,49 @@ class TestHostedVLLMEmbeddingTransformation: assert sent_data["model"] == "BAAI/bge-small-en-v1.5" assert sent_data["input"] == ["Hello world"] + @pytest.mark.parametrize( + "provider_params", + [ + {"extra_body": {"truncate": "END", "input_type": "query"}}, + {"truncate": "END", "input_type": "query"}, + ], + ) + def test_provider_params_are_sent_at_the_top_level_of_the_request(self, provider_params: dict[str, object]) -> None: + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + client = HTTPHandler() + + with patch.object(HTTPHandler, "post") as mock_post: + mock_response = Mock() + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_response.json.return_value = { + "object": "list", + "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}], + "model": "nvidia/nv-embedqa-e5-v5", + "usage": {"prompt_tokens": 2, "total_tokens": 2}, + } + mock_response.text = json.dumps(mock_response.json.return_value) + mock_post.return_value = mock_response + + litellm.embedding( + model="hosted_vllm/nvidia/nv-embedqa-e5-v5", + input=["Hello world"], + api_base="https://integrate.api.nvidia.com/v1", + api_key="fake-key", + client=client, + caching=False, + **provider_params, + ) + + sent_data = json.loads(mock_post.call_args.kwargs["data"]) + + assert sent_data["truncate"] == "END" + assert sent_data["input_type"] == "query" + assert "extra_body" not in sent_data + assert sent_data["model"] == "nvidia/nv-embedqa-e5-v5" + assert sent_data["input"] == ["Hello world"] + if __name__ == "__main__": pytest.main([__file__, "-v", "-s"]) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index 53921e7e74a..e3516b6eda7 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -5079,24 +5079,200 @@ async def test_apply_guardrail_failure_logs_a_dict_not_a_bare_string(): assert "error" in logged -def test_build_tracing_detail_surfaces_usage_counters(): - """LIT-5650: the billable usage block Bedrock returns per ApplyGuardrail call must - land on the tracing detail as guardrail_usage so it reaches spend logs as a - sibling of guardrail_response (which default redaction replaces wholesale).""" +def test_build_tracing_detail_surfaces_usage_counters_and_cost(monkeypatch): + """LIT-5650/LIT-5651: AWS-billed usage must land as guardrail_usage priced into guardrail_cost.""" + monkeypatch.setattr( + litellm, + "model_cost", + { + "bedrock/guardrails": { + "guardrail_cost_per_unit": { + "topicPolicyUnits": 0.00015, + "contentPolicyUnits": 0.00015, + "wordPolicyUnits": 0.0, + } + } + }, + ) guardrail = BedrockGuardrail(guardrailIdentifier="test-guardrail", guardrailVersion="DRAFT") detail = guardrail._build_tracing_detail( { "action": "GUARDRAIL_INTERVENED", "usage": {"topicPolicyUnits": 1, "contentPolicyUnits": 2, "wordPolicyUnits": 0, "oddball": "not-an-int"}, - } + }, + aws_region_name="us-east-1", ) assert detail["guardrail_usage"] == {"topicPolicyUnits": 1, "contentPolicyUnits": 2, "wordPolicyUnits": 0} + assert detail["guardrail_cost"] == pytest.approx(0.00045) def test_build_tracing_detail_omits_guardrail_usage_when_bedrock_reports_none(): guardrail = BedrockGuardrail(guardrailIdentifier="test-guardrail", guardrailVersion="DRAFT") - assert "guardrail_usage" not in guardrail._build_tracing_detail({"action": "NONE"}) - assert "guardrail_usage" not in guardrail._build_tracing_detail({"action": "NONE", "usage": {}}) + for detail in ( + guardrail._build_tracing_detail({"action": "NONE"}, aws_region_name="us-east-1"), + guardrail._build_tracing_detail({"action": "NONE", "usage": {}}, aws_region_name="us-east-1"), + ): + assert "guardrail_usage" not in detail + assert "guardrail_cost" not in detail + + +@pytest.mark.asyncio +async def test_blocked_chunk_logs_usage_and_cost_of_prior_passed_chunks(monkeypatch): + """LIT-5651 regression: a block on a later chunk must still bill the chunks AWS already processed.""" + monkeypatch.setattr( + litellm, + "model_cost", + { + "bedrock/guardrails": { + "guardrail_cost_per_unit": { + "contentPolicyUnits": 0.00015, + "wordPolicyUnits": 0.0, + } + } + }, + ) + guardrail = BedrockGuardrail( + guardrailIdentifier="test-guardrail", + guardrailVersion="DRAFT", + chunk_budget_chars=40, + ) + + too_large_response = MagicMock() + too_large_response.status_code = 429 + too_large_response.json.return_value = { + "message": "Input text size (60 text units) exceeds the maximum allowed (1 text units) for the content filter policy" + } + + passed_chunk_response = MagicMock() + passed_chunk_response.status_code = 200 + passed_chunk_response.json.return_value = { + "action": "NONE", + "outputs": [], + "assessments": [], + "usage": {"contentPolicyUnits": 2, "wordPolicyUnits": 1}, + } + + blocked_chunk_response = MagicMock() + blocked_chunk_response.status_code = 200 + blocked_chunk_response.json.return_value = { + "action": "GUARDRAIL_INTERVENED", + "assessments": [{"contentPolicy": {"filters": [{"type": "HATE", "confidence": "HIGH", "action": "BLOCKED"}]}}], + "outputs": [{"text": "Content blocked"}], + "usage": {"contentPolicyUnits": 3}, + } + + mock_credentials = MagicMock() + mock_credentials.access_key = "test-access-key" + mock_credentials.secret_key = "test-secret-key" + mock_credentials.token = None + + request_data = { + "model": "gpt-4o", + "messages": [ + {"role": "user", "content": "a" * 30}, + {"role": "user", "content": "b" * 30}, + ], + } + + with ( + patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post, + patch.object(guardrail, "_load_credentials", return_value=(mock_credentials, "us-east-1")), + patch.object(guardrail, "_prepare_request", return_value=MagicMock()), + ): + mock_post.side_effect = [too_large_response, passed_chunk_response, blocked_chunk_response] + + with pytest.raises(HTTPException): + await guardrail.make_bedrock_api_request( + source="INPUT", + messages=request_data["messages"], + request_data=request_data, + ) + + assert mock_post.call_count == 3 + logged_entries = request_data["metadata"]["standard_logging_guardrail_information"] + assert len(logged_entries) == 1 + logged = logged_entries[0] + assert logged["guardrail_usage"] == {"contentPolicyUnits": 5, "wordPolicyUnits": 1} + assert logged["guardrail_cost"] == pytest.approx(0.00075) + assert logged["guardrail_response"]["usage"] == {"contentPolicyUnits": 5, "wordPolicyUnits": 1} + + +@pytest.mark.asyncio +async def test_terminal_failure_logs_usage_and_cost_of_prior_passed_chunks(monkeypatch): + """LIT-5651 regression: a terminal failure on a later chunk must still bill the chunks AWS already processed.""" + monkeypatch.setattr( + litellm, + "model_cost", + { + "bedrock/guardrails": { + "guardrail_cost_per_unit": { + "contentPolicyUnits": 0.00015, + "wordPolicyUnits": 0.0, + } + } + }, + ) + guardrail = BedrockGuardrail( + guardrailIdentifier="test-guardrail", + guardrailVersion="DRAFT", + chunk_budget_chars=40, + ) + + too_large_response = MagicMock() + too_large_response.status_code = 429 + too_large_response.json.return_value = { + "message": "Input text size (60 text units) exceeds the maximum allowed (1 text units) for the content filter policy" + } + + passed_chunk_response = MagicMock() + passed_chunk_response.status_code = 200 + passed_chunk_response.json.return_value = { + "action": "NONE", + "outputs": [], + "assessments": [], + "usage": {"contentPolicyUnits": 2, "wordPolicyUnits": 1}, + } + + failed_chunk_response = MagicMock() + failed_chunk_response.status_code = 400 + failed_chunk_response.json.return_value = {"message": "ValidationException: guardrail is in a failed state"} + + mock_credentials = MagicMock() + mock_credentials.access_key = "test-access-key" + mock_credentials.secret_key = "test-secret-key" + mock_credentials.token = None + + request_data = { + "model": "gpt-4o", + "messages": [ + {"role": "user", "content": "a" * 30}, + {"role": "user", "content": "b" * 30}, + ], + } + + with ( + patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post, + patch.object(guardrail, "_load_credentials", return_value=(mock_credentials, "us-east-1")), + patch.object(guardrail, "_prepare_request", return_value=MagicMock()), + ): + mock_post.side_effect = [too_large_response, passed_chunk_response, failed_chunk_response] + + with pytest.raises(HTTPException): + await guardrail.make_bedrock_api_request( + source="INPUT", + messages=request_data["messages"], + request_data=request_data, + ) + + assert mock_post.call_count == 3 + logged_entries = request_data["metadata"]["standard_logging_guardrail_information"] + assert len(logged_entries) == 1 + logged = logged_entries[0] + assert logged["guardrail_status"] == "guardrail_failed_to_respond" + assert logged["guardrail_usage"] == {"contentPolicyUnits": 2, "wordPolicyUnits": 1} + assert logged["guardrail_cost"] == pytest.approx(0.0003) + assert logged["guardrail_response"]["usage"] == {"contentPolicyUnits": 2, "wordPolicyUnits": 1} + assert "error" in logged["guardrail_response"] diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index bca8210baa6..50c93ed5275 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -17,7 +17,7 @@ from litellm.proxy.hooks.proxy_track_cost_callback import ( _should_track_cost_callback, _update_database_and_spend_counters, ) -from litellm.types.utils import CallTypes +from litellm.types.utils import CallTypes, Usage @pytest.mark.asyncio @@ -152,6 +152,70 @@ async def test_async_post_call_failure_hook_does_not_clobber_guardrail_info_in_m assert metadata["standard_logging_guardrail_information"] == metadata_bucket_info +@pytest.mark.asyncio +async def test_async_post_call_failure_hook_bills_guardrail_cost_on_blocked_request(): + """LIT-5651: a request blocked by a guardrail never reaches the LLM, but the + guardrail invocation itself is billed by the provider. The failure row must + charge that cost against the key instead of recording zero spend.""" + logger = _ProxyDBLogger() + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "metadata": { + "standard_logging_guardrail_information": [ + { + "guardrail_name": "bedrock-guard", + "guardrail_status": "guardrail_intervened", + "guardrail_usage": {"topicPolicyUnits": 1, "contentPolicyUnits": 1}, + "guardrail_cost": 0.0003, + } + ] + }, + "proxy_server_request": {"request_id": "test_request_id"}, + } + + with patch( + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, + ) as mock_update_database: + await logger.async_post_call_failure_hook( + request_data=request_data, + original_exception=Exception("Violated guardrail policy"), + user_api_key_dict=UserAPIKeyAuth(api_key="test_api_key"), + ) + + assert mock_update_database.call_args[1]["response_cost"] == pytest.approx(0.0003) + + +@pytest.mark.asyncio +async def test_async_post_call_failure_hook_adds_guardrail_cost_to_recovered_stream_cost(): + logger = _ProxyDBLogger() + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "metadata": { + "standard_logging_guardrail_information": [ + {"guardrail_name": "bedrock-guard", "guardrail_status": "success", "guardrail_cost": 0.0003} + ] + }, + "proxy_server_request": {"request_id": "test_request_id"}, + "combined_usage_object": Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), + "response_cost": 0.001, + } + + with patch( + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, + ) as mock_update_database: + await logger.async_post_call_failure_hook( + request_data=request_data, + original_exception=Exception("stream broke mid-flight"), + user_api_key_dict=UserAPIKeyAuth(api_key="test_api_key"), + ) + + assert mock_update_database.call_args[1]["response_cost"] == pytest.approx(0.0013) + + @pytest.mark.asyncio async def test_async_post_call_failure_hook_non_llm_route(): # Setup diff --git a/tests/test_litellm/proxy/test_pricing_field_strip.py b/tests/test_litellm/proxy/test_pricing_field_strip.py index 25377a6d209..bbdddd1cd8c 100644 --- a/tests/test_litellm/proxy/test_pricing_field_strip.py +++ b/tests/test_litellm/proxy/test_pricing_field_strip.py @@ -102,6 +102,30 @@ class TestStripClientPricingOverrides: assert data["metadata"] == {"user_session": "keep-me"} assert data["litellm_metadata"] == {} + def test_metadata_guardrail_information_dropped(self): + # Client-seeded guardrail entries would otherwise be summed into + # response_cost and spend, letting a caller forge (even negative) + # guardrail cost against their own budget. + data = { + "model": "gpt-4", + "metadata": { + "user_session": "keep-me", + "standard_logging_guardrail_information": [ + { + "guardrail_name": "forged", + "guardrail_status": "success", + "guardrail_cost": -0.005, + } + ], + }, + "litellm_metadata": { + "standard_logging_guardrail_information": [{"guardrail_cost": 5.0}], + }, + } + _strip_client_pricing_overrides(data) + assert data["metadata"] == {"user_session": "keep-me"} + assert data["litellm_metadata"] == {} + def test_non_pricing_fields_untouched(self): data = { "model": "gpt-4", @@ -129,6 +153,7 @@ class TestStripClientPricingOverrides: def test_metadata_field_set_contains_model_info(self): assert "model_info" in _CLIENT_PRICING_METADATA_FIELDS + assert "standard_logging_guardrail_information" in _CLIENT_PRICING_METADATA_FIELDS def test_strip_emits_debug_log_listing_dropped_fields(self, caplog): # Operators need a paper trail so they can diagnose why a previously diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 99a824511f6..afdfdf170ac 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -869,6 +869,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "container", "image_edit", "embedding", + "guardrail", "image_generation", "video_generation", "moderation", @@ -976,6 +977,10 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "type": "string", }, }, + "guardrail_cost_per_unit": { + "type": "object", + "additionalProperties": {"type": "number"}, + }, "search_context_cost_per_query": { "type": "object", "properties": { diff --git a/type-discipline-budget.json b/type-discipline-budget.json index b5c07de8cd3..94c1f9b86f9 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,9 +1,9 @@ { "LIT001": { - "limit": 22900 + "limit": 22897 }, "LIT002": { - "limit": 26889 + "limit": 26888 }, "LIT003": { "limit": 269 diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index e467ddc5ec4..bc23c2433d3 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -16,7 +16,7 @@ }, "src/app/(dashboard)/admin-panel/_components/AdminPanel.tsx": { "no-restricted-imports": { - "count": 2 + "count": 1 }, "react-hooks/set-state-in-effect": { "count": 1 @@ -1647,9 +1647,6 @@ } }, "src/components/SCIM.tsx": { - "no-restricted-imports": { - "count": 2 - }, "react-hooks/set-state-in-effect": { "count": 1 } @@ -1661,7 +1658,7 @@ }, "src/components/SSOModals.tsx": { "no-restricted-imports": { - "count": 2 + "count": 1 } }, "src/components/Settings/AdminSettings/HashicorpVault/EditHashicorpVaultModal.tsx": { @@ -1705,11 +1702,6 @@ "count": 1 } }, - "src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx": { - "no-restricted-imports": { - "count": 2 - } - }, "src/components/Settings/AdminSettings/SSOSettings/Modals/EditSSOSettingsModal.test.tsx": { "max-nested-callbacks": { "count": 1 @@ -1760,7 +1752,7 @@ }, "src/components/Settings/RouterSettings/Fallbacks/FallbackSelectionForm.tsx": { "no-restricted-imports": { - "count": 2 + "count": 1 }, "react-hooks/set-state-in-effect": { "count": 1 @@ -2024,10 +2016,7 @@ "count": 1 }, "no-nested-ternary": { - "count": 4 - }, - "no-restricted-imports": { - "count": 2 + "count": 1 } }, "src/components/bulk_create_users_button.tsx": { @@ -2091,7 +2080,7 @@ "count": 1 }, "no-restricted-imports": { - "count": 2 + "count": 1 }, "no-restricted-syntax": { "count": 3 @@ -2461,7 +2450,7 @@ "count": 1 }, "no-restricted-imports": { - "count": 2 + "count": 1 } }, "src/components/organisms/RegenerateKeyModal.tsx": { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/admin-panel/_components/AdminPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/admin-panel/_components/AdminPanel.tsx index 3f3f7fb5db7..48eaa62e37b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/admin-panel/_components/AdminPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/admin-panel/_components/AdminPanel.tsx @@ -3,18 +3,12 @@ * Use this to avoid sharing master key with others */ import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; -import { - Button, - Callout, - Card, - Table, - TableBody, - TableCell, - TableHead, - TableHeaderCell, - TableRow, -} from "@tremor/react"; -import { Alert, Modal, Space, Tabs, Typography } from "antd"; +import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert"; +import { Button } from "@/components/ui/button"; +import { Card } from "@/components/ui/card"; +import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table"; +import { Alert as AntdAlert, Modal, Space, Tabs, Typography } from "antd"; +import { Info } from "lucide-react"; import React, { useEffect, useState } from "react"; import NewBadge from "@/components/common_components/NewBadge"; import { useBaseUrl } from "@/components/constants"; @@ -37,7 +31,6 @@ import UIAccessControlForm from "@/components/UIAccessControlForm"; import { z } from "zod/v4"; import { FieldGroup } from "@/components/shared/form/field"; import { FormField } from "@/components/shared/form/FormField"; -import { Button as ShadcnButton } from "@/components/ui/button"; import { Input } from "@/components/ui/input"; import { useZodForm } from "@/lib/forms/useZodForm"; @@ -59,7 +52,7 @@ const AddAllowedIPForm = ({ onSubmit }: { onSubmit: (values: AllowedIPFormValues {({ ref, ...field }) => }