Merge branch 'main' into litellm_release_day_02_10_2026

This commit is contained in:
Alexsander Hamir 2026-02-10 16:25:54 -08:00
commit 01a07903f6
80 changed files with 2090 additions and 1463 deletions

View file

@ -2277,6 +2277,7 @@ jobs:
- run: python ./tests/code_coverage_tests/router_code_coverage.py
- run: python ./tests/code_coverage_tests/test_chat_completion_imports.py
- run: python ./tests/code_coverage_tests/info_log_check.py
- run: python ./tests/code_coverage_tests/check_guardrail_apply_decorator.py
- run: python ./tests/code_coverage_tests/test_ban_set_verbose.py
- run: python ./tests/code_coverage_tests/code_qa_check_tests.py
- run: python ./tests/code_coverage_tests/check_get_model_cost_key_performance.py

View file

@ -0,0 +1,95 @@
---
slug: model-cost-map-incident
title: "Incident Report: Invalid model cost map on main"
date: 2026-02-10T10:00:00
authors:
- name: Ishaan Jaffer
title: "CTO, LiteLLM"
url: https://www.linkedin.com/in/ishaanjaffer/
image_url: https://pbs.twimg.com/profile_images/1613813310264340481/lz54oEiB_400x400.jpg
tags: [incident-report, stability]
hide_table_of_contents: false
---
**Date:** January 27, 2026
**Duration:** ~20 minutes
**Severity:** Low
**Status:** Resolved
## Summary
A malformed JSON entry in `model_prices_and_context_window.json` was merged to `main` ([`562f0a0`](https://github.com/BerriAI/litellm/commit/562f0a028251750e3d75386bee0e630d9796d0df)). This caused LiteLLM to silently fall back to a stale local copy of the model cost map. Users on older package versions lost cost tracking for newer models only (e.g. `azure/gpt-5.2`). No LLM calls were blocked.
- **LLM calls and proxy routing:** No impact.
- **Cost tracking:** Impacted for newer models not present in the local backup. Older models were unaffected. The incident lasted ~20 minutes until the commit was reverted.
{/* truncate */}
---
## Background
The model cost map is not in the request path. It is used after the LLM response comes back, inside a try/catch, to calculate spend. A missing entry never blocks a call.
```mermaid
flowchart TD
A["1. litellm.completion() receives request
litellm/main.py"] --> B["2. Route to provider
litellm/litellm_core_utils/get_llm_provider_logic.py"]
B --> C["3. LLM returns response
litellm/main.py"]
C --> D["4. Post-call: look up model in cost map
litellm/cost_calculator.py"]
D -->|"found"| E["5a. Attach cost to response"]
D -->|"not found (try/catch)"| F["5b. Log warning, set cost=0"]
E --> G["6. Return response to caller"]
F --> G
style D fill:#fff3cd,stroke:#ffc107
style F fill:#fff3cd,stroke:#ffc107
style E fill:#d4edda,stroke:#28a745
style G fill:#d4edda,stroke:#28a745
```
Both paths return a response to the caller. When the cost map lookup fails, the only difference is `cost=0` on that request.
---
## Root cause
LiteLLM fetches the model cost map from GitHub `main` at import time. If the fetch fails, it falls back to a local backup bundled with the package. Before this incident, the fallback was completely silent -- no warning was logged.
A contributor PR introduced an extra `{` bracket, producing invalid JSON. The remote fetch failed with `JSONDecodeError`, triggering the silent fallback. Users on older package versions had backup files missing newer models.
**Timeline:**
1. Malformed JSON merged to `main`
2. LiteLLM installations fall back to local backup on next import
3. Users report `"This model isn't mapped yet"` for newer models
4. Bad commit identified and reverted (~20 minutes)
---
## Remediation
| # | Action | Status | Code |
|---|---|---|---|
| 1 | CI validation on `model_prices_and_context_window.json` | ✅ Done | [`test-model-map.yaml`](https://github.com/BerriAI/litellm/blob/main/.github/workflows/test-model-map.yaml) |
| 2 | Warning log on fallback to local backup | ✅ Done | [`get_model_cost_map.py#L57-L68`](https://github.com/BerriAI/litellm/blob/main/litellm/litellm_core_utils/get_model_cost_map.py#L57-L68) |
| 3 | `GetModelCostMap` class with integrity validation helpers | ✅ Done | [`get_model_cost_map.py#L24-L149`](https://github.com/BerriAI/litellm/blob/main/litellm/litellm_core_utils/get_model_cost_map.py#L24-L149) |
| 4 | Resilience test suite (bad hosted map, fallback, completion) | ✅ Done | [`test_model_cost_map_resilience.py#L150-L291`](https://github.com/BerriAI/litellm/blob/main/tests/llm_translation/test_model_cost_map_resilience.py#L150-L291) |
| 5 | Test that backup model cost map always exists and contains common models | ✅ Done | [`test_model_cost_map_resilience.py#L213-L228`](https://github.com/BerriAI/litellm/blob/main/tests/llm_translation/test_model_cost_map_resilience.py#L213-L228) |
Enterprises that require zero external dependencies at import time can set `LITELLM_LOCAL_MODEL_COST_MAP=True` to skip the GitHub fetch entirely.
---
## Other dependencies on external resources
| Dependency | Impact if unavailable | Fallback |
|---|---|---|
| Model cost map (GitHub) | Cost tracking for newer models | Local backup (now with warning) |
| JWT public keys (IDP/SSO) | Auth fails | None |
| OIDC UserInfo (IDP/SSO) | Auth fails | None |
| HuggingFace model API | HF provider calls fail | None |
| Ollama tags (localhost) | Ollama model list stale | Static list |

View file

@ -1085,6 +1085,17 @@ const sidebars = {
"troubleshoot/max_callbacks",
],
},
{
type: "category",
label: "Blog",
items: [
{
type: "link",
label: "Incident: Broken Model Cost Map",
href: "/blog/model-cost-map-incident",
},
],
},
],
};

Binary file not shown.

View file

@ -48,6 +48,14 @@ DEFAULT_REPLICATE_POLLING_DELAY_SECONDS = int(
os.getenv("DEFAULT_REPLICATE_POLLING_DELAY_SECONDS", 1)
)
DEFAULT_IMAGE_TOKEN_COUNT = int(os.getenv("DEFAULT_IMAGE_TOKEN_COUNT", 250))
# Model cost map validation constants
MODEL_COST_MAP_MIN_MODEL_COUNT = int(
os.getenv("MODEL_COST_MAP_MIN_MODEL_COUNT", 50)
) # Minimum number of models a fetched cost map must contain to be considered valid
MODEL_COST_MAP_MAX_SHRINK_RATIO = float(
os.getenv("MODEL_COST_MAP_MAX_SHRINK_RATIO", 0.5)
) # Maximum allowed shrinkage ratio vs local backup (0.5 = reject if fetched map is <50% of backup)
DEFAULT_IMAGE_WIDTH = int(os.getenv("DEFAULT_IMAGE_WIDTH", 300))
DEFAULT_IMAGE_HEIGHT = int(os.getenv("DEFAULT_IMAGE_HEIGHT", 300))
# Maximum size for image URL downloads in MB (default 50MB, set to 0 to disable limit)

View file

@ -616,6 +616,7 @@ class CustomGuardrail(CustomLogger):
end_time: Optional[float] = None,
duration: Optional[float] = None,
event_type: Optional[GuardrailEventHooks] = None,
original_inputs: Optional[Dict] = None,
):
"""
Add StandardLoggingGuardrailInformation to the request data
@ -625,6 +626,17 @@ class CustomGuardrail(CustomLogger):
# Convert None to empty dict to satisfy type requirements
guardrail_response = {} if response is None else response
# For apply_guardrail functions in custom_code_guardrail scenario,
# simplify the logged response to "allow", "deny", or "mask"
if original_inputs is not None and isinstance(response, dict):
# Check if inputs were modified by comparing them
if self._inputs_were_modified(original_inputs, response):
guardrail_response = "mask"
else:
guardrail_response = "allow"
verbose_logger.debug(f"Guardrail response: {response}")
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response=guardrail_response,
request_data=request_data,
@ -650,8 +662,14 @@ class CustomGuardrail(CustomLogger):
This gets logged on downsteam Langfuse, DataDog, etc.
"""
# For custom_code_guardrail scenario, log as "deny" instead of full exception
# Check if this is from custom_code_guardrail by checking the class name
guardrail_response: Union[Exception, str] = e
if "CustomCodeGuardrail" in self.__class__.__name__:
guardrail_response = "deny"
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response=e,
guardrail_json_response=guardrail_response,
request_data=request_data,
guardrail_status="guardrail_failed_to_respond",
duration=duration,
@ -661,6 +679,25 @@ class CustomGuardrail(CustomLogger):
)
raise e
def _inputs_were_modified(self, original_inputs: Dict, response: Dict) -> bool:
"""
Compare original inputs with response to determine if content was modified.
Returns True if the inputs were modified (mask scenario), False otherwise (allow scenario).
"""
# Get all keys from both dictionaries
all_keys = set(original_inputs.keys()) | set(response.keys())
# Compare each key's value
for key in all_keys:
original_value = original_inputs.get(key)
response_value = response.get(key)
if original_value != response_value:
return True
# No modifications detected
return False
def mask_content_in_string(
self,
content_string: str,
@ -768,6 +805,12 @@ def log_guardrail_information(func):
self: CustomGuardrail = args[0]
request_data: dict = kwargs.get("data") or kwargs.get("request_data") or {}
event_type = _infer_event_type_from_function_name(func.__name__)
# Store original inputs for comparison (for apply_guardrail functions)
original_inputs = None
if func.__name__ == "apply_guardrail" and "inputs" in kwargs:
original_inputs = kwargs.get("inputs")
try:
response = await func(*args, **kwargs)
return self._process_response(
@ -777,6 +820,7 @@ def log_guardrail_information(func):
end_time=datetime.now().timestamp(),
duration=(datetime.now() - start_time).total_seconds(),
event_type=event_type,
original_inputs=original_inputs,
)
except Exception as e:
return self._process_error(
@ -794,6 +838,12 @@ def log_guardrail_information(func):
self: CustomGuardrail = args[0]
request_data: dict = kwargs.get("data") or kwargs.get("request_data") or {}
event_type = _infer_event_type_from_function_name(func.__name__)
# Store original inputs for comparison (for apply_guardrail functions)
original_inputs = None
if func.__name__ == "apply_guardrail" and "inputs" in kwargs:
original_inputs = kwargs.get("inputs")
try:
response = func(*args, **kwargs)
return self._process_response(
@ -801,6 +851,7 @@ def log_guardrail_information(func):
request_data=request_data,
duration=(datetime.now() - start_time).total_seconds(),
event_type=event_type,
original_inputs=original_inputs,
)
except Exception as e:
return self._process_error(

View file

@ -72,6 +72,13 @@ class OpenTelemetryConfig:
model_id: Optional[str] = None
def __post_init__(self) -> None:
# If endpoint is specified but exporter is still the default "console",
# automatically infer "otlp_http" to send traces to the endpoint.
# This fixes an issue where UI-configured OTEL settings would default
# to console output instead of sending traces to the configured endpoint.
if self.endpoint and isinstance(self.exporter, str) and self.exporter == "console":
self.exporter = "otlp_http"
if not self.service_name:
self.service_name = os.getenv("OTEL_SERVICE_NAME", "litellm")
if not self.deployment_environment:

View file

@ -8,40 +8,187 @@ export LITELLM_LOCAL_MODEL_COST_MAP=True
```
"""
import json
import os
from importlib.resources import files
import httpx
from litellm import verbose_logger
from litellm.constants import (
MODEL_COST_MAP_MAX_SHRINK_RATIO,
MODEL_COST_MAP_MIN_MODEL_COUNT,
)
class GetModelCostMap:
"""
Handles fetching, validating, and loading the model cost map.
Only the backup model *count* is cached (a single int). The full
backup dict is never held in memory — it is only parsed when it
needs to be *returned* as a fallback.
"""
_backup_model_count: int = -1 # -1 = not yet loaded
@staticmethod
def load_local_model_cost_map() -> dict:
"""Load the local backup model cost map bundled with the package."""
content = json.loads(
files("litellm")
.joinpath("model_prices_and_context_window_backup.json")
.read_text(encoding="utf-8")
)
return content
@classmethod
def _get_backup_model_count(cls) -> int:
"""Return the number of models in the local backup (cached int)."""
if cls._backup_model_count < 0:
backup = cls.load_local_model_cost_map()
cls._backup_model_count = len(backup)
return cls._backup_model_count
@staticmethod
def _check_is_valid_dict(fetched_map: dict) -> bool:
"""Check 1: fetched map is a non-empty dict."""
if not isinstance(fetched_map, dict):
verbose_logger.warning(
"LiteLLM: Fetched model cost map is not a dict (type=%s). "
"Falling back to local backup.",
type(fetched_map).__name__,
)
return False
if len(fetched_map) == 0:
verbose_logger.warning(
"LiteLLM: Fetched model cost map is empty. "
"Falling back to local backup.",
)
return False
return True
@classmethod
def _check_model_count_not_reduced(
cls,
fetched_map: dict,
backup_model_count: int,
min_model_count: int = MODEL_COST_MAP_MIN_MODEL_COUNT,
max_shrink_ratio: float = MODEL_COST_MAP_MAX_SHRINK_RATIO,
) -> bool:
"""Check 2: model count has not reduced significantly vs backup."""
fetched_count = len(fetched_map)
if fetched_count < min_model_count:
verbose_logger.warning(
"LiteLLM: Fetched model cost map has only %d models (minimum=%d). "
"This may indicate a corrupted upstream file. "
"Falling back to local backup.",
fetched_count,
min_model_count,
)
return False
if backup_model_count > 0 and fetched_count < backup_model_count * max_shrink_ratio:
verbose_logger.warning(
"LiteLLM: Fetched model cost map shrank significantly "
"(fetched=%d, backup=%d, threshold=%.0f%%). "
"This may indicate a corrupted upstream file. "
"Falling back to local backup.",
fetched_count,
backup_model_count,
max_shrink_ratio * 100,
)
return False
return True
@classmethod
def validate_model_cost_map(
cls,
fetched_map: dict,
backup_model_count: int,
min_model_count: int = MODEL_COST_MAP_MIN_MODEL_COUNT,
max_shrink_ratio: float = MODEL_COST_MAP_MAX_SHRINK_RATIO,
) -> bool:
"""
Validate the integrity of a fetched model cost map.
Runs each check in order and returns False on the first failure.
Checks:
1. ``_check_is_valid_dict`` -- fetched map is a non-empty dict.
2. ``_check_model_count_not_reduced`` -- model count meets minimum
and has not shrunk >``max_shrink_ratio`` vs backup.
Returns True if all checks pass, False otherwise.
"""
if not cls._check_is_valid_dict(fetched_map):
return False
if not cls._check_model_count_not_reduced(
fetched_map=fetched_map,
backup_model_count=backup_model_count,
min_model_count=min_model_count,
max_shrink_ratio=max_shrink_ratio,
):
return False
return True
@staticmethod
def fetch_remote_model_cost_map(url: str, timeout: int = 5) -> dict:
"""
Fetch the model cost map from a remote URL.
Returns the parsed JSON dict. Raises on network/parse errors
(caller is expected to handle).
"""
response = httpx.get(url, timeout=timeout)
response.raise_for_status()
return response.json()
def get_model_cost_map(url: str) -> dict:
if (
os.getenv("LITELLM_LOCAL_MODEL_COST_MAP", False)
or os.getenv("LITELLM_LOCAL_MODEL_COST_MAP", False) == "True"
):
from importlib.resources import files
import json
"""
Public entry point — returns the model cost map dict.
content = json.loads(
files("litellm")
.joinpath("model_prices_and_context_window_backup.json")
.read_text(encoding="utf-8")
)
return content
1. If ``LITELLM_LOCAL_MODEL_COST_MAP`` is set, uses the local backup only.
2. Otherwise fetches from ``url``, validates integrity, and falls back
to the local backup on any failure.
Only the backup model count is cached (a single int) for validation.
The full backup dict is only parsed when it must be *returned* as a
fallback — it is never held in memory long-term.
"""
# Note: can't use get_secret_bool here — this runs during litellm.__init__
# before litellm._key_management_settings is set.
if os.getenv("LITELLM_LOCAL_MODEL_COST_MAP", "").lower() == "true":
return GetModelCostMap.load_local_model_cost_map()
try:
response = httpx.get(
url, timeout=5
) # set a 5 second timeout for the get request
response.raise_for_status() # Raise an exception if the request is unsuccessful
content = response.json()
return content
except Exception:
from importlib.resources import files
import json
content = json.loads(
files("litellm")
.joinpath("model_prices_and_context_window_backup.json")
.read_text(encoding="utf-8")
content = GetModelCostMap.fetch_remote_model_cost_map(url)
except Exception as e:
verbose_logger.warning(
"LiteLLM: Failed to fetch remote model cost map from %s: %s. "
"Falling back to local backup.",
url,
str(e),
)
return content
return GetModelCostMap.load_local_model_cost_map()
# Validate using cached count (cheap int comparison, no file I/O)
if not GetModelCostMap.validate_model_cost_map(
fetched_map=content,
backup_model_count=GetModelCostMap._get_backup_model_count(),
):
verbose_logger.warning(
"LiteLLM: Fetched model cost map failed integrity check. "
"Using local backup instead. url=%s",
url,
)
return GetModelCostMap.load_local_model_cost_map()
return content

View file

@ -1,6 +1,7 @@
import asyncio
import contextlib
import os
import ssl
import typing
import urllib.request
from typing import Callable, Dict, Optional, Union
@ -139,8 +140,13 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
Credit to: https://github.com/karpetrosyan/httpx-aiohttp for this implementation
"""
def __init__(self, client: Union[ClientSession, Callable[[], ClientSession]]):
def __init__(
self,
client: Union[ClientSession, Callable[[], ClientSession]],
ssl_verify: Optional[Union[bool, ssl.SSLContext]] = None,
):
self.client = client
self._ssl_verify = ssl_verify # Store for per-request SSL override
super().__init__(client=client)
# Store the client factory for recreating sessions when needed
if callable(client):
@ -214,6 +220,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
timeout: dict,
proxy: Optional[str],
sni_hostname: Optional[str],
ssl_verify: Optional[Union[bool, ssl.SSLContext]] = None,
) -> ClientResponse:
"""
Helper function to make an aiohttp request with the given parameters.
@ -224,6 +231,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
timeout: Timeout settings dict with 'connect', 'read', 'pool' keys
proxy: Optional proxy URL
sni_hostname: Optional SNI hostname for SSL
ssl_verify: Optional SSL verification setting (False to disable, SSLContext for custom)
Returns:
ClientResponse from aiohttp
@ -237,6 +245,13 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
data = request.stream # type: ignore
request.headers.pop("transfer-encoding", None) # handled by aiohttp
# Only pass ssl kwarg when explicitly configured, to avoid
# overriding the session/connector defaults with None (which is
# not a valid value for aiohttp's ssl parameter).
ssl_kwargs: Dict[str, Union[bool, ssl.SSLContext]] = {}
if ssl_verify is not None:
ssl_kwargs["ssl"] = ssl_verify
response = await client_session.request(
method=request.method,
url=YarlURL(str(request.url), encoded=True),
@ -251,6 +266,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
),
proxy=proxy,
server_hostname=sni_hostname,
**ssl_kwargs,
).__aenter__()
return response
@ -268,6 +284,9 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
# Resolve proxy settings from environment variables
proxy = await self._get_proxy_settings(request)
# Use stored SSL configuration for per-request override
ssl_config = self._ssl_verify
try:
with map_aiohttp_exceptions():
response = await self._make_aiohttp_request(
@ -276,6 +295,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
timeout=timeout,
proxy=proxy,
sni_hostname=sni_hostname,
ssl_verify=ssl_config,
)
except RuntimeError as e:
# Handle the case where session was closed between our check and actual use
@ -296,6 +316,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
timeout=timeout,
proxy=proxy,
sni_hostname=sni_hostname,
ssl_verify=ssl_config,
)
else:
# Re-raise if it's a different RuntimeError

View file

@ -846,6 +846,16 @@ class AsyncHTTPHandler:
if str_to_bool(os.getenv("AIOHTTP_TRUST_ENV", "False")) is True:
trust_env = True
#########################################################
# Determine SSL config to pass to transport for per-request override
# This ensures ssl_verify works even with shared sessions
#########################################################
ssl_for_transport: Optional[Union[bool, ssl.SSLContext]] = None
if ssl_context is not None:
ssl_for_transport = ssl_context
elif ssl_verify is False:
ssl_for_transport = False
verbose_logger.debug("Creating AiohttpTransport...")
# Use shared session if provided and valid
@ -853,7 +863,10 @@ class AsyncHTTPHandler:
verbose_logger.debug(
f"SHARED SESSION: Reusing existing ClientSession (ID: {id(shared_session)})"
)
return LiteLLMAiohttpTransport(client=shared_session)
return LiteLLMAiohttpTransport(
client=shared_session,
ssl_verify=ssl_for_transport,
)
# Create new session only if none provided or existing one is invalid
verbose_logger.debug(
@ -877,6 +890,7 @@ class AsyncHTTPHandler:
connector=TCPConnector(**transport_connector_kwargs),
trust_env=trust_env,
),
ssl_verify=ssl_for_transport,
)
@staticmethod

View file

@ -5848,6 +5848,19 @@
"output_cost_per_token": 7e-07,
"supports_tool_choice": true
},
"azure_ai/kimi-k2.5": {
"input_cost_per_token": 6e-07,
"litellm_provider": "azure_ai",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 3e-06,
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/kimi-k2-5-now-in-microsoft-foundry/4492321",
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_vision": true
},
"azure_ai/ministral-3b": {
"input_cost_per_token": 4e-08,
"litellm_provider": "azure_ai",

View file

@ -795,9 +795,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
#########################################################
########## 1. Make the Bedrock API request ##########
#########################################################
bedrock_guardrail_response: Optional[
Union[BedrockGuardrailResponse, str]
] = None
bedrock_guardrail_response: Optional[Union[BedrockGuardrailResponse, str]] = (
None
)
try:
bedrock_guardrail_response = await self.make_bedrock_api_request(
source="INPUT", messages=filtered_messages, request_data=data
@ -867,9 +867,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
#########################################################
########## 1. Make the Bedrock API request ##########
#########################################################
bedrock_guardrail_response: Optional[
Union[BedrockGuardrailResponse, str]
] = None
bedrock_guardrail_response: Optional[Union[BedrockGuardrailResponse, str]] = (
None
)
try:
bedrock_guardrail_response = await self.make_bedrock_api_request(
source="INPUT", messages=filtered_messages, request_data=data

View file

@ -35,7 +35,10 @@ from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Type, cast
from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
from litellm.types.utils import GenericGuardrailAPIInputs
@ -179,6 +182,7 @@ class CustomCodeGuardrail(CustomGuardrail):
self._compile_error = f"Failed to compile custom code: {e}"
raise CustomCodeCompilationError(self._compile_error) from e
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,

View file

@ -23,7 +23,10 @@ import httpx
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
@ -483,6 +486,7 @@ class EnkryptAIGuardrails(CustomGuardrail):
request_data=data, guardrail_name=self.guardrail_name
)
@log_guardrail_information
async def apply_guardrail(
self,
inputs: "GenericGuardrailAPIInputs",

View file

@ -10,7 +10,10 @@ from typing import TYPE_CHECKING, Any, Dict, Literal, Optional
from litellm._logging import verbose_proxy_logger
from litellm.exceptions import GuardrailRaisedException
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
@ -150,6 +153,7 @@ class GenericGuardrailAPI(CustomGuardrail):
return result_metadata
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,

View file

@ -9,7 +9,8 @@ from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
ModifyResponseException
ModifyResponseException,
log_guardrail_information,
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
@ -108,7 +109,9 @@ class GraySwanGuardrail(CustomGuardrail):
self.categories = categories
self.policy_id = policy_id
self.fail_open = True if fail_open is None else bool(fail_open)
self.guardrail_timeout = 30.0 if guardrail_timeout is None else float(guardrail_timeout)
self.guardrail_timeout = (
30.0 if guardrail_timeout is None else float(guardrail_timeout)
)
# Streaming configuration
self.streaming_end_of_stream_only = streaming_end_of_stream_only
@ -155,6 +158,7 @@ class GraySwanGuardrail(CustomGuardrail):
# Unified Guardrail Interface (works with ALL endpoints automatically)
# ------------------------------------------------------------------
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
@ -208,7 +212,9 @@ class GraySwanGuardrail(CustomGuardrail):
messages = [{"role": role, "content": text} for text in texts]
# Get dynamic params from request metadata
dynamic_body = self.get_guardrail_dynamic_request_body_params(request_data) or {}
dynamic_body = (
self.get_guardrail_dynamic_request_body_params(request_data) or {}
)
if dynamic_body:
verbose_proxy_logger.debug(
"Gray Swan Guardrail: dynamic extra_body=%s", safe_dumps(dynamic_body)
@ -271,12 +277,12 @@ class GraySwanGuardrail(CustomGuardrail):
async def run_grayswan_guardrail(self, payload: dict) -> Dict[str, Any]:
"""
Run the GraySwan guardrail on a payload.
This is a legacy method for testing purposes.
Args:
payload: The payload to scan
Returns:
Dict containing the GraySwan API response
"""
@ -293,11 +299,11 @@ class GraySwanGuardrail(CustomGuardrail):
) -> None:
"""
Legacy method for processing GraySwan API responses.
This method is maintained for backward compatibility with existing tests.
It handles the test scenarios where responses need to be processed with
knowledge of the request context (pre/during/post call hooks).
Args:
response_json: Response from GraySwan API
data: Optional request data (for passthrough exceptions)
@ -365,7 +371,10 @@ class GraySwanGuardrail(CustomGuardrail):
)
# If hook_type is provided and in pre/during call, raise exception
if hook_type in [GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call]:
if hook_type in [
GuardrailEventHooks.pre_call,
GuardrailEventHooks.during_call,
]:
# Raise ModifyResponseException to short-circuit LLM call
if data is None:
data = {}
@ -540,7 +549,9 @@ class GraySwanGuardrail(CustomGuardrail):
if isinstance(litellm_metadata, dict) and litellm_metadata:
cleaned_litellm_metadata = dict(litellm_metadata)
# cleaned_litellm_metadata.pop("user_api_key_auth", None)
sanitized = safe_json_loads(safe_dumps(cleaned_litellm_metadata), default={})
sanitized = safe_json_loads(
safe_dumps(cleaned_litellm_metadata), default={}
)
if isinstance(sanitized, dict) and sanitized:
payload["litellm_metadata"] = sanitized
@ -566,7 +577,9 @@ class GraySwanGuardrail(CustomGuardrail):
detection_info = detection_info[0]
# Extract fields from detection_info dict
detection_dict: dict = detection_info if isinstance(detection_info, dict) else {}
detection_dict: dict = (
detection_info if isinstance(detection_info, dict) else {}
)
violation_score = detection_dict.get("violation_score", 0.0)
violated_rules = detection_dict.get("violated_rules", [])
mutation = detection_dict.get("mutation", False)
@ -582,7 +595,9 @@ class GraySwanGuardrail(CustomGuardrail):
if violated_rules:
formatted_rules = self._format_violated_rules(violated_rules)
if formatted_rules:
message_parts.append(f"It was violating the rule(s): {formatted_rules}.")
message_parts.append(
f"It was violating the rule(s): {formatted_rules}."
)
if mutation:
message_parts.append(
@ -590,9 +605,7 @@ class GraySwanGuardrail(CustomGuardrail):
)
if ipi:
message_parts.append(
"Indirect Prompt Injection was DETECTED."
)
message_parts.append("Indirect Prompt Injection was DETECTED.")
return "\n".join(message_parts)

View file

@ -10,7 +10,10 @@ from httpx import HTTPStatusError
from requests.auth import HTTPBasicAuth
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
@ -110,6 +113,7 @@ class HiddenlayerGuardrail(CustomGuardrail):
)
super().__init__(**kwargs)
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,

View file

@ -28,7 +28,10 @@ from fastapi import HTTPException
from litellm import Router
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.utils import ModelResponseStream
@ -50,6 +53,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter impor
ContentFilterDetection,
PatternDetection,
)
from .patterns import PATTERN_EXTRA_CONFIG, get_compiled_pattern
MAX_KEYWORD_VALUE_GAP_WORDS = 1
@ -168,9 +172,9 @@ class ContentFilterGuardrail(CustomGuardrail):
self.image_model = image_model
# Store loaded categories
self.loaded_categories: Dict[str, CategoryConfig] = {}
self.category_keywords: Dict[
str, Tuple[str, str, ContentFilterAction]
] = {} # keyword -> (category, severity, action)
self.category_keywords: Dict[str, Tuple[str, str, ContentFilterAction]] = (
{}
) # keyword -> (category, severity, action)
# Load categories if provided
if categories:
@ -994,6 +998,7 @@ class ContentFilterGuardrail(CustomGuardrail):
masked_entity_count=masked_entity_count,
)
@log_guardrail_information
async def apply_guardrail(
self,
inputs: "GenericGuardrailAPIInputs",

View file

@ -12,7 +12,10 @@ import httpx
from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
@ -26,7 +29,11 @@ if TYPE_CHECKING:
class OnyxGuardrail(CustomGuardrail):
def __init__(
self, api_base: Optional[str] = None, api_key: Optional[str] = None, timeout: Optional[float] = 10.0, **kwargs
self,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
timeout: Optional[float] = 10.0,
**kwargs,
):
timeout = timeout or int(os.getenv("ONYX_TIMEOUT", 10.0))
self.async_handler = get_async_httpx_client(
@ -79,6 +86,7 @@ class OnyxGuardrail(CustomGuardrail):
)
return result
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,

View file

@ -58,7 +58,9 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
guardrail_name: str,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
model: Optional[Literal["omni-moderation-latest", "text-moderation-latest"]] = None,
model: Optional[
Literal["omni-moderation-latest", "text-moderation-latest"]
] = None,
**kwargs,
):
"""Initialize OpenAI Moderation guardrail handler."""
@ -75,7 +77,7 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
supported_event_hooks=supported_event_hooks,
**kwargs,
)
self.async_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback
)
@ -83,10 +85,14 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
# Store configuration
self.api_key = api_key or self._get_api_key()
self.api_base = api_base or "https://api.openai.com/v1"
self.model: Literal["omni-moderation-latest", "text-moderation-latest"] = model or "omni-moderation-latest"
self.model: Literal["omni-moderation-latest", "text-moderation-latest"] = (
model or "omni-moderation-latest"
)
if not self.api_key:
raise ValueError("OpenAI Moderation: api_key is required. Set OPENAI_API_KEY environment variable or pass it in configuration.")
raise ValueError(
"OpenAI Moderation: api_key is required. Set OPENAI_API_KEY environment variable or pass it in configuration."
)
verbose_proxy_logger.debug(
f"Initialized OpenAI Moderation Guardrail: {guardrail_name} with model: {self.model}"
@ -98,7 +104,7 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
import litellm
from litellm.secret_managers.main import get_secret_str
return (
os.environ.get("OPENAI_API_KEY")
or litellm.api_key
@ -106,21 +112,14 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
or get_secret_str("OPENAI_API_KEY")
)
async def async_make_request(
self, input_text: str
) -> "OpenAIModerationResponse":
async def async_make_request(self, input_text: str) -> "OpenAIModerationResponse":
"""
Make a request to the OpenAI Moderation API.
"""
request_body = {
"model": self.model,
"input": input_text
}
verbose_proxy_logger.debug(
"OpenAI Moderation guard request: %s", request_body
)
request_body = {"model": self.model, "input": input_text}
verbose_proxy_logger.debug("OpenAI Moderation guard request: %s", request_body)
response = await self.async_handler.post(
url=f"{self.api_base}/moderations",
headers={
@ -133,7 +132,7 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
verbose_proxy_logger.debug(
"OpenAI Moderation guard response: %s", response.json()
)
if response.status_code != 200:
raise HTTPException(
status_code=response.status_code,
@ -144,9 +143,12 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
)
from litellm.types.llms.openai import OpenAIModerationResponse
return OpenAIModerationResponse(**response.json())
def _check_moderation_result(self, moderation_response: "OpenAIModerationResponse") -> None:
def _check_moderation_result(
self, moderation_response: "OpenAIModerationResponse"
) -> None:
"""
Check if the moderation response indicates harmful content and raise exception if needed.
"""
@ -168,10 +170,10 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
}
verbose_proxy_logger.warning(
"OpenAI Moderation: Content flagged for violations: %s",
violation_details
"OpenAI Moderation: Content flagged for violations: %s",
violation_details,
)
raise HTTPException(
status_code=400,
detail={
@ -180,6 +182,7 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
},
)
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
@ -189,51 +192,50 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
) -> GenericGuardrailAPIInputs:
"""
Apply OpenAI moderation guardrail using the unified guardrail interface.
This method is called by the UnifiedLLMGuardrails system for all endpoint types
(chat completions, embeddings, responses API, etc.).
Args:
inputs: GenericGuardrailAPIInputs containing texts and/or structured_messages
request_data: The original request data
input_type: Whether this is a "request" (pre-call) or "response" (post-call)
logging_obj: Optional logging object
Returns:
The inputs unchanged (moderation doesn't modify content, only blocks)
Raises:
HTTPException: If content violates moderation policy
"""
# Extract text to moderate from inputs
text_to_moderate: Optional[str] = None
# Prefer structured_messages if available (has role context)
if structured_messages := inputs.get("structured_messages"):
text_to_moderate = self.get_user_prompt(structured_messages)
# Fall back to texts
if not text_to_moderate:
if texts := inputs.get("texts"):
# Join all texts for moderation
text_to_moderate = "\n".join(texts)
if not text_to_moderate:
verbose_proxy_logger.debug(
"OpenAI Moderation: No text content to moderate in inputs"
)
return inputs
# Make moderation request
moderation_response = await self.async_make_request(input_text=text_to_moderate)
# Check if content is flagged and raise exception if needed
self._check_moderation_result(moderation_response)
# Moderation doesn't modify content, just blocks - return inputs unchanged
return inputs
@log_guardrail_information
async def async_post_call_streaming_iterator_hook(
self,
@ -252,9 +254,7 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
from litellm.main import stream_chunk_builder
from litellm.types.utils import TextCompletionResponse
verbose_proxy_logger.debug(
"OpenAI Moderation: Running streaming response scan"
)
verbose_proxy_logger.debug("OpenAI Moderation: Running streaming response scan")
# Collect all chunks to process them together
all_chunks: List["ModelResponseStream"] = []
@ -269,7 +269,7 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
)
if isinstance(assembled_model_response, (type(None), TextCompletionResponse)):
# If we can't assemble a ModelResponse or it's a text completion,
# If we can't assemble a ModelResponse or it's a text completion,
# just yield the original chunks without moderation
verbose_proxy_logger.warning(
"OpenAI Moderation: Could not assemble ModelResponse from chunks, skipping moderation"
@ -284,19 +284,17 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
verbose_proxy_logger.debug(
f"OpenAI Moderation: Streaming response text: {response_text[:100]}..." # Log first 100 chars
)
# Make moderation request - this will raise HTTPException if content is flagged
moderation_response = await self.async_make_request(
input_text=response_text,
)
# Check if content is flagged and raise exception if needed
self._check_moderation_result(moderation_response)
# If we reach here, content passed moderation - yield the original chunks
mock_response = MockResponseIterator(
model_response=assembled_model_response
)
mock_response = MockResponseIterator(model_response=assembled_model_response)
# Return the reconstructed stream
async for chunk in mock_response:
@ -306,34 +304,34 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
"""
Extract text content from the model response for moderation.
"""
if not hasattr(response, 'choices') or not response.choices:
if not hasattr(response, "choices") or not response.choices:
return None
response_texts = []
for choice in response.choices:
try:
# Try to get content from message (chat completion)
message = getattr(choice, 'message', None)
message = getattr(choice, "message", None)
if message:
content = getattr(message, 'content', None)
content = getattr(message, "content", None)
if content and isinstance(content, str):
response_texts.append(content)
continue
# Try to get text (text completion)
text = getattr(choice, 'text', None)
text = getattr(choice, "text", None)
if text and isinstance(text, str):
response_texts.append(text)
continue
# Try to get content from delta (streaming)
delta = getattr(choice, 'delta', None)
delta = getattr(choice, "delta", None)
if delta:
content = getattr(delta, 'content', None)
content = getattr(delta, "content", None)
if content and isinstance(content, str):
response_texts.append(content)
continue
except (AttributeError, TypeError):
# Skip choices that don't have expected attributes
continue

View file

@ -9,10 +9,10 @@
import asyncio
import threading
import json
from datetime import datetime
import threading
from contextlib import asynccontextmanager
from datetime import datetime
from typing import (
TYPE_CHECKING,
Any,
@ -39,7 +39,10 @@ if TYPE_CHECKING:
from litellm._uuid import uuid
from litellm.caching.caching import DualCache
from litellm.exceptions import BlockedPiiEntityError
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import (
GuardrailEventHooks,
@ -568,9 +571,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
if messages is None:
return data
tasks = []
task_mappings: List[
Tuple[int, Optional[int]]
] = [] # Track (message_index, content_index) for each task
task_mappings: List[Tuple[int, Optional[int]]] = (
[]
) # Track (message_index, content_index) for each task
for msg_idx, m in enumerate(messages):
content = m.get("content", None)
@ -671,9 +674,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
): # /chat/completions requests
messages: Optional[List] = kwargs.get("messages", None)
tasks = []
task_mappings: List[
Tuple[int, Optional[int]]
] = [] # Track (message_index, content_index) for each task
task_mappings: List[Tuple[int, Optional[int]]] = (
[]
) # Track (message_index, content_index) for each task
if messages is None:
return kwargs, result
@ -792,11 +795,11 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
# Type narrowing: StreamingChoices doesn't have .message attribute
if not hasattr(choice, "message"):
continue
content = getattr(choice.message, "content", None)
content = getattr(choice.message, "content", None) # type: ignore
if content is None:
continue
if isinstance(content, str):
choice.message.content = await self.check_pii(
choice.message.content = await self.check_pii( # type: ignore
text=content,
output_parse_pii=False,
presidio_config=presidio_config,
@ -989,6 +992,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
except Exception:
pass
@log_guardrail_information
async def apply_guardrail(
self,
inputs: "GenericGuardrailAPIInputs",

View file

@ -6,7 +6,10 @@ from typing import TYPE_CHECKING, Any, List, Literal, Optional, Type
from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
@ -67,6 +70,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
super().__init__(**kwargs)
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,

View file

@ -12,10 +12,11 @@ from typing import Any, Dict, List, Literal, Optional, Type
from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.litellm_logging import (
Logging as LiteLLMLoggingObj,
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
@ -343,9 +344,7 @@ class QualifireGuardrail(CustomGuardrail):
)
url = f"{self.qualifire_api_base}/api/evaluation/evaluate"
verbose_proxy_logger.debug(
f"Qualifire Guardrail: Making request to {url}"
)
verbose_proxy_logger.debug(f"Qualifire Guardrail: Making request to {url}")
# Make the API request
response = await self.async_handler.post(
@ -393,6 +392,7 @@ class QualifireGuardrail(CustomGuardrail):
verbose_proxy_logger.exception(f"Qualifire Guardrail error: {e}")
raise
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,

View file

@ -9,7 +9,10 @@ from typing import TYPE_CHECKING, Literal, Optional
from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
@ -70,6 +73,7 @@ class ZscalerAIGuard(CustomGuardrail):
return str(value).strip()
return "N/A"
@log_guardrail_information
async def apply_guardrail(
self,
inputs: "GenericGuardrailAPIInputs",
@ -92,7 +96,7 @@ class ZscalerAIGuard(CustomGuardrail):
Raises:
Exception: If content is blocked by Zscaler AI Guard
"""
texts = inputs.get("texts", [])
try:
verbose_proxy_logger.debug(f"ZscalerAIGuard: Checking {len(texts)} text(s)")
@ -102,8 +106,8 @@ class ZscalerAIGuard(CustomGuardrail):
team_metadata = metadata.get("team_metadata", {}) or {}
# Precedence for policy_id:
# 1. metadata.zguard_policy_id # request level
# 2. user_api_key_metadata.zguard_policy_id # Key level
# 1. metadata.zguard_policy_id # request level
# 2. user_api_key_metadata.zguard_policy_id # Key level
# 3. team_metadata.zguard_policy_id # Team level
# 4. self.policy_id (from environment) # Global
policy_id = (
@ -154,9 +158,7 @@ class ZscalerAIGuard(CustomGuardrail):
zscaler_ai_guard_result
and zscaler_ai_guard_result.get("action") == "BLOCK"
):
blocking_info = zscaler_ai_guard_result.get(
"zscaler_ai_guard_response"
)
blocking_info = zscaler_ai_guard_result.get("zscaler_ai_guard_response")
error_message = f"Content blocked by Zscaler AI Guard: {self.extract_blocking_info(blocking_info)}"
raise Exception(error_message)
except Exception as e:

View file

@ -1672,6 +1672,9 @@ async def ui_view_spend_logs( # noqa: PLR0915
model: Optional[str] = fastapi.Query(
default=None, description="Filter logs by model"
),
model_id: Optional[str] = fastapi.Query(
default=None, description="Filter logs by model ID (litellm model deployment id)"
),
key_alias: Optional[str] = fastapi.Query(
default=None, description="Filter logs by key alias"
),
@ -1763,6 +1766,9 @@ async def ui_view_spend_logs( # noqa: PLR0915
if model is not None:
where_conditions["model"] = model
if model_id is not None:
where_conditions["model_id"] = model_id
# Build metadata filters
metadata_filters = []
if key_alias is not None:

View file

@ -1924,6 +1924,17 @@ class Router:
"deployment_model_name": deployment_model_name,
}
)
## DEPLOYMENT-LEVEL TAGS
deployment_tags = deployment.get("litellm_params", {}).get("tags")
if deployment_tags:
existing_tags = kwargs[metadata_variable_name].get("tags") or []
merged_tags = list(existing_tags)
for tag in deployment_tags:
if tag not in merged_tags:
merged_tags.append(tag)
kwargs[metadata_variable_name]["tags"] = merged_tags
kwargs["model_info"] = model_info
kwargs["timeout"] = self._get_timeout(

View file

@ -5848,6 +5848,19 @@
"output_cost_per_token": 7e-07,
"supports_tool_choice": true
},
"azure_ai/kimi-k2.5": {
"input_cost_per_token": 6e-07,
"litellm_provider": "azure_ai",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 3e-06,
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/kimi-k2-5-now-in-microsoft-foundry/4492321",
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_vision": true
},
"azure_ai/ministral-3b": {
"input_cost_per_token": 4e-08,
"litellm_provider": "azure_ai",

View file

@ -0,0 +1,126 @@
"""
Test that all guardrail hooks with async def apply_guardrail use @log_guardrail_information decorator.
This ensures consistent logging and observability across all guardrail implementations.
"""
import ast
from pathlib import Path
from typing import List, Tuple
def find_apply_guardrail_methods(file_path: Path) -> List[Tuple[str, int, bool]]:
"""
Find all apply_guardrail methods and check if they have the decorator.
Returns:
List of tuples: (class_name, line_number, has_decorator)
"""
with open(file_path, "r") as f:
content = f.read()
try:
tree = ast.parse(content)
except SyntaxError:
return []
results = []
for node in ast.walk(tree):
if isinstance(node, ast.ClassDef):
class_name = node.name
# Check if this class has apply_guardrail method
for item in node.body:
if (
isinstance(item, ast.AsyncFunctionDef)
and item.name == "apply_guardrail"
):
# Check if it has the log_guardrail_information decorator
has_decorator = False
for decorator in item.decorator_list:
if (
isinstance(decorator, ast.Name)
and decorator.id == "log_guardrail_information"
):
has_decorator = True
break
results.append((class_name, item.lineno, has_decorator))
return results
def test_guardrail_apply_decorator():
"""Test that all guardrail hooks with apply_guardrail have the decorator."""
# Path to the guardrail hooks directory
guardrail_hooks_dir = (
Path(__file__).parent.parent.parent
/ "litellm"
/ "proxy"
/ "guardrails"
/ "guardrail_hooks"
)
# Find all Python files in the guardrail hooks directory
python_files = list(guardrail_hooks_dir.rglob("*.py"))
# Track violations
violations = []
for python_file in python_files:
# Skip __init__.py files and test files
if python_file.name == "__init__.py" or python_file.name.startswith("test_"):
continue
# Skip base files and primitives
if python_file.name in ["base.py", "primitives.py", "patterns.py"]:
continue
# Skip bedrock_guardrails.py - it implements logging differently via
# add_standard_logging_guardrail_information_to_request_data calls
# in make_bedrock_api_request method instead of using the decorator
if python_file.name == "bedrock_guardrails.py":
continue
results = find_apply_guardrail_methods(python_file)
for class_name, line_num, has_decorator in results:
if not has_decorator:
relative_path = python_file.relative_to(
Path(__file__).parent.parent.parent
)
violations.append((relative_path, class_name, line_num))
# Assert no violations found
if violations:
print(
f"\nFound {len(violations)} guardrail hook(s) without @log_guardrail_information decorator:"
)
print(
"\nAll guardrail hooks must use @log_guardrail_information decorator on their apply_guardrail method."
)
print(
"This ensures consistent logging and observability across all guardrails.\n"
)
for file_path, class_name, line_num in violations:
print(f" - {file_path}:{line_num} ({class_name}.apply_guardrail)")
print("\nTo fix, add the decorator:")
print(
" from litellm.integrations.custom_guardrail import log_guardrail_information"
)
print(" ")
print(" @log_guardrail_information")
print(" async def apply_guardrail(self, ...):")
print(" ...")
raise AssertionError(
f"Found {len(violations)} guardrail hook(s) without @log_guardrail_information decorator"
)
if __name__ == "__main__":
test_guardrail_apply_decorator()
print("✓ All guardrail hooks have @log_guardrail_information decorator")

View file

@ -0,0 +1,291 @@
"""
Tests for model cost map resilience.
Simulates:
- A bad (invalid JSON) model cost map upstream
- A bad (empty/missing) backup model cost map
- Verifies litellm.completion() still works even with a broken cost map
- Verifies litellm.get_model_info() raises the expected error for unmapped models
- Verifies the integrity validation helper catches corrupted maps
"""
import json
import os
import sys
from unittest.mock import MagicMock, patch
import pytest
sys.path.insert(
0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../.."))
)
import litellm
from litellm.litellm_core_utils.get_model_cost_map import (
GetModelCostMap,
get_model_cost_map,
)
class TestCheckIsValidDict:
"""Unit tests for _check_is_valid_dict."""
def test_should_reject_non_dict(self):
"""Non-dict should fail."""
assert GetModelCostMap._check_is_valid_dict("not a dict") is False
def test_should_reject_empty_dict(self):
"""Empty dict should fail."""
assert GetModelCostMap._check_is_valid_dict({}) is False
def test_should_reject_list(self):
"""List should fail."""
assert GetModelCostMap._check_is_valid_dict([1, 2, 3]) is False
def test_should_reject_none(self):
"""None should fail."""
assert GetModelCostMap._check_is_valid_dict(None) is False
def test_should_accept_non_empty_dict(self):
"""Non-empty dict should pass."""
assert GetModelCostMap._check_is_valid_dict({"model": {}}) is True
class TestCheckModelCountNotReduced:
"""Unit tests for _check_model_count_not_reduced."""
def test_should_reject_too_few_models(self):
"""Fetched map with fewer models than min_model_count should fail."""
small_map = {f"model-{i}": {} for i in range(5)}
assert (
GetModelCostMap._check_model_count_not_reduced(
fetched_map=small_map, backup_model_count=0, min_model_count=10
)
is False
)
def test_should_reject_significant_shrinkage(self):
"""Fetched map that shrunk >50% vs backup should fail."""
fetched = {f"model-{i}": {} for i in range(40)} # 40% of 100
assert (
GetModelCostMap._check_model_count_not_reduced(
fetched_map=fetched, backup_model_count=100, min_model_count=10
)
is False
)
def test_should_accept_when_above_threshold(self):
"""Fetched map at 60% of backup (above 50% threshold) should pass."""
fetched = {f"model-{i}": {} for i in range(60)}
assert (
GetModelCostMap._check_model_count_not_reduced(
fetched_map=fetched, backup_model_count=100, min_model_count=10
)
is True
)
def test_should_accept_growth(self):
"""Fetched map larger than backup should pass."""
fetched = {f"model-{i}": {} for i in range(120)}
assert (
GetModelCostMap._check_model_count_not_reduced(
fetched_map=fetched, backup_model_count=100, min_model_count=10
)
is True
)
def test_should_accept_with_empty_backup(self):
"""When backup is empty, only min_model_count matters."""
fetched = {f"model-{i}": {} for i in range(15)}
assert (
GetModelCostMap._check_model_count_not_reduced(
fetched_map=fetched, backup_model_count=0, min_model_count=10
)
is True
)
class TestValidateModelCostMap:
"""Unit tests for validate_model_cost_map (combines both checks)."""
def test_should_reject_non_dict(self):
"""Non-dict should fail at check 1."""
assert GetModelCostMap.validate_model_cost_map(fetched_map="not a dict", backup_model_count=0) is False
def test_should_reject_empty_map(self):
"""Empty dict should fail at check 1."""
assert GetModelCostMap.validate_model_cost_map(fetched_map={}, backup_model_count=0) is False
def test_should_reject_significant_shrinkage(self):
"""Should fail at check 2 (shrinkage)."""
fetched = {f"model-{i}": {} for i in range(40)}
assert (
GetModelCostMap.validate_model_cost_map(
fetched_map=fetched, backup_model_count=100, min_model_count=10
)
is False
)
def test_should_accept_valid_map(self):
"""Should pass both checks."""
fetched = {f"model-{i}": {} for i in range(120)}
assert (
GetModelCostMap.validate_model_cost_map(
fetched_map=fetched, backup_model_count=100, min_model_count=10
)
is True
)
def test_should_accept_equal_size_map(self):
"""Equal size should pass both checks."""
fetched = {f"model-{i}": {} for i in range(100)}
assert (
GetModelCostMap.validate_model_cost_map(
fetched_map=fetched, backup_model_count=100, min_model_count=10
)
is True
)
class TestGetModelCostMapFallback:
"""Tests for get_model_cost_map fallback behavior with bad upstream."""
def test_should_fallback_to_backup_on_invalid_json(self):
"""When upstream returns invalid JSON, should fall back to local backup."""
mock_response = MagicMock()
mock_response.raise_for_status = MagicMock()
mock_response.json.side_effect = json.JSONDecodeError("bad json", "", 0)
with patch("httpx.get", return_value=mock_response):
result = get_model_cost_map("https://fake-url.com/model_prices.json")
# Should have fallen back to backup — backup always has models
assert isinstance(result, dict)
assert len(result) > 0
def test_should_fallback_to_backup_on_network_error(self):
"""When upstream is unreachable, should fall back to local backup."""
with patch("httpx.get", side_effect=Exception("Connection refused")):
result = get_model_cost_map("https://fake-url.com/model_prices.json")
assert isinstance(result, dict)
assert len(result) > 0
def test_should_fallback_when_fetched_map_is_empty(self):
"""When upstream returns valid JSON but empty dict, should fall back."""
mock_response = MagicMock()
mock_response.raise_for_status = MagicMock()
mock_response.json.return_value = {} # empty map
with patch("httpx.get", return_value=mock_response):
result = get_model_cost_map("https://fake-url.com/model_prices.json")
# Should have fallen back to backup since empty map fails validation
assert isinstance(result, dict)
assert len(result) > 0
def test_should_fallback_when_fetched_map_shrinks_dramatically(self):
"""When upstream returns far fewer models than backup, should fall back."""
tiny_map = {f"model-{i}": {"litellm_provider": "test"} for i in range(11)}
mock_response = MagicMock()
mock_response.raise_for_status = MagicMock()
mock_response.json.return_value = tiny_map
with patch("httpx.get", return_value=mock_response):
result = get_model_cost_map("https://fake-url.com/model_prices.json")
# Backup has thousands of models; 11 is a massive shrinkage → fallback
assert len(result) > 11
def test_should_use_local_map_when_env_var_set(self):
"""LITELLM_LOCAL_MODEL_COST_MAP=True should skip remote fetch entirely."""
with patch.dict(os.environ, {"LITELLM_LOCAL_MODEL_COST_MAP": "True"}):
with patch("httpx.get") as mock_get:
result = get_model_cost_map(
"https://fake-url.com/model_prices.json"
)
mock_get.assert_not_called()
assert isinstance(result, dict)
assert len(result) > 0
class TestBackupModelCostMapExists:
"""Validates the local backup file is always present and valid."""
def test_should_have_backup_file(self):
"""The backup model cost map must exist and be loadable."""
backup = GetModelCostMap.load_local_model_cost_map()
assert isinstance(backup, dict)
assert len(backup) > 0, "Backup model cost map is empty"
def test_should_have_minimum_models_in_backup(self):
"""The backup must contain a reasonable number of models."""
backup = GetModelCostMap.load_local_model_cost_map()
assert len(backup) > 100, (
f"Backup has only {len(backup)} models, expected > 100"
)
class TestBadHostedModelCostMap:
"""
Simulates the hosted model cost map being bad (invalid JSON / corrupted).
When the hosted map is bad, get_model_cost_map() falls back to the local
backup. These tests verify that after fallback:
- get_model_info() still works for models in the backup
- litellm.completion() still works
"""
def test_should_model_info_pass_after_bad_hosted_map(self):
"""
If the hosted map is bad, get_model_cost_map falls back to the local
backup. get_model_info should still work for models in the backup.
"""
mock_response = MagicMock()
mock_response.raise_for_status = MagicMock()
mock_response.json.side_effect = json.JSONDecodeError("bad json", "", 0)
with patch("httpx.get", return_value=mock_response):
fallback_map = get_model_cost_map("https://fake-url.com/bad.json")
original = litellm.model_cost
litellm.model_cost = fallback_map
try:
# gpt-4o is in every backup — should work fine
info = litellm.get_model_info("gpt-4o")
assert info is not None
assert info["input_cost_per_token"] > 0
finally:
litellm.model_cost = original
def test_should_completion_pass_after_bad_hosted_map(self):
"""
If the hosted map is bad, litellm.completion() should still work.
Uses litellm's built-in mock_response param so the real completion
path is exercised (routing, cost calculator, logging) without
needing API credentials.
"""
# Simulate bad hosted map → fallback to backup
mock_http = MagicMock()
mock_http.raise_for_status = MagicMock()
mock_http.json.side_effect = json.JSONDecodeError("bad json", "", 0)
with patch("httpx.get", return_value=mock_http):
fallback_map = get_model_cost_map("https://fake-url.com/bad.json")
original = litellm.model_cost
litellm.model_cost = fallback_map
try:
# mock_response goes through the real completion path —
# routing, cost calculator, logging — but skips the HTTP call
response = litellm.completion(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "say hi"}],
mock_response="hello from mock",
)
assert response is not None
assert response.choices[0].message.content == "hello from mock"
finally:
litellm.model_cost = original

View file

@ -274,6 +274,26 @@ class TestOpenTelemetry(unittest.TestCase):
self.assertEqual(config.deployment_environment, "production")
self.assertEqual(config.model_id, "custom-service")
@patch.dict(os.environ, {}, clear=True)
def test_open_telemetry_config_auto_infer_otlp_http_when_endpoint_set(self):
"""When endpoint is set but exporter is default 'console', auto-infer 'otlp_http'.
This fixes an issue where UI-configured OTEL settings would default to console
output instead of sending traces to the configured endpoint.
See: https://github.com/BerriAI/litellm/issues/XXXX
"""
# When endpoint is specified without explicit exporter, should auto-infer otlp_http
config = OpenTelemetryConfig(endpoint="https://otel-collector.example.com:443")
self.assertEqual(config.exporter, "otlp_http")
# When exporter is explicitly set to something other than console, should not override
config_grpc = OpenTelemetryConfig(exporter="grpc", endpoint="https://otel-collector.example.com:443")
self.assertEqual(config_grpc.exporter, "grpc")
# When no endpoint is set, should keep console as default
config_no_endpoint = OpenTelemetryConfig()
self.assertEqual(config_no_endpoint.exporter, "console")
def wait_for_spans(self, exporter: InMemorySpanExporter, prefix: str):
"""Poll until we see at least one span with an attribute key starting with `prefix`."""
deadline = time.time() + self.POLL_TIMEOUT

View file

@ -140,6 +140,83 @@ async def test_ssl_verification_with_aiohttp_transport():
litellm.disable_aiohttp_transport = original_disable
@pytest.mark.asyncio
async def test_ssl_verification_with_shared_session():
"""
Test that ssl_verify=False is respected even with shared sessions.
This was a bug where shared sessions bypassed SSL configuration because
_create_aiohttp_transport returned immediately without passing ssl_verify
to the LiteLLMAiohttpTransport constructor.
The fix stores ssl_verify in the transport and passes it per-request.
"""
import aiohttp
# Ensure aiohttp transport is enabled for this test
original_disable = litellm.disable_aiohttp_transport
litellm.disable_aiohttp_transport = False
try:
# Create a shared session (simulating what happens in production)
shared_session = aiohttp.ClientSession()
try:
# Create transport with shared session and ssl_verify=False
transport = AsyncHTTPHandler._create_aiohttp_transport(
ssl_verify=False,
shared_session=shared_session,
)
# Verify the transport uses the shared session
assert transport.client is shared_session
# Verify the SSL setting is stored in the transport for per-request use
assert transport._ssl_verify is False
finally:
await shared_session.close()
finally:
# Restore original setting
litellm.disable_aiohttp_transport = original_disable
@pytest.mark.asyncio
async def test_ssl_context_with_shared_session():
"""
Test that ssl_context is respected even with shared sessions.
"""
import aiohttp
# Ensure aiohttp transport is enabled for this test
original_disable = litellm.disable_aiohttp_transport
litellm.disable_aiohttp_transport = False
try:
# Create a custom SSL context
custom_ssl_context = ssl.create_default_context()
# Create a shared session
shared_session = aiohttp.ClientSession()
try:
# Create transport with shared session and custom ssl_context
transport = AsyncHTTPHandler._create_aiohttp_transport(
ssl_context=custom_ssl_context,
shared_session=shared_session,
)
# Verify the transport uses the shared session
assert transport.client is shared_session
# Verify the SSL context is stored in the transport for per-request use
assert transport._ssl_verify is custom_ssl_context
finally:
await shared_session.close()
finally:
# Restore original setting
litellm.disable_aiohttp_transport = original_disable
@pytest.mark.asyncio
async def test_aiohttp_transport_trust_env_setting(monkeypatch):
"""Test that trust_env setting is properly configured in aiohttp transport"""

View file

@ -1026,6 +1026,85 @@ async def test_ui_view_spend_logs_with_model(client, monkeypatch):
assert data["data"][0]["model"] == "gpt-3.5-turbo"
@pytest.mark.asyncio
async def test_ui_view_spend_logs_with_model_id(client, monkeypatch):
"""Test that the model_id query param filters spend logs by litellm model deployment id."""
mock_spend_logs = [
{
"id": "log1",
"request_id": "req1",
"api_key": "sk-test-key",
"user": "test_user_1",
"team_id": "team1",
"spend": 0.05,
"startTime": datetime.datetime.now(timezone.utc).isoformat(),
"model": "gpt-3.5-turbo",
"model_id": "deployment-id-1",
"status": "success",
},
{
"id": "log2",
"request_id": "req2",
"api_key": "sk-test-key",
"user": "test_user_2",
"team_id": "team1",
"spend": 0.10,
"startTime": datetime.datetime.now(timezone.utc).isoformat(),
"model": "gpt-4",
"model_id": "deployment-id-2",
"status": "success",
},
]
class MockDB:
async def find_many(self, *args, **kwargs):
if (
"where" in kwargs
and "model_id" in kwargs["where"]
and kwargs["where"]["model_id"] == "deployment-id-1"
):
return [mock_spend_logs[0]]
return mock_spend_logs
async def count(self, *args, **kwargs):
if (
"where" in kwargs
and "model_id" in kwargs["where"]
and kwargs["where"]["model_id"] == "deployment-id-1"
):
return 1
return len(mock_spend_logs)
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
self.db.litellm_spendlogs = self.db
mock_prisma_client = MockPrismaClient()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
start_date = (
datetime.datetime.now(timezone.utc) - datetime.timedelta(days=7)
).strftime("%Y-%m-%d %H:%M:%S")
end_date = datetime.datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M:%S")
response = client.get(
"/spend/logs/ui",
params={
"model_id": "deployment-id-1",
"start_date": start_date,
"end_date": end_date,
},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200
data = response.json()
assert data["total"] == 1
assert len(data["data"]) == 1
assert data["data"][0]["model_id"] == "deployment-id-1"
@pytest.mark.asyncio
async def test_ui_view_spend_logs_with_key_hash(client, monkeypatch):
# Mock data for the test

View file

@ -1990,3 +1990,94 @@ async def test_anthropic_messages_call_type_is_cached():
# This assertion will FAIL if anthropic_messages is filtered out
assert cached_result is not None, "Model ID should be cached for anthropic_messages call type"
assert cached_result["model_id"] == test_model_id, f"Expected {test_model_id}, got {cached_result['model_id']}"
def test_update_kwargs_with_deployment_propagates_model_tags():
"""
Test that deployment-level tags from litellm_params are merged into
kwargs metadata when _update_kwargs_with_deployment is called.
This ensures model-level tags defined in config.yaml appear in SpendLogs.
See: https://github.com/BerriAI/litellm/issues/XXXX
"""
router = litellm.Router(
model_list=[
{
"model_name": "gpt-4o-mini",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "fake-key",
"tags": ["openai-account", "production"],
},
},
],
)
kwargs: dict = {"metadata": {}}
deployment = router.get_deployment_by_model_group_name(
model_group_name="gpt-4o-mini"
)
router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs)
# Deployment tags should be propagated to kwargs metadata
assert "tags" in kwargs["metadata"]
assert "openai-account" in kwargs["metadata"]["tags"]
assert "production" in kwargs["metadata"]["tags"]
def test_update_kwargs_with_deployment_merges_tags_without_duplicates():
"""
Test that when both request-level and deployment-level tags exist,
they are merged without duplicates.
"""
router = litellm.Router(
model_list=[
{
"model_name": "gpt-4o-mini",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "fake-key",
"tags": ["openai-account", "shared-tag"],
},
},
],
)
# Simulate request that already has tags (from request body or key/team level)
kwargs: dict = {"metadata": {"tags": ["user-tag", "shared-tag"]}}
deployment = router.get_deployment_by_model_group_name(
model_group_name="gpt-4o-mini"
)
router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs)
# Both sources should be merged, no duplicates
assert "user-tag" in kwargs["metadata"]["tags"]
assert "openai-account" in kwargs["metadata"]["tags"]
assert "shared-tag" in kwargs["metadata"]["tags"]
assert kwargs["metadata"]["tags"].count("shared-tag") == 1
def test_update_kwargs_with_deployment_no_tags():
"""
Test that when deployment has no tags, kwargs metadata is not affected.
"""
router = litellm.Router(
model_list=[
{
"model_name": "gpt-4o-mini",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "fake-key",
},
},
],
)
kwargs: dict = {"metadata": {}}
deployment = router.get_deployment_by_model_group_name(
model_group_name="gpt-4o-mini"
)
router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs)
# No tags key should be added if deployment has no tags
assert "tags" not in kwargs["metadata"]

View file

@ -13159,6 +13159,21 @@
"type": "github",
"url": "https://github.com/sponsors/wooorm"
}
},
"node_modules/@next/swc-win32-ia32-msvc": {
"version": "14.2.33",
"resolved": "https://registry.npmjs.org/@next/swc-win32-ia32-msvc/-/swc-win32-ia32-msvc-14.2.33.tgz",
"integrity": "sha512-pc9LpGNKhJ0dXQhZ5QMmYxtARwwmWLpeocFmVG5Z0DzWq5Uf0izcI8tLc+qOpqxO1PWqZ5A7J1blrUIKrIFc7Q==",
"cpu": [
"ia32"
],
"optional": true,
"os": [
"win32"
],
"engines": {
"node": ">= 10"
}
}
}
}

View file

@ -3,7 +3,7 @@
"version": "0.1.0",
"private": true,
"scripts": {
"dev": "next dev --webpack",
"dev": "next dev",
"build": "next build",
"start": "next start",
"lint": "next lint",

View file

@ -3,13 +3,14 @@ import { renderHook, waitFor } from "@testing-library/react";
import React, { ReactNode } from "react";
import { beforeEach, describe, expect, it, vi } from "vitest";
import {
useModelsInfo,
useModelHub,
useAllProxyModels,
useInfiniteModelInfo,
useModelHub,
useModelsInfo,
useSelectedTeamModels,
type ProxyModel,
type AllProxyModelsResponse,
type PaginatedModelInfoResponse,
type ProxyModel,
} from "./useModels";
vi.mock("@/components/networking", () => ({
@ -23,7 +24,7 @@ vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
default: () => mockUseAuthorized(),
}));
import { modelInfoCall, modelHubCall, modelAvailableCall } from "@/components/networking";
import { modelAvailableCall, modelHubCall, modelInfoCall } from "@/components/networking";
const mockProxyModel: ProxyModel = {
id: "model-1",
@ -106,7 +107,7 @@ describe("useModelsInfo", () => {
undefined,
undefined,
undefined,
undefined
undefined,
);
expect(modelInfoCall).toHaveBeenCalledTimes(1);
});
@ -130,7 +131,7 @@ describe("useModelsInfo", () => {
undefined,
undefined,
undefined,
undefined
undefined,
);
});
@ -393,7 +394,7 @@ describe("useAllProxyModels", () => {
null,
true,
false,
"expand"
"expand",
);
expect(modelAvailableCall).toHaveBeenCalledTimes(1);
});
@ -531,13 +532,7 @@ describe("useSelectedTeamModels", () => {
expect(result.current.data).toEqual(mockAllProxyModelsResponse);
expect(result.current.error).toBeNull();
expect(modelAvailableCall).toHaveBeenCalledWith(
"test-access-token",
"test-user-id",
"Admin",
true,
"team-1"
);
expect(modelAvailableCall).toHaveBeenCalledWith("test-access-token", "test-user-id", "Admin", true, "team-1");
expect(modelAvailableCall).toHaveBeenCalledTimes(1);
});
@ -639,3 +634,222 @@ describe("useSelectedTeamModels", () => {
expect(modelAvailableCall).not.toHaveBeenCalled();
});
});
describe("useInfiniteModelInfo", () => {
let queryClient: QueryClient;
const mockPageOneResponse: PaginatedModelInfoResponse = {
data: [{ model_name: "gpt-4", model_info: { id: "model-1" } }],
total_count: 2,
current_page: 1,
total_pages: 2,
size: 50,
};
const mockPageTwoResponse: PaginatedModelInfoResponse = {
data: [{ model_name: "claude-3", model_info: { id: "model-2" } }],
total_count: 2,
current_page: 2,
total_pages: 2,
size: 50,
};
beforeEach(() => {
queryClient = new QueryClient({
defaultOptions: {
queries: {
retry: false,
},
},
});
vi.clearAllMocks();
mockUseAuthorized.mockReturnValue({
accessToken: "test-access-token",
userId: "test-user-id",
userRole: "Admin",
token: "test-token",
userEmail: "test@example.com",
premiumUser: false,
disabledPersonalKeyCreation: null,
showSSOBanner: false,
});
});
const wrapper = ({ children }: { children: ReactNode }) =>
React.createElement(QueryClientProvider, { client: queryClient }, children);
it("should return defined result", () => {
(modelInfoCall as any).mockResolvedValue(mockPageOneResponse);
const { result } = renderHook(() => useInfiniteModelInfo(), { wrapper });
expect(result.current).toBeDefined();
expect(result.current).toHaveProperty("data");
expect(result.current).toHaveProperty("fetchNextPage");
expect(result.current).toHaveProperty("hasNextPage");
expect(result.current).toHaveProperty("isFetchingNextPage");
expect(result.current).toHaveProperty("isLoading");
});
it("should return paginated data and call modelInfoCall with page 1 initially", async () => {
(modelInfoCall as any).mockResolvedValue(mockPageOneResponse);
const { result } = renderHook(() => useInfiniteModelInfo(), { wrapper });
expect(result.current.isLoading).toBe(true);
await waitFor(() => {
expect(result.current.isLoading).toBe(false);
expect(result.current.isSuccess).toBe(true);
});
expect(result.current.data?.pages).toHaveLength(1);
expect(result.current.data?.pages[0]).toEqual(mockPageOneResponse);
expect(result.current.hasNextPage).toBe(true);
expect(modelInfoCall).toHaveBeenCalledWith("test-access-token", "test-user-id", "Admin", 1, 50, undefined);
expect(modelInfoCall).toHaveBeenCalledTimes(1);
});
it("should use custom size parameter", async () => {
(modelInfoCall as any).mockResolvedValue(mockPageOneResponse);
const { result } = renderHook(() => useInfiniteModelInfo(25), { wrapper });
await waitFor(() => {
expect(result.current.isLoading).toBe(false);
});
expect(modelInfoCall).toHaveBeenCalledWith("test-access-token", "test-user-id", "Admin", 1, 25, undefined);
});
it("should pass search parameter to modelInfoCall", async () => {
(modelInfoCall as any).mockResolvedValue(mockPageOneResponse);
const { result } = renderHook(() => useInfiniteModelInfo(50, "gpt"), { wrapper });
await waitFor(() => {
expect(result.current.isLoading).toBe(false);
});
expect(modelInfoCall).toHaveBeenCalledWith("test-access-token", "test-user-id", "Admin", 1, 50, "gpt");
});
it("should fetch next page when fetchNextPage is called", async () => {
(modelInfoCall as any).mockResolvedValueOnce(mockPageOneResponse).mockResolvedValueOnce(mockPageTwoResponse);
const { result } = renderHook(() => useInfiniteModelInfo(), { wrapper });
await waitFor(() => {
expect(result.current.isSuccess).toBe(true);
expect(result.current.hasNextPage).toBe(true);
});
await result.current.fetchNextPage();
await waitFor(() => {
expect(result.current.data?.pages).toHaveLength(2);
expect(result.current.data?.pages[1]).toEqual(mockPageTwoResponse);
expect(result.current.hasNextPage).toBe(false);
});
expect(modelInfoCall).toHaveBeenNthCalledWith(2, "test-access-token", "test-user-id", "Admin", 2, 50, undefined);
});
it("should return undefined for hasNextPage when on last page", async () => {
const lastPageResponse: PaginatedModelInfoResponse = {
...mockPageOneResponse,
current_page: 1,
total_pages: 1,
};
(modelInfoCall as any).mockResolvedValue(lastPageResponse);
const { result } = renderHook(() => useInfiniteModelInfo(), { wrapper });
await waitFor(() => {
expect(result.current.isSuccess).toBe(true);
});
expect(result.current.hasNextPage).toBe(false);
});
it("should handle error when modelInfoCall fails", async () => {
const errorMessage = "Failed to fetch models";
const testError = new Error(errorMessage);
(modelInfoCall as any).mockRejectedValue(testError);
const { result } = renderHook(() => useInfiniteModelInfo(), { wrapper });
expect(result.current.isLoading).toBe(true);
await waitFor(() => {
expect(result.current.isLoading).toBe(false);
expect(result.current.isError).toBe(true);
});
expect(result.current.error).toEqual(testError);
expect(result.current.data).toBeUndefined();
expect(modelInfoCall).toHaveBeenCalledTimes(1);
});
it("should not execute query when accessToken is missing", () => {
mockUseAuthorized.mockReturnValue({
accessToken: null,
userId: "test-user-id",
userRole: "Admin",
token: null,
userEmail: "test@example.com",
premiumUser: false,
disabledPersonalKeyCreation: null,
showSSOBanner: false,
});
const { result } = renderHook(() => useInfiniteModelInfo(), { wrapper });
expect(result.current.isLoading).toBe(false);
expect(result.current.data).toBeUndefined();
expect(result.current.isFetched).toBe(false);
expect(modelInfoCall).not.toHaveBeenCalled();
});
it("should not execute query when userId is missing", () => {
mockUseAuthorized.mockReturnValue({
accessToken: "test-access-token",
userId: null,
userRole: "Admin",
token: "test-token",
userEmail: "test@example.com",
premiumUser: false,
disabledPersonalKeyCreation: null,
showSSOBanner: false,
});
const { result } = renderHook(() => useInfiniteModelInfo(), { wrapper });
expect(result.current.isLoading).toBe(false);
expect(result.current.data).toBeUndefined();
expect(result.current.isFetched).toBe(false);
expect(modelInfoCall).not.toHaveBeenCalled();
});
it("should not execute query when userRole is missing", () => {
mockUseAuthorized.mockReturnValue({
accessToken: "test-access-token",
userId: "test-user-id",
userRole: null,
token: "test-token",
userEmail: "test@example.com",
premiumUser: false,
disabledPersonalKeyCreation: null,
showSSOBanner: false,
});
const { result } = renderHook(() => useInfiniteModelInfo(), { wrapper });
expect(result.current.isLoading).toBe(false);
expect(result.current.data).toBeUndefined();
expect(result.current.isFetched).toBe(false);
expect(modelInfoCall).not.toHaveBeenCalled();
});
});

View file

@ -1,4 +1,4 @@
import { useQuery } from "@tanstack/react-query";
import { useQuery, useInfiniteQuery } from "@tanstack/react-query";
import { createQueryKeys } from "../common/queryKeysFactory";
import { modelInfoCall, modelHubCall, modelAvailableCall } from "@/components/networking";
import useAuthorized from "../useAuthorized";
@ -26,6 +26,7 @@ const modelKeys = createQueryKeys("models");
const modelHubKeys = createQueryKeys("modelHub");
const allProxyModelsKeys = createQueryKeys("allProxyModels");
const selectedTeamModelsKeys = createQueryKeys("selectedTeamModels");
const infiniteModelKeys = createQueryKeys("infiniteModels");
export const useModelsInfo = (page: number = 1, size: number = 50, search?: string, modelId?: string, teamId?: string, sortBy?: string, sortOrder?: string) => {
const { accessToken, userId, userRole } = useAuthorized();
@ -74,3 +75,38 @@ export const useSelectedTeamModels = (teamID: string | null) => {
enabled: Boolean(accessToken && userId && userRole && teamID),
});
};
export const useInfiniteModelInfo = (
size: number = 50,
search?: string,
) => {
const { accessToken, userId, userRole } = useAuthorized();
return useInfiniteQuery<PaginatedModelInfoResponse>({
queryKey: infiniteModelKeys.list({
filters: {
...(userId && { userId }),
...(userRole && { userRole }),
size,
...(search && { search }),
},
}),
queryFn: async ({ pageParam }) => {
return await modelInfoCall(
accessToken!,
userId!,
userRole!,
pageParam as number,
size,
search,
);
},
initialPageParam: 1,
getNextPageParam: (lastPage) => {
if (lastPage.current_page < lastPage.total_pages) {
return lastPage.current_page + 1;
}
return undefined;
},
enabled: Boolean(accessToken && userId && userRole),
});
};

View file

@ -0,0 +1,301 @@
import { screen, waitFor } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { beforeEach, describe, expect, it, vi } from "vitest";
import { renderWithProviders } from "../../../../tests/test-utils";
import { PaginatedModelSelect } from "./PaginatedModelSelect";
const mockFetchNextPage = vi.fn();
vi.mock("@/app/(dashboard)/hooks/models/useModels", () => ({
useInfiniteModelInfo: vi.fn(),
}));
vi.mock("@tanstack/react-pacer/debouncer", () => {
const React = require("react");
return {
useDebouncedState: (initial: string) => {
const [value, setValue] = React.useState(initial);
return [value, setValue];
},
};
});
import { useInfiniteModelInfo } from "@/app/(dashboard)/hooks/models/useModels";
const mockUseInfiniteModelInfo = vi.mocked(useInfiniteModelInfo);
const mockPagesWithModels = {
pages: [
{
data: [
{ model_name: "GPT-4", model_info: { id: "model-1" } },
{ model_name: "Claude-3", model_info: { id: "model-2" } },
],
total_count: 2,
current_page: 1,
total_pages: 1,
size: 50,
},
],
};
const mockEmptyPages = {
pages: [{ data: [], total_count: 0, current_page: 1, total_pages: 1, size: 50 }],
};
describe("PaginatedModelSelect", () => {
const mockOnChange = vi.fn();
const defaultHookReturn = {
data: mockPagesWithModels,
fetchNextPage: mockFetchNextPage,
hasNextPage: false,
isFetchingNextPage: false,
isLoading: false,
};
beforeEach(() => {
vi.clearAllMocks();
mockUseInfiniteModelInfo.mockReturnValue(defaultHookReturn as any);
});
it("should render", () => {
renderWithProviders(<PaginatedModelSelect onChange={mockOnChange} />);
expect(screen.getByRole("combobox")).toBeInTheDocument();
expect(screen.getByText("Select a model")).toBeInTheDocument();
});
it("should display custom placeholder when provided", () => {
renderWithProviders(
<PaginatedModelSelect onChange={mockOnChange} placeholder="Choose model" />,
);
expect(screen.getByText("Choose model")).toBeInTheDocument();
});
it("should display model options when data is loaded", async () => {
renderWithProviders(<PaginatedModelSelect onChange={mockOnChange} />);
const combobox = screen.getByRole("combobox");
await userEvent.click(combobox);
await waitFor(() => {
expect(screen.getByRole("option", { name: "GPT-4 (model-1)" })).toBeInTheDocument();
expect(screen.getByRole("option", { name: "Claude-3 (model-2)" })).toBeInTheDocument();
});
});
it("should call onChange when user selects a model", async () => {
const user = userEvent.setup({ delay: null });
renderWithProviders(<PaginatedModelSelect onChange={mockOnChange} />);
const combobox = screen.getByRole("combobox");
await user.click(combobox);
const visibleOption = await screen.findByTitle("GPT-4 (model-1)");
await user.click(visibleOption);
await waitFor(() => {
expect(mockOnChange).toHaveBeenCalledWith("model-1");
});
});
it("should display selected value when value prop is provided", async () => {
renderWithProviders(
<PaginatedModelSelect value="model-1" onChange={mockOnChange} />,
);
const combobox = screen.getByRole("combobox");
await userEvent.click(combobox);
await waitFor(() => {
expect(screen.getByRole("option", { name: "GPT-4 (model-1)" })).toBeInTheDocument();
});
});
it("should show loading state when isLoading is true", () => {
mockUseInfiniteModelInfo.mockReturnValue({
...defaultHookReturn,
isLoading: true,
} as any);
renderWithProviders(<PaginatedModelSelect onChange={mockOnChange} />);
expect(screen.getByRole("combobox")).toHaveAttribute("aria-expanded", "false");
});
it("should pass pageSize to useInfiniteModelInfo", () => {
renderWithProviders(<PaginatedModelSelect onChange={mockOnChange} pageSize={25} />);
expect(mockUseInfiniteModelInfo).toHaveBeenCalledWith(25, undefined);
});
it("should pass search to useInfiniteModelInfo when user types", async () => {
const user = userEvent.setup();
renderWithProviders(<PaginatedModelSelect onChange={mockOnChange} />);
const combobox = screen.getByRole("combobox");
await user.click(combobox);
await user.keyboard("gpt");
await waitFor(() => {
expect(mockUseInfiniteModelInfo).toHaveBeenCalledWith(50, "gpt");
});
});
it("should have scroll container for infinite loading when hasNextPage is true", async () => {
mockUseInfiniteModelInfo.mockReturnValue({
...defaultHookReturn,
hasNextPage: true,
isFetchingNextPage: false,
} as any);
renderWithProviders(<PaginatedModelSelect onChange={mockOnChange} />);
const combobox = screen.getByRole("combobox");
await userEvent.click(combobox);
await waitFor(() => {
expect(screen.getByRole("option", { name: "GPT-4 (model-1)" })).toBeInTheDocument();
});
const scrollableContainer = document.querySelector(
".ant-select-dropdown .rc-virtual-list-holder",
);
expect(scrollableContainer).toBeInTheDocument();
expect(scrollableContainer).toHaveAttribute("style");
});
it("should deduplicate models with same id across pages", async () => {
mockUseInfiniteModelInfo.mockReturnValue({
...defaultHookReturn,
data: {
pages: [
{
data: [
{ model_name: "GPT-4", model_info: { id: "model-1" } },
{ model_name: "GPT-4 Dupe", model_info: { id: "model-1" } },
],
total_count: 2,
current_page: 1,
total_pages: 1,
size: 50,
},
],
},
fetchNextPage: mockFetchNextPage,
hasNextPage: false,
isFetchingNextPage: false,
isLoading: false,
} as any);
renderWithProviders(<PaginatedModelSelect onChange={mockOnChange} />);
const combobox = screen.getByRole("combobox");
await userEvent.click(combobox);
await waitFor(() => {
const model1Options = screen.queryAllByRole("option", { name: /model-1/ });
expect(model1Options.length).toBe(1);
});
});
it("should skip models without model_info id", async () => {
mockUseInfiniteModelInfo.mockReturnValue({
...defaultHookReturn,
data: {
pages: [
{
data: [
{ model_name: "Valid Model", model_info: { id: "valid-id" } },
{ model_name: "No ID", model_info: null },
{ model_name: "Empty ID", model_info: { id: "" } },
],
total_count: 3,
current_page: 1,
total_pages: 1,
size: 50,
},
],
},
fetchNextPage: mockFetchNextPage,
hasNextPage: false,
isFetchingNextPage: false,
isLoading: false,
} as any);
renderWithProviders(<PaginatedModelSelect onChange={mockOnChange} />);
const combobox = screen.getByRole("combobox");
await userEvent.click(combobox);
await waitFor(() => {
expect(screen.getByRole("option", { name: "Valid Model (valid-id)" })).toBeInTheDocument();
expect(screen.queryByRole("option", { name: "No ID" })).not.toBeInTheDocument();
expect(screen.queryByRole("option", { name: "Empty ID" })).not.toBeInTheDocument();
});
});
it("should show model ID only when model_name is empty", async () => {
mockUseInfiniteModelInfo.mockReturnValue({
...defaultHookReturn,
data: {
pages: [
{
data: [{ model_name: "", model_info: { id: "id-only" } }],
total_count: 1,
current_page: 1,
total_pages: 1,
size: 50,
},
],
},
fetchNextPage: mockFetchNextPage,
hasNextPage: false,
isFetchingNextPage: false,
isLoading: false,
} as any);
renderWithProviders(<PaginatedModelSelect onChange={mockOnChange} />);
const combobox = screen.getByRole("combobox");
await userEvent.click(combobox);
await waitFor(() => {
expect(screen.getByRole("option", { name: "id-only" })).toBeInTheDocument();
});
});
it("should respect allowClear prop", () => {
renderWithProviders(
<PaginatedModelSelect value="model-1" onChange={mockOnChange} allowClear={false} />,
);
expect(screen.getByRole("combobox")).toBeInTheDocument();
});
it("should respect disabled prop", () => {
renderWithProviders(<PaginatedModelSelect onChange={mockOnChange} disabled />);
const combobox = screen.getByRole("combobox");
expect(combobox.closest(".ant-select")).toHaveClass("ant-select-disabled");
});
it("should not call fetchNextPage when hasNextPage is false", async () => {
mockUseInfiniteModelInfo.mockReturnValue({
...defaultHookReturn,
hasNextPage: false,
} as any);
renderWithProviders(<PaginatedModelSelect onChange={mockOnChange} />);
await userEvent.click(screen.getByRole("combobox"));
await waitFor(() => {
expect(screen.getByRole("option", { name: "GPT-4 (model-1)" })).toBeInTheDocument();
});
expect(mockFetchNextPage).not.toHaveBeenCalled();
});
});

View file

@ -0,0 +1,143 @@
import { useInfiniteModelInfo } from "@/app/(dashboard)/hooks/models/useModels";
import { LoadingOutlined } from "@ant-design/icons";
import { useDebouncedState } from "@tanstack/react-pacer/debouncer";
import { Select, Space, Typography } from "antd";
import { useMemo, useState, type UIEvent } from "react";
const { Text } = Typography;
export interface PaginatedModelSelectProps {
value?: string;
onChange?: (value: string) => void;
placeholder?: string;
style?: React.CSSProperties;
pageSize?: number;
allowClear?: boolean;
disabled?: boolean;
}
const SCROLL_THRESHOLD = 0.8;
const DEBOUNCE_MS = 300;
export const PaginatedModelSelect = ({
value,
onChange,
placeholder = "Select a model",
style,
pageSize = 50,
allowClear = true,
disabled = false,
}: PaginatedModelSelectProps) => {
const [searchInput, setSearchInput] = useState("");
const [debouncedSearch, setDebouncedSearch] = useDebouncedState("", {
wait: DEBOUNCE_MS,
});
const {
data,
fetchNextPage,
hasNextPage,
isFetchingNextPage,
isLoading,
} = useInfiniteModelInfo(pageSize, debouncedSearch || undefined);
const options = useMemo(() => {
if (!data?.pages) return [];
const seen = new Set<string>();
const result: { label: string; value: string; modelName: string; modelId: string }[] = [];
for (const page of data.pages) {
for (const model of page.data) {
const modelId = model.model_info?.id ?? "";
const modelName = model.model_name ?? "";
// Dedupe by id - skip models without id (can't uniquely identify)
if (!modelId || seen.has(modelId)) continue;
seen.add(modelId);
result.push({
label: modelName ? `${modelName} (${modelId})` : modelId,
value: modelId,
modelName,
modelId,
});
}
}
return result;
}, [data]);
const optionRender = (option: { data: { modelName: string; modelId: string; label: string } }) => {
const { modelName, modelId } = option.data;
return (
<>
{modelName ? (
<Space direction="vertical">
<Space direction="horizontal">
<Text strong>Model name:</Text>
<Text ellipsis>{modelName}</Text>
</Space>
<Text ellipsis type="secondary" >
Model ID: {modelId}
</Text>
</Space>
) : (
<Text ellipsis type="secondary">Model ID: {modelId}</Text>
)}
</>
);
};
const handlePopupScroll = (e: UIEvent<HTMLDivElement>) => {
const target = e.currentTarget;
const scrollRatio =
(target.scrollTop + target.clientHeight) / target.scrollHeight;
if (scrollRatio >= SCROLL_THRESHOLD && hasNextPage && !isFetchingNextPage) {
fetchNextPage();
}
};
const handleSearch = (value: string) => {
setSearchInput(value);
setDebouncedSearch(value);
};
const handleChange = (v: string | string[] | null) => {
const normalized =
typeof v === "string" ? v : Array.isArray(v) ? v[0] ?? "" : "";
onChange?.(normalized);
};
return (
<Select
value={value || undefined}
onChange={handleChange}
placeholder={placeholder}
style={{ width: "100%", ...style }}
allowClear={allowClear}
disabled={disabled}
showSearch
filterOption={false}
onSearch={handleSearch}
searchValue={searchInput}
onPopupScroll={handlePopupScroll}
loading={isLoading}
notFoundContent={isLoading ? <LoadingOutlined spin /> : "No models found"}
options={options}
optionRender={optionRender}
popupRender={(menu) => (
<>
{menu}
{isFetchingNextPage && (
<div style={{ textAlign: "center", padding: 8 }}>
<LoadingOutlined spin />
</div>
)}
</>
)}
/>
);
};

View file

@ -1,5 +1,5 @@
import React, { useState, useRef, useEffect } from "react";
import { Modal, Select, Switch, Collapse, Input } from "antd";
import { Modal, Select, Switch, Collapse, Input, Divider } from "antd";
import { Button, TextInput } from "@tremor/react";
import {
CodeOutlined,
@ -8,6 +8,8 @@ import {
CloseCircleOutlined,
CaretRightOutlined,
SaveOutlined,
UsergroupAddOutlined,
ExportOutlined,
} from "@ant-design/icons";
import { createGuardrailCall, updateGuardrailCall, testCustomCodeGuardrail } from "../../networking";
import NotificationsManager from "../../molecules/notifications_manager";
@ -91,6 +93,7 @@ const CODE_TEMPLATES = {
},
};
// Available primitives organized by category
const PRIMITIVES = {
"Return Values": [
@ -241,6 +244,8 @@ const CustomCodeModal: React.FC<CustomCodeModalProps> = ({
// Handle template change
const handleTemplateChange = (templateKey: string) => {
setSelectedTemplate(templateKey);
// Check if it's a standard template
setCode(CODE_TEMPLATES[templateKey as keyof typeof CODE_TEMPLATES].code);
};
@ -486,12 +491,45 @@ const CustomCodeModal: React.FC<CustomCodeModalProps> = ({
onChange={handleTemplateChange}
className="w-full"
size="middle"
dropdownRender={(menu) => (
<>
{menu}
<Divider style={{ margin: '8px 0' }} />
<div
style={{
padding: '8px 12px',
cursor: 'pointer',
color: '#1890ff',
fontSize: '12px',
display: 'flex',
alignItems: 'center',
gap: '4px',
}}
onClick={(e) => {
e.preventDefault();
window.open('https://models.litellm.ai/guardrails', '_blank');
}}
onMouseEnter={(e) => {
e.currentTarget.style.backgroundColor = '#f0f0f0';
}}
onMouseLeave={(e) => {
e.currentTarget.style.backgroundColor = 'transparent';
}}
>
<UsergroupAddOutlined />
<span>Browse Community templates</span>
<ExportOutlined style={{ fontSize: '10px' }} />
</div>
</>
)}
>
{Object.entries(CODE_TEMPLATES).map(([key, template]) => (
<Select.Option key={key} value={key}>
{template.name}
</Select.Option>
))}
<Select.OptGroup label="STANDARD">
{Object.entries(CODE_TEMPLATES).map(([key, template]) => (
<Select.Option key={key} value={key}>
{template.name}
</Select.Option>
))}
</Select.OptGroup>
</Select>
</div>
<div className="flex items-center gap-2 pt-5">
@ -632,6 +670,27 @@ const CustomCodeModal: React.FC<CustomCodeModalProps> = ({
</div>
</Panel>
</Collapse>
{/* Contribution CTA Banner */}
<div className="mt-3 p-4 bg-gradient-to-r from-blue-50 to-indigo-50 border border-blue-200 rounded-lg flex items-center justify-between flex-shrink-0">
<div className="flex items-center gap-3">
<div className="bg-blue-100 rounded-full p-2">
<UsergroupAddOutlined className="text-blue-600 text-lg" />
</div>
<div>
<div className="text-sm font-medium text-gray-900">Built a useful guardrail?</div>
<div className="text-xs text-gray-600">Share it with the community and help others build faster</div>
</div>
</div>
<Button
size="xs"
onClick={() => window.open('https://github.com/BerriAI/litellm-guardrails', '_blank')}
icon={ExportOutlined}
className="bg-blue-600 hover:bg-blue-700 text-white border-0"
>
Contribute Template
</Button>
</div>
</div>
{/* Primitives Panel */}

View file

@ -1,7 +1,13 @@
import React, { useState, useCallback, useEffect } from "react";
import { Button, Input, Select } from "antd";
import { FilterIcon } from "@heroicons/react/outline";
import { Button, Input, Select } from "antd";
import debounce from "lodash/debounce";
import React, { useCallback, useEffect, useState } from "react";
export interface FilterOptionCustomComponentProps {
value?: string;
onChange: (value: string) => void;
placeholder?: string;
}
export interface FilterOption {
name: string;
@ -9,6 +15,7 @@ export interface FilterOption {
isSearchable?: boolean;
searchFn?: (searchText: string) => Promise<Array<{ label: string; value: string }>>;
options?: Array<{ label: string; value: string }>;
customComponent?: React.ComponentType<FilterOptionCustomComponentProps>;
}
interface FilterValues {
@ -194,6 +201,17 @@ const FilterComponent: React.FC<FilterComponentProps> = ({
</Select.Option>
))}
</Select>
) : option.customComponent ? (
(() => {
const CustomComponent = option.customComponent;
return (
<CustomComponent
value={tempValues[option.name] || undefined}
onChange={(value) => handleFilterChange(option.name, value ?? "")}
placeholder={`Select ${option.label || option.name}...`}
/>
);
})()
) : (
<Input
className="w-full"

View file

@ -2593,6 +2593,7 @@ export const uiSpendLogsCall = async (
end_user?: string,
status_filter?: string,
model?: string,
modelId?: string,
keyAlias?: string,
error_code?: string,
error_message?: string,
@ -2614,6 +2615,7 @@ export const uiSpendLogsCall = async (
if (end_user) queryParams.append("end_user", end_user);
if (status_filter) queryParams.append("status_filter", status_filter);
if (model) queryParams.append("model", model);
if (modelId) queryParams.append("model_id", modelId);
if (keyAlias) queryParams.append("key_alias", keyAlias);
if (error_code) queryParams.append("error_code", error_code);
if (error_message) queryParams.append("error_message", error_message);

File diff suppressed because it is too large Load diff

View file

@ -10,29 +10,30 @@ import { Row } from "@tanstack/react-table";
import { Switch, Tab, TabGroup, TabList, TabPanel, TabPanels } from "@tremor/react";
import { Button, Tooltip } from "antd";
import { internalUserRoles } from "../../utils/roles";
import NewBadge from "../common_components/NewBadge";
import DeletedKeysPage from "../DeletedKeysPage/DeletedKeysPage";
import DeletedTeamsPage from "../DeletedTeamsPage/DeletedTeamsPage";
import { fetchAllKeyAliases } from "../key_team_helpers/filter_helpers";
import { KeyResponse, Team } from "../key_team_helpers/key_list";
import { PaginatedModelSelect } from "../ModelSelect/PaginatedModelSelect/PaginatedModelSelect";
import FilterComponent, { FilterOption } from "../molecules/filter";
import { allEndUsersCall, keyInfoV1Call, keyListCall, sessionSpendLogsCall, uiSpendLogsCall } from "../networking";
import KeyInfoView from "../templates/key_info_view";
import AuditLogs from "./audit_logs";
import { columns, LogEntry } from "./columns";
import { ConfigInfoMessage } from "./ConfigInfoMessage";
import { ERROR_CODE_OPTIONS, QUICK_SELECT_OPTIONS } from "./constants";
import { CostBreakdownViewer } from "./CostBreakdownViewer";
import { ErrorViewer } from "./ErrorViewer";
import { useLogFilterLogic } from "./log_filter_logic";
import { LogDetailsDrawer } from "./LogDetailsDrawer";
import { getTimeRangeDisplay } from "./logs_utils";
import { prefetchLogDetails } from "./prefetch";
import { ERROR_CODE_OPTIONS, QUICK_SELECT_OPTIONS } from "./constants";
import { RequestResponsePanel } from "./RequestResponsePanel";
import { SessionView } from "./SessionView";
import SpendLogsSettingsModal from "./SpendLogsSettingsModal/SpendLogsSettingsModal";
import { DataTable } from "./table";
import { VectorStoreViewer } from "./VectorStoreViewer";
import NewBadge from "../common_components/NewBadge";
import { LogDetailsDrawer } from "./LogDetailsDrawer";
interface SpendLogsTableProps {
accessToken: string | null;
@ -205,7 +206,8 @@ export default function SpendLogsTable({
filterByCurrentUser ? userID : undefined,
selectedEndUser,
selectedStatus,
selectedModel,
undefined,
selectedModel || undefined,
);
// Trigger prefetch for all logs
@ -404,7 +406,7 @@ export default function SpendLogsTable({
{
name: "Model",
label: "Model",
isSearchable: false,
customComponent: PaginatedModelSelect,
},
{
name: "Key Alias",

View file

@ -96,6 +96,7 @@ export function useLogFilterLogic({
filters[FILTER_KEYS.USER_ID] || undefined,
filters[FILTER_KEYS.END_USER] || undefined,
filters[FILTER_KEYS.STATUS] || undefined,
undefined,
filters[FILTER_KEYS.MODEL] || undefined,
filters[FILTER_KEYS.KEY_ALIAS] || undefined,
filters[FILTER_KEYS.ERROR_CODE] || undefined,

View file

@ -14,7 +14,7 @@
"moduleResolution": "bundler",
"resolveJsonModule": true,
"isolatedModules": true,
"jsx": "react-jsx",
"jsx": "preserve",
"incremental": true,
"plugins": [
{