mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
Merge branch 'main' into litellm_release_day_02_10_2026
This commit is contained in:
commit
01a07903f6
80 changed files with 2090 additions and 1463 deletions
|
|
@ -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
|
||||
|
|
|
|||
95
docs/my-website/blog/model_cost_map_incident/index.md
Normal file
95
docs/my-website/blog/model_cost_map_incident/index.md
Normal 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 |
|
||||
|
|
@ -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",
|
||||
},
|
||||
],
|
||||
},
|
||||
],
|
||||
};
|
||||
|
||||
|
|
|
|||
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.34-py3-none-any.whl
vendored
Normal file
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.34-py3-none-any.whl
vendored
Normal file
Binary file not shown.
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.34.tar.gz
vendored
Normal file
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.34.tar.gz
vendored
Normal file
Binary file not shown.
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
126
tests/code_coverage_tests/check_guardrail_apply_decorator.py
Normal file
126
tests/code_coverage_tests/check_guardrail_apply_decorator.py
Normal 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")
|
||||
291
tests/llm_translation/test_model_cost_map_resilience.py
Normal file
291
tests/llm_translation/test_model_cost_map_resilience.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
15
ui/litellm-dashboard/package-lock.json
generated
15
ui/litellm-dashboard/package-lock.json
generated
|
|
@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
});
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
)}
|
||||
</>
|
||||
)}
|
||||
/>
|
||||
);
|
||||
};
|
||||
|
|
@ -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 */}
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -14,7 +14,7 @@
|
|||
"moduleResolution": "bundler",
|
||||
"resolveJsonModule": true,
|
||||
"isolatedModules": true,
|
||||
"jsx": "react-jsx",
|
||||
"jsx": "preserve",
|
||||
"incremental": true,
|
||||
"plugins": [
|
||||
{
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue