Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/competent-lewin-1c8fd9

This commit is contained in:
Yuneng Jiang 2026-08-18 15:49:12 -07:00
commit c71b6ed51b
No known key found for this signature in database
44 changed files with 1304 additions and 274 deletions

View file

@ -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).",

View file

@ -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":

View file

@ -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

View file

@ -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()

View file

@ -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},
)

View file

@ -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"]

View file

@ -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

View file

@ -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),

View file

@ -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

View file

@ -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

View file

@ -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",

View file

@ -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(

View file

@ -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]:

View file

@ -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,

View file

@ -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

View file

@ -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)

View file

@ -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:

View file

@ -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",

View file

@ -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",

View file

@ -33,7 +33,7 @@
"limit": 2
},
"B006": {
"limit": 178
"limit": 177
},
"B008": {
"limit": 503

View file

@ -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"}

View file

@ -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

View file

@ -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

View file

@ -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()

View file

@ -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(

View file

@ -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}

View file

@ -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

View file

@ -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"])

View file

@ -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"]

View file

@ -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

View file

@ -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

View file

@ -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": {

View file

@ -1,9 +1,9 @@
{
"LIT001": {
"limit": 22900
"limit": 22897
},
"LIT002": {
"limit": 26889
"limit": 26888
},
"LIT003": {
"limit": 269

View file

@ -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": {

View file

@ -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 }) => <Input ref={ref} placeholder="Enter IP address" {...field} />}
</FormField>
<div>
<ShadcnButton type="submit">Add IP Address</ShadcnButton>
<Button type="submit">Add IP Address</Button>
</div>
</FieldGroup>
</form>
@ -228,9 +221,9 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
label: "Security Settings",
children: (
<>
<Card>
<Card className="block p-6">
<Title level={4}> ✨ Security Settings</Title>
<Alert
<AntdAlert
message="SSO Configuration Deprecated"
description="Editing SSO Settings on this page is deprecated and will be removed in a future version. Please use the SSO Settings tab for SSO configuration."
type="warning"
@ -298,19 +291,19 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
]}
>
<Table>
<TableHead>
<TableHeader>
<TableRow>
<TableHeaderCell>IP Address</TableHeaderCell>
<TableHeaderCell className="text-right">Action</TableHeaderCell>
<TableHead>IP Address</TableHead>
<TableHead className="text-right">Action</TableHead>
</TableRow>
</TableHead>
</TableHeader>
<TableBody>
{allowedIPs.map((ip, index) => (
<TableRow key={index}>
<TableCell>{ip}</TableCell>
<TableCell className="text-right">
{ip !== all_ip_address_allowed && (
<Button onClick={() => handleDeleteIP(ip)} color="red" size="xs">
<Button onClick={() => handleDeleteIP(ip)} variant="destructive" size="sm">
Delete
</Button>
)}
@ -365,12 +358,16 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
/>
</Modal>
</div>
<Callout title="Login without SSO" color="teal">
If you need to login without sso, you can access{" "}
<a href={nonSssoUrl} target="_blank" rel="noopener noreferrer">
<b>{nonSssoUrl}</b>{" "}
</a>
</Callout>
<Alert variant="info">
<Info />
<AlertTitle>Login without SSO</AlertTitle>
<AlertDescription>
If you need to login without sso, you can access{" "}
<a href={nonSssoUrl} target="_blank" rel="noopener noreferrer">
<b>{nonSssoUrl}</b>{" "}
</a>
</AlertDescription>
</Alert>
</>
),
},

View file

@ -49,7 +49,7 @@ const langgraphInfo: AgentCreateInfo = {
const renderForm = () =>
render(<AddAgentForm visible={true} onClose={vi.fn()} accessToken="tok" onSuccess={vi.fn()} />);
const panel = (name: RegExp) => screen.getByRole("button", { name });
const panel = (name: RegExp) => screen.findByRole("button", { name });
const openAgentTypeMenu = async (user: ReturnType<typeof userEvent.setup>) => {
await user.click(screen.getAllByRole("combobox")[0]);
@ -102,7 +102,7 @@ describe("AddAgentForm submit payload", () => {
await user.clear(screen.getByLabelText("Version"));
await user.type(screen.getByLabelText("Version"), "2.0.0");
await user.click(panel(/Skills/));
await user.click(await panel(/Skills/));
await user.click(screen.getByRole("button", { name: /Add Skill/ }));
await user.type(await screen.findByLabelText("Skill ID"), "hello");
await user.type(screen.getByLabelText("Skill Name"), "Hello");
@ -111,22 +111,22 @@ describe("AddAgentForm submit payload", () => {
await user.type(screen.getByLabelText("Examples"), "say hi");
await user.click(screen.getByLabelText("Agent Name"));
await user.click(panel(/Capabilities/));
await user.click(await panel(/Capabilities/));
await user.click(await screen.findByRole("switch", { name: "Streaming" }));
await user.click(screen.getByRole("switch", { name: "Push Notifications" }));
await user.click(panel(/Optional Settings/));
await user.click(await panel(/Optional Settings/));
await user.type(await screen.findByLabelText("Icon URL"), "https://example.com/icon.png");
await user.click(panel(/Cost Configuration/));
await user.click(await panel(/Cost Configuration/));
await user.type(await screen.findByLabelText("Cost Per Query ($)"), "0.25");
await user.type(screen.getByLabelText("Input Cost Per Token ($)"), "0.000002");
await user.click(panel(/LiteLLM Parameters/));
await user.click(await panel(/LiteLLM Parameters/));
await user.type(await screen.findByLabelText("Model (Optional)"), "gpt-4o");
await user.click(screen.getByRole("switch", { name: "Make Public" }));
await user.click(panel(/Authentication Headers/));
await user.click(await panel(/Authentication Headers/));
await user.click(await screen.findByRole("button", { name: /Add Static Header/ }));
await user.type(await screen.findByPlaceholderText("Header name (e.g. Authorization)"), "X-Tenant");
await user.type(screen.getByPlaceholderText("Value (e.g. Bearer token123)"), "acme");
@ -176,9 +176,9 @@ describe("AddAgentForm submit payload", () => {
await user.type(screen.getByLabelText("Display Name"), "Collapsed");
await user.type(screen.getByPlaceholderText("Describe what this agent does..."), "d");
await user.click(panel(/Cost Configuration/));
await user.click(await panel(/Cost Configuration/));
await user.type(await screen.findByLabelText("Cost Per Query ($)"), "0.75");
await user.click(panel(/Cost Configuration/));
await user.click(await panel(/Cost Configuration/));
await goToLastStepAndCreate(user);
@ -189,10 +189,10 @@ describe("AddAgentForm submit payload", () => {
const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never });
renderForm();
await user.click(panel(/Cost Configuration/));
await user.click(await panel(/Cost Configuration/));
await user.type(await screen.findByLabelText("Cost Per Query ($)"), "0.75");
await user.click(panel(/Cost Configuration/));
await user.click(panel(/Cost Configuration/));
await user.click(await panel(/Cost Configuration/));
await user.click(await panel(/Cost Configuration/));
expect(await screen.findByLabelText("Cost Per Query ($)")).toHaveValue(0.75);
});

View file

@ -1,21 +1,17 @@
import React, { useState, useEffect } from "react";
import { Card, Title, Text, Grid, Callout, Divider } from "@tremor/react";
import { z } from "zod/v4";
import { keyCreateCall } from "./networking";
import { CopyToClipboard } from "react-copy-to-clipboard";
import {
LinkOutlined,
KeyOutlined,
CopyOutlined,
ExclamationCircleOutlined,
PlusCircleOutlined,
} from "@ant-design/icons";
import { CircleAlert, CirclePlus, Copy, Info, KeyRound, Link } from "lucide-react";
import { parseErrorMessage } from "./shared/errorUtils";
import { toast } from "@/lib/toast";
import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert";
import { FieldGroup } from "@/components/shared/form/field";
import { FormField } from "@/components/shared/form/FormField";
import { Button } from "@/components/ui/button";
import { Card, CardContent, CardTitle } from "@/components/ui/card";
import { Input } from "@/components/ui/input";
import { Separator } from "@/components/ui/separator";
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
import { useZodForm } from "@/lib/forms/useZodForm";
@ -80,114 +76,121 @@ const SCIMConfig: React.FC<SCIMConfigProps> = ({ accessToken, userID, proxySetti
};
return (
<Grid numItems={1}>
<div className="grid grid-cols-1">
<Card>
<div className="flex items-center mb-4">
<Title>SCIM Configuration</Title>
</div>
<Text className="text-muted-foreground">
System for Cross-domain Identity Management (SCIM) allows you to automatically provision and manage users and
groups in LiteLLM.
</Text>
<Divider />
<div className="space-y-8">
{/* Step 1: SCIM URL */}
<div>
<div className="flex items-center mb-2">
<div className="flex items-center justify-center w-6 h-6 rounded-full bg-blue-100 text-blue-700 dark:bg-blue-950 dark:text-blue-300 mr-2">
1
</div>
<Title className="text-lg flex items-center">
<LinkOutlined className="h-5 w-5 mr-2" />
SCIM Tenant URL
</Title>
</div>
<Text className="text-muted-foreground mb-3">
Use this URL in your identity provider SCIM integration settings.
</Text>
<div className="flex items-center">
<Input value={scimBaseUrl} disabled={true} readOnly className="grow" />
<CopyToClipboard text={scimBaseUrl} onCopy={() => toast.success("URL copied to clipboard")}>
<Button type="button" className="ml-2 flex items-center">
<CopyOutlined className="h-4 w-4 mr-1" />
Copy
</Button>
</CopyToClipboard>
</div>
<CardContent>
<div className="flex items-center mb-4">
<CardTitle>SCIM Configuration</CardTitle>
</div>
<p className="text-muted-foreground">
System for Cross-domain Identity Management (SCIM) allows you to automatically provision and manage users
and groups in LiteLLM.
</p>
{/* Step 2: SCIM Token */}
<div>
<div className="flex items-center mb-2">
<div className="flex items-center justify-center w-6 h-6 rounded-full bg-blue-100 text-blue-700 dark:bg-blue-950 dark:text-blue-300 mr-2">
2
<Separator className="my-6" />
<div className="space-y-8">
{/* Step 1: SCIM URL */}
<div>
<div className="flex items-center mb-2">
<div className="flex items-center justify-center w-6 h-6 rounded-full bg-blue-100 text-blue-700 dark:bg-blue-950 dark:text-blue-300 mr-2">
1
</div>
<h3 className="text-lg font-medium flex items-center">
<Link className="h-5 w-5 mr-2" />
SCIM Tenant URL
</h3>
</div>
<p className="text-muted-foreground mb-3">
Use this URL in your identity provider SCIM integration settings.
</p>
<div className="flex items-center">
<Input value={scimBaseUrl} disabled={true} readOnly className="grow" />
<CopyToClipboard text={scimBaseUrl} onCopy={() => toast.success("URL copied to clipboard")}>
<Button type="button" className="ml-2 flex items-center">
<Copy />
Copy
</Button>
</CopyToClipboard>
</div>
<Title className="text-lg flex items-center">
<KeyOutlined className="h-5 w-5 mr-2" />
Authentication Token
</Title>
</div>
<Callout title="Using SCIM" color="blue" className="mb-4">
You need a SCIM token to authenticate with the SCIM API. Create one below and use it in your SCIM provider
configuration.
</Callout>
{/* Step 2: SCIM Token */}
<div>
<div className="flex items-center mb-2">
<div className="flex items-center justify-center w-6 h-6 rounded-full bg-blue-100 text-blue-700 dark:bg-blue-950 dark:text-blue-300 mr-2">
2
</div>
<h3 className="text-lg font-medium flex items-center">
<KeyRound className="h-5 w-5 mr-2" />
Authentication Token
</h3>
</div>
{!tokenData ? (
<div className="bg-muted p-4 rounded-lg">
<form onSubmit={form.handleSubmit(handleCreateSCIMToken)}>
<FieldGroup>
<FormField control={form.control} name="key_alias" label="Token Name">
{({ ref, ...field }) => <Input {...field} ref={ref} placeholder="SCIM Access Token" />}
</FormField>
<div>
<Button type="submit" disabled={isCreatingToken} className="flex items-center">
{isCreatingToken ? (
<UiLoadingSpinner className="size-4 mr-1" />
) : (
<KeyOutlined className="h-4 w-4 mr-1" />
)}
Create SCIM Token
<Alert variant="info" className="mb-4">
<Info />
<AlertTitle>Using SCIM</AlertTitle>
<AlertDescription>
You need a SCIM token to authenticate with the SCIM API. Create one below and use it in your SCIM
provider configuration.
</AlertDescription>
</Alert>
{!tokenData ? (
<div className="bg-muted p-4 rounded-lg">
<form onSubmit={form.handleSubmit(handleCreateSCIMToken)}>
<FieldGroup>
<FormField control={form.control} name="key_alias" label="Token Name">
{({ ref, ...field }) => <Input {...field} ref={ref} placeholder="SCIM Access Token" />}
</FormField>
<div>
<Button
type="submit"
disabled={isCreatingToken}
aria-busy={isCreatingToken}
className="flex items-center"
>
{isCreatingToken ? <UiLoadingSpinner className="size-4" /> : <KeyRound />}
Create SCIM Token
</Button>
</div>
</FieldGroup>
</form>
</div>
) : (
<Card className="block p-6 border border-yellow-300 bg-yellow-50 dark:border-yellow-800 dark:bg-yellow-950">
<div className="flex items-center mb-2 text-yellow-800 dark:text-yellow-300">
<CircleAlert className="h-5 w-5 mr-2" />
<h4 className="text-lg font-medium text-yellow-800 dark:text-yellow-300">Your SCIM Token</h4>
</div>
<p className="text-yellow-800 dark:text-yellow-300 mb-4 font-medium">
Make sure to copy this token now. You will not be able to see it again.
</p>
<div className="flex items-center">
<Input value={tokenData.key} className="grow mr-2" type="password" disabled={true} readOnly />
<CopyToClipboard text={tokenData.key} onCopy={() => toast.success("Token copied to clipboard")}>
<Button type="button" className="flex items-center">
<Copy />
Copy
</Button>
</div>
</FieldGroup>
</form>
</div>
) : (
<Card className="border border-yellow-300 bg-yellow-50 dark:border-yellow-800 dark:bg-yellow-950">
<div className="flex items-center mb-2 text-yellow-800 dark:text-yellow-300">
<ExclamationCircleOutlined className="h-5 w-5 mr-2" />
<Title className="text-lg text-yellow-800 dark:text-yellow-300">Your SCIM Token</Title>
</div>
<Text className="text-yellow-800 dark:text-yellow-300 mb-4 font-medium">
Make sure to copy this token now. You will not be able to see it again.
</Text>
<div className="flex items-center">
<Input value={tokenData.key} className="grow mr-2" type="password" disabled={true} readOnly />
<CopyToClipboard text={tokenData.key} onCopy={() => toast.success("Token copied to clipboard")}>
<Button type="button" className="flex items-center">
<CopyOutlined className="h-4 w-4 mr-1" />
Copy
</Button>
</CopyToClipboard>
</div>
<Button
type="button"
variant="secondary"
className="mt-4 flex items-center"
onClick={() => setTokenData(null)}
>
<PlusCircleOutlined className="h-4 w-4 mr-1" />
Create Another Token
</Button>
</Card>
)}
</CopyToClipboard>
</div>
<Button
type="button"
variant="secondary"
className="mt-4 flex items-center"
onClick={() => setTokenData(null)}
>
<CirclePlus />
Create Another Token
</Button>
</Card>
)}
</div>
</div>
</div>
</CardContent>
</Card>
</Grid>
</div>
);
};

View file

@ -1,7 +1,6 @@
import React, { useEffect, useState } from "react";
import { FormProvider, useWatch, type UseFormReturn } from "react-hook-form";
import { Modal } from "antd";
import { Text } from "@tremor/react";
import { getSSOSettings, updateSSOSettings } from "./networking";
import { toast } from "@/lib/toast";
import { parseErrorMessage } from "./shared/errorUtils";
@ -317,10 +316,10 @@ const SSOModals: React.FC<SSOModalsProps> = ({
onCancel={handleInstructionsCancel}
>
<p>Follow these steps to complete the SSO setup:</p>
<Text className="mt-2">1. DO NOT Exit this TAB</Text>
<Text className="mt-2">2. Open a new tab, visit your proxy base url</Text>
<Text className="mt-2">3. Confirm your SSO is configured correctly and you can login on the new Tab</Text>
<Text className="mt-2">4. If Step 3 is successful, you can close this tab</Text>
<p className="text-sm mt-2">1. DO NOT Exit this TAB</p>
<p className="text-sm mt-2">2. Open a new tab, visit your proxy base url</p>
<p className="text-sm mt-2">3. Confirm your SSO is configured correctly and you can login on the new Tab</p>
<p className="text-sm mt-2">4. If Step 3 is successful, you can close this tab</p>
<div style={{ textAlign: "right", marginTop: "10px" }}>
<Button type="button" onClick={handleInstructionsOk}>
Done

View file

@ -4,7 +4,7 @@
* Decoupled from form submission logic
*/
import { Button } from "@tremor/react";
import { Button } from "@/components/ui/button";
import { Tabs } from "antd";
import { Plus } from "lucide-react";
import React, { useEffect, useState } from "react";
@ -98,7 +98,8 @@ export function FallbackSelectionForm({
return (
<div className="text-center py-12 bg-gray-50 rounded-lg border border-dashed border-gray-300">
<p className="text-gray-500 mb-4">No fallback groups configured</p>
<Button variant="primary" onClick={handleAddGroup} icon={() => <Plus className="w-4 h-4" />}>
<Button onClick={handleAddGroup}>
<Plus className="w-4 h-4" />
Create First Group
</Button>
</div>

View file

@ -174,9 +174,7 @@ describe("DynamicForm change notifications", () => {
const user = userEvent.setup();
const { handleResetField } = renderForm();
const row = screen.getByText("region_name").closest("tr");
const reset = row?.querySelector(".tremor-Icon-root");
await user.click(reset as Element);
await user.click(screen.getByRole("button", { name: "Reset region_name" }));
expect(handleResetField).toHaveBeenCalledWith("region_name", 1);
});

View file

@ -1,8 +1,11 @@
import React from "react";
import { useForm } from "react-hook-form";
import { TrashIcon, CheckCircleIcon } from "@heroicons/react/outline";
import { Button, Badge, Icon, Text, TableRow, TableCell, Switch } from "@tremor/react";
import { CircleCheck, Trash2 } from "lucide-react";
import { Badge } from "@/components/ui/badge";
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import { Switch } from "@/components/ui/switch";
import { TableCell, TableRow } from "@/components/ui/table";
interface AlertingSetting {
field_name: string;
@ -71,7 +74,13 @@ const DynamicForm: React.FC<DynamicFormProps> = ({
);
}
if (setting.field_type === "Boolean") {
return <Switch checked={setting.field_value} onChange={(checked) => handleToggle(setting, checked)} />;
return (
<Switch
aria-label={setting.field_name}
checked={setting.field_value}
onCheckedChange={(checked) => handleToggle(setting, checked)}
/>
);
}
return <Input value={setting.field_value ?? ""} onChange={(event) => handleTextChange(setting, event)} />;
};
@ -80,8 +89,8 @@ const DynamicForm: React.FC<DynamicFormProps> = ({
<form onSubmit={form.handleSubmit(onFinish)} noValidate>
{alertingSettings.map((value, index) => (
<TableRow key={index}>
<TableCell align="center">
<Text>{value.field_name}</Text>
<TableCell>
<p className="text-sm">{value.field_name}</p>
<p className="mt-1 text-[0.65rem] italic text-muted-foreground">{value.field_description}</p>
</TableCell>
{value.premium_field && !premiumUser ? (
@ -97,19 +106,27 @@ const DynamicForm: React.FC<DynamicFormProps> = ({
)}
<TableCell>
{value.stored_in_db == true ? (
<Badge icon={CheckCircleIcon} className="text-white">
<Badge variant="secondary">
<CircleCheck />
In DB
</Badge>
) : value.stored_in_db == false ? (
<Badge className="text-gray bg-background outline-solid">In Config</Badge>
<Badge variant="outline">In Config</Badge>
) : (
<Badge className="text-gray bg-background outline-solid">Not Set</Badge>
<Badge variant="outline">Not Set</Badge>
)}
</TableCell>
<TableCell>
<Icon icon={TrashIcon} color="red" onClick={() => handleResetField(value.field_name, index)}>
Reset
</Icon>
<Button
type="button"
variant="ghost"
size="icon-sm"
aria-label={`Reset ${value.field_name}`}
onClick={() => handleResetField(value.field_name, index)}
className="text-red-500"
>
<Trash2 className="size-5" />
</Button>
</TableCell>
</TableRow>
))}

View file

@ -1,13 +1,16 @@
import React, { useState, useEffect } from "react";
import { Text, Button, Callout } from "@tremor/react";
import { Modal, Spin, Select } from "antd";
import { CircleCheck, FileDown } from "lucide-react";
import { z } from "zod/v4";
import { getGlobalLitellmHeaderName } from "@/components/networking";
import { toast } from "@/lib/toast";
import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert";
import { PasswordInput } from "@/components/shared/PasswordInput";
import { FieldGroup } from "@/components/shared/form/field";
import { FormField } from "@/components/shared/form/FormField";
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
import { useZodForm } from "@/lib/forms/useZodForm";
const cloudZeroSettingsSchema = z.object({
@ -242,7 +245,7 @@ const CloudZeroExportModal: React.FC<CloudZeroExportModalProps> = ({ isOpen, onC
<div className="space-y-4">
{/* Export Type Selection */}
<div>
<Text className="font-medium mb-2 block">Export Destination</Text>
<p className="text-sm font-medium mb-2 block">Export Destination</p>
<Select value={exportType} onChange={setExportType} options={exportOptions} className="w-full" size="large" />
</div>
@ -256,27 +259,15 @@ const CloudZeroExportModal: React.FC<CloudZeroExportModalProps> = ({ isOpen, onC
) : (
<>
{existingSettings && (
<Callout
title="Existing CloudZero Configuration"
icon={() => (
<svg className="w-5 h-5" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M9 12l2 2 4-4m6 2a9 9 0 11-18 0 9 9 0 0118 0z"
/>
</svg>
)}
color="green"
className="mb-4"
>
<Text>
<Alert className="mb-4">
<CircleCheck />
<AlertTitle>Existing CloudZero Configuration</AlertTitle>
<AlertDescription>
API Key: {existingSettings.api_key_masked}
<br />
Connection ID: {existingSettings.connection_id}
</Text>
</Callout>
</AlertDescription>
</Alert>
)}
{!existingSettings && (
@ -303,25 +294,27 @@ const CloudZeroExportModal: React.FC<CloudZeroExportModalProps> = ({ isOpen, onC
{/* CSV Export Info */}
{exportType === "csv" && (
<Callout
title="CSV Export"
icon={() => (
<svg className="w-5 h-5" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M12 6v6m0 0v6m0-6h6m-6 0H6" />
</svg>
)}
color="blue"
>
<Text>Export your usage data as a CSV file for analysis in spreadsheet applications.</Text>
</Callout>
<Alert variant="info">
<FileDown />
<AlertTitle>CSV Export</AlertTitle>
<AlertDescription>
Export your usage data as a CSV file for analysis in spreadsheet applications.
</AlertDescription>
</Alert>
)}
{/* Action Buttons */}
<div className="flex justify-end space-x-2 pt-4">
<Button variant="secondary" onClick={handleModalClose}>
<Button type="button" variant="secondary" onClick={handleModalClose}>
Cancel
</Button>
<Button onClick={handleExport} loading={loading || exportLoading} disabled={loading || exportLoading}>
<Button
type="button"
onClick={handleExport}
disabled={loading || exportLoading}
aria-busy={loading || exportLoading}
>
{(loading || exportLoading) && <UiLoadingSpinner className="size-4" />}
{exportType === "cloudzero" ? "Export to CloudZero" : "Export CSV"}
</Button>
</div>

View file

@ -1,7 +1,6 @@
import React from "react";
import { Button, Modal, Typography } from "antd";
import { CopyToClipboard } from "react-copy-to-clipboard";
import { Text } from "@tremor/react";
import { toast } from "@/lib/toast";
export interface InvitationLink {
@ -90,14 +89,12 @@ export default function OnboardingModal({
: "Copy and send the generated link to the user to reset their password."}
</Paragraph>
<div className="flex justify-between pt-5 pb-2">
<Text className="text-base">User ID</Text>
<Text>{invitationLinkData?.user_id}</Text>
<p className="text-base">User ID</p>
<p className="text-sm">{invitationLinkData?.user_id}</p>
</div>
<div className="flex justify-between pt-5 pb-2">
<Text>{modalType === "invitation" ? "Invitation Link" : "Reset Password Link"}</Text>
<Text>
<Text>{getInvitationUrl()}</Text>
</Text>
<p className="text-sm">{modalType === "invitation" ? "Invitation Link" : "Reset Password Link"}</p>
<p className="text-sm">{getInvitationUrl()}</p>
</div>
<div className="flex justify-end mt-5">
<CopyToClipboard text={getInvitationUrl()} onCopy={() => toast.success("Copied!")}>

View file

@ -19,7 +19,8 @@ const config: ViteUserConfig = {
setupFiles: ["tests/setupTests.ts"],
globals: true,
css: true, // lets you import CSS/modules without extra mocks
testTimeout: 30000,
testTimeout: 60000,
hookTimeout: 30000,
silent: process.env.CI ? "passed-only" : false,
teardownTimeout: 60000,
coverage: {