diff --git a/.circleci/config.yml b/.circleci/config.yml index f48b3769e0c..34c3f05cd25 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -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 diff --git a/docs/my-website/blog/model_cost_map_incident/index.md b/docs/my-website/blog/model_cost_map_incident/index.md new file mode 100644 index 00000000000..b9ff20e4128 --- /dev/null +++ b/docs/my-website/blog/model_cost_map_incident/index.md @@ -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 | diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index 2c3dfb2b863..28e1724ef82 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -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", + }, + ], + }, ], }; diff --git a/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.34-py3-none-any.whl b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.34-py3-none-any.whl new file mode 100644 index 00000000000..175d84543ec Binary files /dev/null and b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.34-py3-none-any.whl differ diff --git a/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.34.tar.gz b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.34.tar.gz new file mode 100644 index 00000000000..e1fcc0c603f Binary files /dev/null and b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.34.tar.gz differ diff --git a/litellm/constants.py b/litellm/constants.py index 9c25cf77906..180315ace0e 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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) diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index bbd55a59bce..407bc581f71 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -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( diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 138d508db4b..b847180174a 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -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: diff --git a/litellm/litellm_core_utils/get_model_cost_map.py b/litellm/litellm_core_utils/get_model_cost_map.py index 9b86f4ca2f0..e622a317454 100644 --- a/litellm/litellm_core_utils/get_model_cost_map.py +++ b/litellm/litellm_core_utils/get_model_cost_map.py @@ -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 diff --git a/litellm/llms/custom_httpx/aiohttp_transport.py b/litellm/llms/custom_httpx/aiohttp_transport.py index a7b83d8c802..1b03ec47643 100644 --- a/litellm/llms/custom_httpx/aiohttp_transport.py +++ b/litellm/llms/custom_httpx/aiohttp_transport.py @@ -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 diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index ac9dd5998e2..95f411c397c 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -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 diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 4812a2d8a10..d794aa50d2e 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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", diff --git a/litellm/proxy/_experimental/out/404.html b/litellm/proxy/_experimental/out/404/index.html similarity index 100% rename from litellm/proxy/_experimental/out/404.html rename to litellm/proxy/_experimental/out/404/index.html diff --git a/litellm/proxy/_experimental/out/_not-found.html b/litellm/proxy/_experimental/out/_not-found/index.html similarity index 100% rename from litellm/proxy/_experimental/out/_not-found.html rename to litellm/proxy/_experimental/out/_not-found/index.html diff --git a/litellm/proxy/_experimental/out/api-reference.html b/litellm/proxy/_experimental/out/api-reference/index.html similarity index 100% rename from litellm/proxy/_experimental/out/api-reference.html rename to litellm/proxy/_experimental/out/api-reference/index.html diff --git a/litellm/proxy/_experimental/out/experimental/api-playground.html b/litellm/proxy/_experimental/out/experimental/api-playground/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/api-playground.html rename to litellm/proxy/_experimental/out/experimental/api-playground/index.html diff --git a/litellm/proxy/_experimental/out/experimental/budgets.html b/litellm/proxy/_experimental/out/experimental/budgets/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/budgets.html rename to litellm/proxy/_experimental/out/experimental/budgets/index.html diff --git a/litellm/proxy/_experimental/out/experimental/caching.html b/litellm/proxy/_experimental/out/experimental/caching/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/caching.html rename to litellm/proxy/_experimental/out/experimental/caching/index.html diff --git a/litellm/proxy/_experimental/out/experimental/claude-code-plugins.html b/litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/claude-code-plugins.html rename to litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html diff --git a/litellm/proxy/_experimental/out/experimental/old-usage.html b/litellm/proxy/_experimental/out/experimental/old-usage/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/old-usage.html rename to litellm/proxy/_experimental/out/experimental/old-usage/index.html diff --git a/litellm/proxy/_experimental/out/experimental/prompts.html b/litellm/proxy/_experimental/out/experimental/prompts/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/prompts.html rename to litellm/proxy/_experimental/out/experimental/prompts/index.html diff --git a/litellm/proxy/_experimental/out/experimental/tag-management.html b/litellm/proxy/_experimental/out/experimental/tag-management/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/tag-management.html rename to litellm/proxy/_experimental/out/experimental/tag-management/index.html diff --git a/litellm/proxy/_experimental/out/guardrails.html b/litellm/proxy/_experimental/out/guardrails/index.html similarity index 100% rename from litellm/proxy/_experimental/out/guardrails.html rename to litellm/proxy/_experimental/out/guardrails/index.html diff --git a/litellm/proxy/_experimental/out/login.html b/litellm/proxy/_experimental/out/login/index.html similarity index 100% rename from litellm/proxy/_experimental/out/login.html rename to litellm/proxy/_experimental/out/login/index.html diff --git a/litellm/proxy/_experimental/out/logs.html b/litellm/proxy/_experimental/out/logs/index.html similarity index 100% rename from litellm/proxy/_experimental/out/logs.html rename to litellm/proxy/_experimental/out/logs/index.html diff --git a/litellm/proxy/_experimental/out/mcp/oauth/callback.html b/litellm/proxy/_experimental/out/mcp/oauth/callback/index.html similarity index 100% rename from litellm/proxy/_experimental/out/mcp/oauth/callback.html rename to litellm/proxy/_experimental/out/mcp/oauth/callback/index.html diff --git a/litellm/proxy/_experimental/out/model-hub.html b/litellm/proxy/_experimental/out/model-hub/index.html similarity index 100% rename from litellm/proxy/_experimental/out/model-hub.html rename to litellm/proxy/_experimental/out/model-hub/index.html diff --git a/litellm/proxy/_experimental/out/model_hub.html b/litellm/proxy/_experimental/out/model_hub/index.html similarity index 100% rename from litellm/proxy/_experimental/out/model_hub.html rename to litellm/proxy/_experimental/out/model_hub/index.html diff --git a/litellm/proxy/_experimental/out/model_hub_table.html b/litellm/proxy/_experimental/out/model_hub_table/index.html similarity index 100% rename from litellm/proxy/_experimental/out/model_hub_table.html rename to litellm/proxy/_experimental/out/model_hub_table/index.html diff --git a/litellm/proxy/_experimental/out/models-and-endpoints.html b/litellm/proxy/_experimental/out/models-and-endpoints/index.html similarity index 100% rename from litellm/proxy/_experimental/out/models-and-endpoints.html rename to litellm/proxy/_experimental/out/models-and-endpoints/index.html diff --git a/litellm/proxy/_experimental/out/onboarding.html b/litellm/proxy/_experimental/out/onboarding/index.html similarity index 100% rename from litellm/proxy/_experimental/out/onboarding.html rename to litellm/proxy/_experimental/out/onboarding/index.html diff --git a/litellm/proxy/_experimental/out/organizations.html b/litellm/proxy/_experimental/out/organizations/index.html similarity index 100% rename from litellm/proxy/_experimental/out/organizations.html rename to litellm/proxy/_experimental/out/organizations/index.html diff --git a/litellm/proxy/_experimental/out/playground.html b/litellm/proxy/_experimental/out/playground/index.html similarity index 100% rename from litellm/proxy/_experimental/out/playground.html rename to litellm/proxy/_experimental/out/playground/index.html diff --git a/litellm/proxy/_experimental/out/policies.html b/litellm/proxy/_experimental/out/policies/index.html similarity index 100% rename from litellm/proxy/_experimental/out/policies.html rename to litellm/proxy/_experimental/out/policies/index.html diff --git a/litellm/proxy/_experimental/out/settings/admin-settings.html b/litellm/proxy/_experimental/out/settings/admin-settings/index.html similarity index 100% rename from litellm/proxy/_experimental/out/settings/admin-settings.html rename to litellm/proxy/_experimental/out/settings/admin-settings/index.html diff --git a/litellm/proxy/_experimental/out/settings/logging-and-alerts.html b/litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html similarity index 100% rename from litellm/proxy/_experimental/out/settings/logging-and-alerts.html rename to litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html diff --git a/litellm/proxy/_experimental/out/settings/router-settings.html b/litellm/proxy/_experimental/out/settings/router-settings/index.html similarity index 100% rename from litellm/proxy/_experimental/out/settings/router-settings.html rename to litellm/proxy/_experimental/out/settings/router-settings/index.html diff --git a/litellm/proxy/_experimental/out/settings/ui-theme.html b/litellm/proxy/_experimental/out/settings/ui-theme/index.html similarity index 100% rename from litellm/proxy/_experimental/out/settings/ui-theme.html rename to litellm/proxy/_experimental/out/settings/ui-theme/index.html diff --git a/litellm/proxy/_experimental/out/teams.html b/litellm/proxy/_experimental/out/teams/index.html similarity index 100% rename from litellm/proxy/_experimental/out/teams.html rename to litellm/proxy/_experimental/out/teams/index.html diff --git a/litellm/proxy/_experimental/out/test-key.html b/litellm/proxy/_experimental/out/test-key/index.html similarity index 100% rename from litellm/proxy/_experimental/out/test-key.html rename to litellm/proxy/_experimental/out/test-key/index.html diff --git a/litellm/proxy/_experimental/out/tools/mcp-servers.html b/litellm/proxy/_experimental/out/tools/mcp-servers/index.html similarity index 100% rename from litellm/proxy/_experimental/out/tools/mcp-servers.html rename to litellm/proxy/_experimental/out/tools/mcp-servers/index.html diff --git a/litellm/proxy/_experimental/out/tools/vector-stores.html b/litellm/proxy/_experimental/out/tools/vector-stores/index.html similarity index 100% rename from litellm/proxy/_experimental/out/tools/vector-stores.html rename to litellm/proxy/_experimental/out/tools/vector-stores/index.html diff --git a/litellm/proxy/_experimental/out/usage.html b/litellm/proxy/_experimental/out/usage/index.html similarity index 100% rename from litellm/proxy/_experimental/out/usage.html rename to litellm/proxy/_experimental/out/usage/index.html diff --git a/litellm/proxy/_experimental/out/users.html b/litellm/proxy/_experimental/out/users/index.html similarity index 100% rename from litellm/proxy/_experimental/out/users.html rename to litellm/proxy/_experimental/out/users/index.html diff --git a/litellm/proxy/_experimental/out/virtual-keys.html b/litellm/proxy/_experimental/out/virtual-keys/index.html similarity index 100% rename from litellm/proxy/_experimental/out/virtual-keys.html rename to litellm/proxy/_experimental/out/virtual-keys/index.html diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index f8fba5f5984..6800dff55ac 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py index 68f9dfd7abc..66b80c10f18 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_hooks/enkryptai/enkryptai.py b/litellm/proxy/guardrails/guardrail_hooks/enkryptai/enkryptai.py index 8e992297e5d..63541a1e2f9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/enkryptai/enkryptai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/enkryptai/enkryptai.py @@ -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", diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py index b37074e25e7..9018675d7a5 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py b/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py index 90f689ed23c..8955bffc125 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py +++ b/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/hiddenlayer.py b/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/hiddenlayer.py index e2c20604880..b907fbbcbda 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/hiddenlayer.py +++ b/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/hiddenlayer.py @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py index 083a407e9cf..263b6eee768 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py @@ -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", diff --git a/litellm/proxy/guardrails/guardrail_hooks/onyx/onyx.py b/litellm/proxy/guardrails/guardrail_hooks/onyx/onyx.py index 3598dbe741e..1cfc805dbf9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/onyx/onyx.py +++ b/litellm/proxy/guardrails/guardrail_hooks/onyx/onyx.py @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py b/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py index 030b6036815..a196937ef6c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py +++ b/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index 71ad9819146..d71b8449f94 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -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", diff --git a/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py b/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py index 5ebc7b96eb8..b3e761869b0 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py +++ b/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py b/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py index 87da11efad0..6486da7f714 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py +++ b/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_hooks/zscaler_ai_guard/zscaler_ai_guard.py b/litellm/proxy/guardrails/guardrail_hooks/zscaler_ai_guard/zscaler_ai_guard.py index c60752d7952..ff00cd73ca5 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/zscaler_ai_guard/zscaler_ai_guard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/zscaler_ai_guard/zscaler_ai_guard.py @@ -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: diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 6e49e4244e7..c05eda85c81 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -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: diff --git a/litellm/router.py b/litellm/router.py index 42058c79c17..d9de7e7fc5a 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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( diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 19b6b4cd772..35538ab1003 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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", diff --git a/tests/code_coverage_tests/check_guardrail_apply_decorator.py b/tests/code_coverage_tests/check_guardrail_apply_decorator.py new file mode 100644 index 00000000000..18a86277aa9 --- /dev/null +++ b/tests/code_coverage_tests/check_guardrail_apply_decorator.py @@ -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") diff --git a/tests/llm_translation/test_model_cost_map_resilience.py b/tests/llm_translation/test_model_cost_map_resilience.py new file mode 100644 index 00000000000..61e375eabeb --- /dev/null +++ b/tests/llm_translation/test_model_cost_map_resilience.py @@ -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 diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/test_litellm/integrations/test_opentelemetry.py index 95fa6ed8f60..3da9d9857c9 100644 --- a/tests/test_litellm/integrations/test_opentelemetry.py +++ b/tests/test_litellm/integrations/test_opentelemetry.py @@ -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 diff --git a/tests/test_litellm/llms/custom_httpx/test_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_http_handler.py index 65f08ef5021..c249bd9970c 100644 --- a/tests/test_litellm/llms/custom_httpx/test_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_http_handler.py @@ -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""" diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index ea5e15e548a..08205cd2d9d 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -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 diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 75ec806ee17..9dcb16b545e 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -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"] diff --git a/ui/litellm-dashboard/package-lock.json b/ui/litellm-dashboard/package-lock.json index 4205657ca8c..3a21813fbf4 100644 --- a/ui/litellm-dashboard/package-lock.json +++ b/ui/litellm-dashboard/package-lock.json @@ -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" + } } } } diff --git a/ui/litellm-dashboard/package.json b/ui/litellm-dashboard/package.json index 42b90d27333..164368eb6ba 100644 --- a/ui/litellm-dashboard/package.json +++ b/ui/litellm-dashboard/package.json @@ -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", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts index 4985206092f..2539cc63f95 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts @@ -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(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts index c57de675e0e..fe1afdcc39f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts @@ -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({ + 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), + }); +}; diff --git a/ui/litellm-dashboard/src/components/ModelSelect/PaginatedModelSelect/PaginatedModelSelect.test.tsx b/ui/litellm-dashboard/src/components/ModelSelect/PaginatedModelSelect/PaginatedModelSelect.test.tsx new file mode 100644 index 00000000000..b91f97885e6 --- /dev/null +++ b/ui/litellm-dashboard/src/components/ModelSelect/PaginatedModelSelect/PaginatedModelSelect.test.tsx @@ -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(); + + expect(screen.getByRole("combobox")).toBeInTheDocument(); + expect(screen.getByText("Select a model")).toBeInTheDocument(); + }); + + it("should display custom placeholder when provided", () => { + renderWithProviders( + , + ); + + expect(screen.getByText("Choose model")).toBeInTheDocument(); + }); + + it("should display model options when data is loaded", async () => { + renderWithProviders(); + + 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(); + + 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( + , + ); + + 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(); + + expect(screen.getByRole("combobox")).toHaveAttribute("aria-expanded", "false"); + }); + + it("should pass pageSize to useInfiniteModelInfo", () => { + renderWithProviders(); + + expect(mockUseInfiniteModelInfo).toHaveBeenCalledWith(25, undefined); + }); + + it("should pass search to useInfiniteModelInfo when user types", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + 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(); + + 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(); + + 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(); + + 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(); + + 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( + , + ); + + expect(screen.getByRole("combobox")).toBeInTheDocument(); + }); + + it("should respect disabled prop", () => { + renderWithProviders(); + + 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(); + + await userEvent.click(screen.getByRole("combobox")); + + await waitFor(() => { + expect(screen.getByRole("option", { name: "GPT-4 (model-1)" })).toBeInTheDocument(); + }); + + expect(mockFetchNextPage).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/ModelSelect/PaginatedModelSelect/PaginatedModelSelect.tsx b/ui/litellm-dashboard/src/components/ModelSelect/PaginatedModelSelect/PaginatedModelSelect.tsx new file mode 100644 index 00000000000..9b22fd1bd87 --- /dev/null +++ b/ui/litellm-dashboard/src/components/ModelSelect/PaginatedModelSelect/PaginatedModelSelect.tsx @@ -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(); + 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 ? ( + + + Model name: + {modelName} + + + Model ID: {modelId} + + + ) : ( + Model ID: {modelId} + )} + + ); + }; + + const handlePopupScroll = (e: UIEvent) => { + 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 ( +
@@ -632,6 +670,27 @@ const CustomCodeModal: React.FC = ({
+ {/* Contribution CTA Banner */} +
+
+
+ +
+
+
Built a useful guardrail?
+
Share it with the community and help others build faster
+
+
+ +
+ {/* Primitives Panel */} diff --git a/ui/litellm-dashboard/src/components/molecules/filter.tsx b/ui/litellm-dashboard/src/components/molecules/filter.tsx index a3a12fdf759..dcf22293f86 100644 --- a/ui/litellm-dashboard/src/components/molecules/filter.tsx +++ b/ui/litellm-dashboard/src/components/molecules/filter.tsx @@ -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>; options?: Array<{ label: string; value: string }>; + customComponent?: React.ComponentType; } interface FilterValues { @@ -194,6 +201,17 @@ const FilterComponent: React.FC = ({ ))} + ) : option.customComponent ? ( + (() => { + const CustomComponent = option.customComponent; + return ( + handleFilterChange(option.name, value ?? "")} + placeholder={`Select ${option.label || option.name}...`} + /> + ); + })() ) : ( | null; - budget_duration: string | null; - }; -} - -export interface TeamData { - team_id: string; - team_info: { - team_alias: string; - team_id: string; - organization_id: string | null; - admins: string[]; - members: string[]; - members_with_roles: Member[]; - metadata: Record; - tpm_limit: number | null; - rpm_limit: number | null; - max_budget: number | null; - soft_budget?: number | null; - budget_duration: string | null; - models: string[]; - blocked: boolean; - spend: number; - max_parallel_requests: number | null; - budget_reset_at: string | null; - model_id: string | null; - litellm_model_table: { - model_aliases: Record; - } | null; - created_at: string; - guardrails?: string[]; - policies?: string[]; - object_permission?: { - object_permission_id: string; - mcp_servers: string[]; - mcp_access_groups?: string[]; - mcp_tool_permissions?: Record; - vector_stores: string[]; - agents?: string[]; - agent_access_groups?: string[]; - }; - team_member_budget_table: { - max_budget: number; - budget_duration: string; - tpm_limit: number | null; - rpm_limit: number | null; - } | null; - }; - keys: any[]; - team_memberships: TeamMembership[]; -} - -export interface TeamInfoProps { - teamId: string; - onUpdate: (data: any) => void; - onClose: () => void; - accessToken: string | null; - is_team_admin: boolean; - is_proxy_admin: boolean; - userModels: string[]; - editTeam: boolean; - premiumUser?: boolean; -} - -const getOrganizationModels = (organization: Organization | null, userModels: string[]) => { - let tempModelsToPick = []; - - if (organization) { - // Check if organization has "all-proxy-models" in its models array - if (organization.models.includes("all-proxy-models")) { - // Treat as all-proxy-models (use userModels) - tempModelsToPick = userModels; - } else if (organization.models.length > 0) { - // Organization has specific models - tempModelsToPick = organization.models; - } else { - // Empty array [] is treated as all-proxy-models - tempModelsToPick = userModels; - } - } else { - // No organization, show all available models - tempModelsToPick = userModels; - } - - return unfurlWildcardModelsInList(tempModelsToPick, userModels); -}; - -const TeamInfoView: React.FC = ({ - teamId, - onClose, - accessToken, - is_team_admin, - is_proxy_admin, - userModels, - editTeam, - premiumUser = false, - onUpdate, -}) => { - const [teamData, setTeamData] = useState(null); - const [loading, setLoading] = useState(true); - const [isAddMemberModalVisible, setIsAddMemberModalVisible] = useState(false); - const [form] = Form.useForm(); - const [isEditMemberModalVisible, setIsEditMemberModalVisible] = useState(false); - const [selectedEditMember, setSelectedEditMember] = useState(null); - const [isEditing, setIsEditing] = useState(false); - const [mcpAccessGroups, setMcpAccessGroups] = useState([]); - const [mcpAccessGroupsLoaded, setMcpAccessGroupsLoaded] = useState(false); - const [copiedStates, setCopiedStates] = useState>({}); - const [guardrailsList, setGuardrailsList] = useState([]); - const [policiesList, setPoliciesList] = useState([]); - const [policyGuardrails, setPolicyGuardrails] = useState>({}); - const [loadingPolicies, setLoadingPolicies] = useState(false); - const [memberToDelete, setMemberToDelete] = useState(null); - const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false); - const [isDeleting, setIsDeleting] = useState(false); - const [isTeamSaving, setIsTeamSaving] = useState(false); - const [organization, setOrganization] = useState(null); - const { userRole } = useAuthorized(); - - const canEditTeam = is_team_admin || is_proxy_admin; - - const fetchTeamInfo = async () => { - try { - setLoading(true); - if (!accessToken) return; - const response = await teamInfoCall(accessToken, teamId); - setTeamData(response); - } catch (error) { - NotificationsManager.fromBackend("Failed to load team information"); - console.error("Error fetching team info:", error); - } finally { - setLoading(false); - } - }; - - useEffect(() => { - fetchTeamInfo(); - }, [teamId, accessToken]); - - // Fetch organization data when team has organization_id - useEffect(() => { - const fetchOrganization = async () => { - if (!accessToken || !teamData?.team_info?.organization_id) { - setOrganization(null); - return; - } - - try { - const orgData = await organizationInfoCall(accessToken, teamData.team_info.organization_id); - setOrganization(orgData); - } catch (error) { - console.error("Error fetching organization info:", error); - setOrganization(null); - } - }; - - fetchOrganization(); - }, [accessToken, teamData?.team_info?.organization_id]); - - // Compute modelsToPick based on organization and userModels - const modelsToPick = useMemo(() => { - return getOrganizationModels(organization, userModels); - }, [organization, userModels]); - - const fetchMcpAccessGroups = async () => { - if (!accessToken) return; - if (mcpAccessGroupsLoaded) return; - try { - const groups = await fetchMCPAccessGroups(accessToken); - setMcpAccessGroups(groups); - setMcpAccessGroupsLoaded(true); - } catch (error) { - console.error("Failed to fetch MCP access groups:", error); - } - }; - - useEffect(() => { - const fetchGuardrails = async () => { - try { - if (!accessToken) return; - const response = await getGuardrailsList(accessToken); - const guardrailNames = response.guardrails.map((g: { guardrail_name: string }) => g.guardrail_name); - setGuardrailsList(guardrailNames); - } catch (error) { - console.error("Failed to fetch guardrails:", error); - } - }; - - const fetchPolicies = async () => { - try { - if (!accessToken) return; - const response = await getPoliciesList(accessToken); - const policyNames = response.policies.map((p: { policy_name: string }) => p.policy_name); - setPoliciesList(policyNames); - } catch (error) { - console.error("Failed to fetch policies:", error); - } - }; - - fetchGuardrails(); - fetchPolicies(); - }, [accessToken]); - - // Fetch resolved guardrails for all policies - useEffect(() => { - const fetchPolicyGuardrails = async () => { - if (!accessToken || !teamData?.team_info?.policies || teamData.team_info.policies.length === 0) { - return; - } - - setLoadingPolicies(true); - const guardrailsMap: Record = {}; - - try { - await Promise.all( - teamData.team_info.policies.map(async (policyName: string) => { - try { - const policyInfo = await getPolicyInfoWithGuardrails(accessToken, policyName); - guardrailsMap[policyName] = policyInfo.resolved_guardrails || []; - } catch (error) { - console.error(`Failed to fetch guardrails for policy ${policyName}:`, error); - guardrailsMap[policyName] = []; - } - }) - ); - setPolicyGuardrails(guardrailsMap); - } catch (error) { - console.error("Failed to fetch policy guardrails:", error); - } finally { - setLoadingPolicies(false); - } - }; - - fetchPolicyGuardrails(); - }, [accessToken, teamData?.team_info?.policies]); - - const handleMemberCreate = async (values: any) => { - try { - if (accessToken == null) return; - - const member: Member = { - user_email: values.user_email, - user_id: values.user_id, - role: values.role, - }; - - await teamMemberAddCall(accessToken, teamId, member); - - NotificationsManager.success("Team member added successfully"); - setIsAddMemberModalVisible(false); - form.resetFields(); - - // Fetch updated team info - const updatedTeamData = await teamInfoCall(accessToken, teamId); - setTeamData(updatedTeamData); - - // Notify parent component of the update - onUpdate(updatedTeamData); - } catch (error: any) { - let errMsg = "Failed to add team member"; - - if (error?.raw?.detail?.error?.includes("Assigning team admins is a premium feature")) { - errMsg = "Assigning admins is an enterprise-only feature. Please upgrade your LiteLLM plan to enable this."; - } else if (error?.message) { - errMsg = error.message; - } - - NotificationsManager.fromBackend(errMsg); - console.error("Error adding team member:", error); - } - }; - - const handleMemberUpdate = async (values: any) => { - try { - if (accessToken == null) { - return; - } - - const member: Member = { - user_email: values.user_email, - user_id: values.user_id, - role: values.role, - max_budget_in_team: values.max_budget_in_team, - tpm_limit: values.tpm_limit, - rpm_limit: values.rpm_limit, - }; - console.log("Updating member with values:", member); - message.destroy(); // Remove all existing toasts - - await teamMemberUpdateCall(accessToken, teamId, member); - - NotificationsManager.success("Team member updated successfully"); - setIsEditMemberModalVisible(false); - - // Fetch updated team info - const updatedTeamData = await teamInfoCall(accessToken, teamId); - setTeamData(updatedTeamData); - - // Notify parent component of the update - onUpdate(updatedTeamData); - } catch (error: any) { - let errMsg = "Failed to update team member"; - if (error?.raw?.detail?.includes("Assigning team admins is a premium feature")) { - errMsg = "Assigning admins is an enterprise-only feature. Please upgrade your LiteLLM plan to enable this."; - } else if (error?.message) { - errMsg = error.message; - } - setIsEditMemberModalVisible(false); - - message.destroy(); // Remove all existing toasts - - NotificationsManager.fromBackend(errMsg); - console.error("Error updating team member:", error); - } - }; - - const handleMemberDelete = (member: Member) => { - setMemberToDelete(member); - setIsDeleteModalOpen(true); - }; - - const handleDeleteConfirm = async () => { - if (!memberToDelete || !accessToken) return; - - setIsDeleting(true); - try { - await teamMemberDeleteCall(accessToken, teamId, memberToDelete); - - NotificationsManager.success("Team member removed successfully"); - - // Fetch updated team info - const updatedTeamData = await teamInfoCall(accessToken, teamId); - setTeamData(updatedTeamData); - - // Notify parent component of the update - onUpdate(updatedTeamData); - } catch (error) { - NotificationsManager.fromBackend("Failed to remove team member"); - console.error("Error removing team member:", error); - } finally { - setIsDeleting(false); - setIsDeleteModalOpen(false); - setMemberToDelete(null); - } - }; - - const handleDeleteCancel = () => { - setIsDeleteModalOpen(false); - setMemberToDelete(null); - }; - - const handleTeamUpdate = async (values: any) => { - try { - if (!accessToken) return; - setIsTeamSaving(true); - - let parsedMetadata = {}; - try { - const rawMetadata = values.metadata ? JSON.parse(values.metadata) : {}; - // Exclude soft_budget_alerting_emails from parsed metadata since it's handled separately - const { soft_budget_alerting_emails, ...rest } = rawMetadata; - parsedMetadata = rest; - } catch (e) { - NotificationsManager.fromBackend("Invalid JSON in metadata field"); - return; - } - - let secretManagerSettings: Record | undefined; - if (typeof values.secret_manager_settings === "string") { - const trimmedSecretConfig = values.secret_manager_settings.trim(); - if (trimmedSecretConfig.length > 0) { - try { - secretManagerSettings = JSON.parse(values.secret_manager_settings); - } catch (e) { - NotificationsManager.fromBackend("Invalid JSON in secret manager settings"); - return; - } - } - } - - const sanitizeNumeric = (v: any) => { - if (v === null || v === undefined) return null; - if (typeof v === "string" && v.trim() === "") return null; - if (typeof v === "number" && Number.isNaN(v)) return null; - return v; - }; - - const updateData: any = { - team_id: teamId, - team_alias: values.team_alias, - models: values.models, - tpm_limit: sanitizeNumeric(values.tpm_limit), - rpm_limit: sanitizeNumeric(values.rpm_limit), - max_budget: values.max_budget, - soft_budget: sanitizeNumeric(values.soft_budget), - budget_duration: values.budget_duration, - metadata: { - ...parsedMetadata, - ...(values.guardrails?.length > 0 ? { guardrails: values.guardrails } : {}), - ...(values.logging_settings?.length > 0 ? { logging: values.logging_settings } : {}), - disable_global_guardrails: values.disable_global_guardrails || false, - soft_budget_alerting_emails: - typeof values.soft_budget_alerting_emails === "string" - ? values.soft_budget_alerting_emails - .split(",") - .map((email: string) => email.trim()) - .filter((email: string) => email.length > 0) - : values.soft_budget_alerting_emails || [], - ...(secretManagerSettings !== undefined ? { secret_manager_settings: secretManagerSettings } : {}), - }, - ...(values.policies?.length > 0 ? { policies: values.policies } : {}), - organization_id: values.organization_id, - }; - - updateData.max_budget = mapEmptyStringToNull(updateData.max_budget); - updateData.team_member_budget_duration = values.team_member_budget_duration; - - if (values.team_member_budget !== undefined) { - updateData.team_member_budget = Number(values.team_member_budget); - } - - if (values.team_member_key_duration !== undefined) { - updateData.team_member_key_duration = values.team_member_key_duration; - } - - if (values.team_member_tpm_limit !== undefined || values.team_member_rpm_limit !== undefined) { - updateData.team_member_tpm_limit = sanitizeNumeric(values.team_member_tpm_limit); - updateData.team_member_rpm_limit = sanitizeNumeric(values.team_member_rpm_limit); - } - - // Handle object_permission updates - const { servers, accessGroups } = values.mcp_servers_and_groups || { - servers: [], - accessGroups: [], - }; - const serverIds = new Set(servers || []); - const mcpToolPermissions = Object.fromEntries( - Object.entries(values.mcp_tool_permissions || {}).filter(([serverId]) => serverIds.has(serverId)), - ); - - updateData.object_permission = {}; - if (servers) { - updateData.object_permission.mcp_servers = servers; - } - if (accessGroups) { - updateData.object_permission.mcp_access_groups = accessGroups; - } - if (mcpToolPermissions) { - updateData.object_permission.mcp_tool_permissions = mcpToolPermissions; - } - delete values.mcp_servers_and_groups; - delete values.mcp_tool_permissions; - - // Handle agent permissions - const { agents, accessGroups: agentAccessGroups } = values.agents_and_groups || { - agents: [], - accessGroups: [], - }; - if (agents && agents.length > 0) { - updateData.object_permission.agents = agents; - } - if (agentAccessGroups && agentAccessGroups.length > 0) { - updateData.object_permission.agent_access_groups = agentAccessGroups; - } - delete values.agents_and_groups; - - // Handle vector stores permissions - if (values.vector_stores && values.vector_stores.length > 0) { - updateData.object_permission.vector_stores = values.vector_stores; - } - - const response = await teamUpdateCall(accessToken, updateData); - - NotificationsManager.success("Team settings updated successfully"); - setIsEditing(false); - fetchTeamInfo(); - } catch (error) { - console.error("Error updating team:", error); - } finally { - setIsTeamSaving(false); - } - }; - - if (loading) { - return
Loading...
; - } - - if (!teamData?.team_info) { - return
Team not found
; - } - - const { team_info: info } = teamData; - - const copyToClipboard = async (text: string, key: string) => { - const success = await utilCopyToClipboard(text); - if (success) { - setCopiedStates((prev) => ({ ...prev, [key]: true })); - setTimeout(() => { - setCopiedStates((prev) => ({ ...prev, [key]: false })); - }, 2000); - } - }; - - return ( -
-
-
- - Back to Teams - - {info.team_alias} -
- {info.team_id} -
-
-
- - - - {[ - Overview, - ...(canEditTeam - ? [ - Members, - Member Permissions, - Settings, - ] - : []), - ]} - - - - {/* Overview Panel */} - - - - Budget Status -
- ${formatNumberWithCommas(info.spend, 4)} - - of {info.max_budget === null ? "Unlimited" : `$${formatNumberWithCommas(info.max_budget, 4)}`} - - {info.budget_duration && Reset: {info.budget_duration}} -
- {info.team_member_budget_table && ( - - Team Member Budget: ${formatNumberWithCommas(info.team_member_budget_table.max_budget, 4)} - - )} -
-
- - - Rate Limits -
- TPM: {info.tpm_limit || "Unlimited"} - RPM: {info.rpm_limit || "Unlimited"} - {info.max_parallel_requests && Max Parallel Requests: {info.max_parallel_requests}} -
-
- - - Models -
- {info.models.length === 0 ? ( - All proxy models - ) : ( - info.models.map((model, index) => ( - - {model} - - )) - )} -
-
- - - Virtual Keys -
- User Keys: {teamData.keys.filter((key) => key.user_id).length} - Service Account Keys: {teamData.keys.filter((key) => !key.user_id).length} - Total: {teamData.keys.length} -
-
- - - - - Guardrails - {info.guardrails && info.guardrails.length > 0 ? ( -
- {info.guardrails.map((guardrail: string, index: number) => ( - - {guardrail} - - ))} -
- ) : ( - No guardrails configured - )} - {info.metadata?.disable_global_guardrails && ( -
- Global Guardrails Disabled -
- )} -
- - - Policies - {info.policies && info.policies.length > 0 ? ( -
- {info.policies.map((policy: string, index: number) => ( -
-
- {policy} - {loadingPolicies && Loading guardrails...} -
- {!loadingPolicies && policyGuardrails[policy] && policyGuardrails[policy].length > 0 && ( -
- Resolved Guardrails: -
- {policyGuardrails[policy].map((guardrail: string, gIndex: number) => ( - - {guardrail} - - ))} -
-
- )} -
- ))} -
- ) : ( - No policies configured - )} -
- - -
-
- - {/* Members Panel */} - - - - - {/* Member Permissions Panel */} - {canEditTeam && ( - - - - )} - - {/* Settings Panel */} - - -
- Team Settings - {canEditTeam && !isEditing && ( - setIsEditing(true)}>Edit Settings - )} -
- - {isEditing ? ( -
rest)(info.metadata), - null, - 2, - ) - : "", - logging_settings: info.metadata?.logging || [], - secret_manager_settings: info.metadata?.secret_manager_settings - ? JSON.stringify(info.metadata.secret_manager_settings, null, 2) - : "", - organization_id: info.organization_id, - vector_stores: info.object_permission?.vector_stores || [], - mcp_servers: info.object_permission?.mcp_servers || [], - mcp_access_groups: info.object_permission?.mcp_access_groups || [], - mcp_servers_and_groups: { - servers: info.object_permission?.mcp_servers || [], - accessGroups: info.object_permission?.mcp_access_groups || [], - }, - mcp_tool_permissions: info.object_permission?.mcp_tool_permissions || {}, - agents_and_groups: { - agents: info.object_permission?.agents || [], - accessGroups: info.object_permission?.agent_access_groups || [], - }, - }} - layout="vertical" - > - - - - - - form.setFieldValue("models", values)} - teamID={teamId} - organizationID={teamData?.team_info?.organization_id || undefined} - options={{ - includeSpecialOptions: true, - includeUserModels: !teamData?.team_info?.organization_id, - showAllProxyModelsOverride: isProxyAdminRole(userRole) && !teamData?.team_info?.organization_id, - }} - context="team" - dataTestId="models-select" - /> - - - - - - - - - - - - - - - - - - - - form.setFieldValue("team_member_budget_duration", value)} - value={form.getFieldValue("team_member_budget_duration")} - /> - - - - - - - - - - - - - - - - - - - - - - - - - - - - Guardrails{" "} - - e.stopPropagation()} - > - - - - - } - name="guardrails" - help="Select existing guardrails or enter new ones" - > - ({ value: name, label: name }))} - /> - - - - form.setFieldValue("vector_stores", values)} - value={form.getFieldValue("vector_stores")} - accessToken={accessToken || ""} - placeholder="Select vector stores" - /> - - - - form.setFieldValue("allowed_passthrough_routes", values)} - value={form.getFieldValue("allowed_passthrough_routes")} - accessToken={accessToken || ""} - placeholder="Select pass through routes" - /> - - - - form.setFieldValue("mcp_servers_and_groups", val)} - value={form.getFieldValue("mcp_servers_and_groups")} - accessToken={accessToken || ""} - placeholder="Select MCP servers or access groups (optional)" - /> - - - {/* Hidden field to register mcp_tool_permissions with the form */} - - - - prevValues.mcp_servers_and_groups !== currentValues.mcp_servers_and_groups || - prevValues.mcp_tool_permissions !== currentValues.mcp_tool_permissions - } - > - {() => ( -
- form.setFieldsValue({ mcp_tool_permissions: toolPerms })} - /> -
- )} -
- - - form.setFieldValue("agents_and_groups", val)} - value={form.getFieldValue("agents_and_groups")} - accessToken={accessToken || ""} - placeholder="Select agents or access groups (optional)" - /> - - - - - - - - form.setFieldValue("logging_settings", values)} - /> - - - { - if (!value) { - return Promise.resolve(); - } - try { - JSON.parse(value); - return Promise.resolve(); - } catch (error) { - return Promise.reject(new Error("Please enter valid JSON")); - } - }, - }, - ]} - > - - - - - - - -
-
- setIsEditing(false)} disabled={isTeamSaving}> - Cancel - - - Save Changes - -
-
-
- ) : ( -
-
- Team Name -
{info.team_alias}
-
-
- Team ID -
{info.team_id}
-
-
- Created At -
{new Date(info.created_at).toLocaleString()}
-
-
- Models -
- {info.models.map((model, index) => ( - - {model} - - ))} -
-
-
- Rate Limits -
TPM: {info.tpm_limit || "Unlimited"}
-
RPM: {info.rpm_limit || "Unlimited"}
-
-
- Team Budget -
- Max Budget:{" "} - {info.max_budget !== null ? `$${formatNumberWithCommas(info.max_budget, 4)}` : "No Limit"} -
-
- Soft Budget:{" "} - {info.soft_budget !== null && info.soft_budget !== undefined - ? `$${formatNumberWithCommas(info.soft_budget, 4)}` - : "No Limit"} -
-
Budget Reset: {info.budget_duration || "Never"}
- {info.metadata?.soft_budget_alerting_emails && - Array.isArray(info.metadata.soft_budget_alerting_emails) && - info.metadata.soft_budget_alerting_emails.length > 0 && ( -
- Soft Budget Alerting Emails: {info.metadata.soft_budget_alerting_emails.join(", ")} -
- )} -
-
- - Team Member Settings{" "} - - - - -
Max Budget: {info.team_member_budget_table?.max_budget || "No Limit"}
-
Budget Duration: {info.team_member_budget_table?.budget_duration || "No Limit"}
-
Key Duration: {info.metadata?.team_member_key_duration || "No Limit"}
-
TPM Limit: {info.team_member_budget_table?.tpm_limit || "No Limit"}
-
RPM Limit: {info.team_member_budget_table?.rpm_limit || "No Limit"}
-
-
- Organization ID -
{info.organization_id}
-
-
- Status - {info.blocked ? "Blocked" : "Active"} -
- -
- Disable Global Guardrails -
- {info.metadata?.disable_global_guardrails === true ? ( - Enabled - Global guardrails bypassed - ) : ( - Disabled - Global guardrails active - )} -
-
- - - - - - {info.metadata?.secret_manager_settings && ( -
- Secret Manager Settings -
-                        {JSON.stringify(info.metadata.secret_manager_settings, null, 2)}
-                      
-
- )} -
- )} -
-
-
-
- - setIsEditMemberModalVisible(false)} - onSubmit={handleMemberUpdate} - initialData={selectedEditMember} - mode="edit" - config={{ - title: "Edit Member", - showEmail: true, - showUserId: true, - roleOptions: [ - { label: "Admin", value: "admin" }, - { label: "User", value: "user" }, - ], - additionalFields: [ - { - name: "max_budget_in_team", - label: ( - - Team Member Budget (USD){" "} - - - - - ), - type: "numerical" as const, - step: 0.01, - min: 0, - placeholder: "Budget limit for this member within this team", - }, - { - name: "tpm_limit", - label: ( - - Team Member TPM Limit{" "} - - - - - ), - type: "numerical" as const, - step: 1, - min: 0, - placeholder: "Tokens per minute limit for this member in this team", - }, - { - name: "rpm_limit", - label: ( - - Team Member RPM Limit{" "} - - - - - ), - type: "numerical" as const, - step: 1, - min: 0, - placeholder: "Requests per minute limit for this member in this team", - }, - ], - }} - /> - - setIsAddMemberModalVisible(false)} - onSubmit={handleMemberCreate} - accessToken={accessToken} - /> - - {/* Delete Member Confirmation Modal */} - -
- ); -}; - -export default TeamInfoView; diff --git a/ui/litellm-dashboard/src/components/view_logs/index.tsx b/ui/litellm-dashboard/src/components/view_logs/index.tsx index b791663cee3..b43ff5220db 100644 --- a/ui/litellm-dashboard/src/components/view_logs/index.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/index.tsx @@ -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", diff --git a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx index 628a2ecfd57..85eea0a8272 100644 --- a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx @@ -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, diff --git a/ui/litellm-dashboard/tsconfig.json b/ui/litellm-dashboard/tsconfig.json index d24bdd340f7..5b0352feb98 100644 --- a/ui/litellm-dashboard/tsconfig.json +++ b/ui/litellm-dashboard/tsconfig.json @@ -14,7 +14,7 @@ "moduleResolution": "bundler", "resolveJsonModule": true, "isolatedModules": true, - "jsx": "react-jsx", + "jsx": "preserve", "incremental": true, "plugins": [ {