mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/competent-lewin-1c8fd9
This commit is contained in:
commit
c71b6ed51b
44 changed files with 1304 additions and 274 deletions
|
|
@ -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).",
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
78
litellm/litellm_core_utils/llm_cost_calc/guardrail_cost.py
Normal file
78
litellm/litellm_core_utils/llm_cost_calc/guardrail_cost.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -33,7 +33,7 @@
|
|||
"limit": 2
|
||||
},
|
||||
"B006": {
|
||||
"limit": 178
|
||||
"limit": 177
|
||||
},
|
||||
"B008": {
|
||||
"limit": 503
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 22900
|
||||
"limit": 22897
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26889
|
||||
"limit": 26888
|
||||
},
|
||||
"LIT003": {
|
||||
"limit": 269
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
</>
|
||||
),
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
))}
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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!")}>
|
||||
|
|
|
|||
|
|
@ -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: {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue