diff --git a/.github/workflows/check-lazy-openapi-snapshot.yml b/.github/workflows/check-lazy-openapi-snapshot.yml deleted file mode 100644 index 2e4ed3637f1..00000000000 --- a/.github/workflows/check-lazy-openapi-snapshot.yml +++ /dev/null @@ -1,75 +0,0 @@ -name: Check Lazy OpenAPI Snapshot - -on: - pull_request: - branches: - - main - - litellm_internal_staging - - "litellm_**" - -permissions: - contents: read - checks: write - -concurrency: - group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }} - cancel-in-progress: true - -jobs: - verify: - runs-on: ubuntu-latest - timeout-minutes: 10 - steps: - - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 - with: - persist-credentials: false - - - name: Set up Python - uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 - with: - python-version: "3.12" - - - name: Set up uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - version: "0.10.9" - - - name: Cache uv dependencies - uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0 - with: - path: | - ~/.cache/uv - .venv - key: ${{ runner.os }}-uv-${{ hashFiles('uv.lock') }} - restore-keys: | - ${{ runner.os }}-uv- - - - name: Install dependencies - run: uv sync --frozen --all-groups --all-extras - - - name: Regenerate snapshot to /tmp - id: regen - run: | - cp litellm/proxy/_lazy_openapi_snapshot.json /tmp/snapshot.committed.json - uv run --no-sync python -m litellm.proxy._lazy_openapi_snapshot - mv litellm/proxy/_lazy_openapi_snapshot.json /tmp/snapshot.fresh.json - mv /tmp/snapshot.committed.json litellm/proxy/_lazy_openapi_snapshot.json - - - name: Compare - id: diff - continue-on-error: true - run: | - diff -q /tmp/snapshot.fresh.json litellm/proxy/_lazy_openapi_snapshot.json - - - name: Mark neutral if drift - if: steps.diff.outcome == 'failure' - uses: LouisBrunner/checks-action@6b626ffbad7cc56fd58627f774b9067e6118af23 # v2.0.0 - with: - token: ${{ secrets.GITHUB_TOKEN }} - name: lazy-openapi-snapshot - conclusion: neutral - output: | - { - "title": "Lazy openapi snapshot is stale", - "summary": "Run `python -m litellm.proxy._lazy_openapi_snapshot` and commit the regenerated `litellm/proxy/_lazy_openapi_snapshot.json`. Not blocking — the snapshot will regenerate at release if not committed." - } diff --git a/.gitignore b/.gitignore index 38bf9554b5b..59812ed6ed4 100644 --- a/.gitignore +++ b/.gitignore @@ -90,7 +90,6 @@ test.py litellm_config.yaml !.github/observatory/litellm_config.yaml .cursor -.vscode/launch.json litellm/proxy/to_delete_loadtest_work/* update_model_cost_map.py tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -100,4 +99,5 @@ STABILIZATION_TODO.md **/test-results **/playwright-report **/*.storageState.json -**/coverage \ No newline at end of file +**/coverage +test-config \ No newline at end of file diff --git a/Makefile b/Makefile index b6b674ff3b1..5dbd308a3e2 100644 --- a/Makefile +++ b/Makefile @@ -185,3 +185,6 @@ test-llm-translation-single: install-test-deps $(UV_RUN) pytest tests/llm_translation/$(FILE) \ --junitxml=test-results/junit.xml \ -v --tb=short --maxfail=100 --timeout=300 + +test-llm-translation-flush-vcr-cache: + $(UV_RUN) python tests/_flush_vcr_cache.py diff --git a/README.md b/README.md index d72fb746ed4..72fd43925c9 100644 --- a/README.md +++ b/README.md @@ -68,7 +68,7 @@ Managing LLM calls across providers gets complicated fast — different SDKs, au Stripe image Google ADK - Greptile + Greptile OpenHands

Netflix

OpenAI Agents SDK diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py index f6ed7767c46..4bfe9d31874 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py @@ -857,10 +857,16 @@ async def project_info( where={"team_id": project.team_id} ) if team: - is_team_member = ( - user_api_key_dict.user_id in team.admins - or user_api_key_dict.user_id in team.members - ) + caller_user_id = user_api_key_dict.user_id + for m in team.members_with_roles or []: + m_user_id = ( + m.get("user_id") + if isinstance(m, dict) + else getattr(m, "user_id", None) + ) + if m_user_id == caller_user_id: + is_team_member = True + break if not (is_admin or is_team_member): raise HTTPException( @@ -911,20 +917,20 @@ async def list_projects( include={"litellm_budget_table": True, "object_permission": True} ) else: - # Get projects for teams the user belongs to - user_teams = await prisma_client.db.litellm_teamtable.find_many( - where={ - "OR": [ - {"members": {"has": user_api_key_dict.user_id}}, - {"admins": {"has": user_api_key_dict.user_id}}, - ] - } + # Look up the user's team memberships via the reverse-index on + # LiteLLM_UserTable.teams (maintained by team_member_add alongside + # members_with_roles). This avoids a full scan of all team rows. + user_record = await prisma_client.db.litellm_usertable.find_unique( + where={"user_id": user_api_key_dict.user_id}, + ) + user_team_ids = ( + user_record.teams + if user_record is not None and user_record.teams + else [] ) - team_ids = [team.team_id for team in user_teams] - projects = await prisma_client.db.litellm_projecttable.find_many( - where={"team_id": {"in": team_ids}}, + where={"team_id": {"in": user_team_ids}}, include={"litellm_budget_table": True, "object_permission": True}, ) diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index ce1bc26c5e0..11733ce4cee 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -432,9 +432,10 @@ class Cache: str: The final hashed cache key with the redis namespace. """ dynamic_cache_control: DynamicCacheControl = kwargs.get("cache", {}) + metadata = kwargs.get("metadata") or {} namespace = ( dynamic_cache_control.get("namespace") - or kwargs.get("metadata", {}).get("redis_namespace") + or metadata.get("redis_namespace") or self.namespace ) if namespace: diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index 7d514e648fe..3cf1d911d7f 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -87,6 +87,18 @@ class CachingHandlerResponse(BaseModel): in_memory_cache_obj = InMemoryCache() +def _should_defer_streaming_cache_hit_callbacks(*, kwargs: Dict[str, Any]) -> bool: + """ + When stream=True, do not run success callbacks at cache-hit time. + + Cached chat/text completion replay uses CustomStreamWrapper; cached Responses + replay uses CachedResponsesAPIStreamingIterator. Both invoke logging success + handlers when the stream finishes; firing them here too would double-count + spend and callback records. + """ + return kwargs.get("stream", False) is True + + class LLMCachingHandler: def __init__( self, @@ -99,6 +111,7 @@ class LLMCachingHandler: self.async_streaming_chunks: List[ModelResponse] = [] self.sync_streaming_chunks: List[ModelResponse] = [] self.request_kwargs = request_kwargs + self.preset_cache_key: Optional[str] = None self.original_function = original_function self.start_time = start_time if litellm.cache is not None and isinstance(litellm.cache.cache, RedisCache): @@ -206,7 +219,7 @@ class LLMCachingHandler: custom_llm_provider=kwargs.get("custom_llm_provider", None), args=args, ) - if kwargs.get("stream", False) is False: + if not _should_defer_streaming_cache_hit_callbacks(kwargs=kwargs): # LOG SUCCESS self._async_log_cache_hit_on_callbacks( logging_obj=logging_obj, @@ -215,11 +228,12 @@ class LLMCachingHandler: end_time=end_time, cache_hit=cache_hit, ) - cache_key = litellm.cache.get_cache_key(**kwargs) - if ( - isinstance(cached_result, BaseModel) - or isinstance(cached_result, CustomStreamWrapper) - ) and hasattr(cached_result, "_hidden_params"): + cache_key = ( + self.preset_cache_key + or self.request_kwargs.get("cache_key") + or litellm.cache.get_cache_key(**self.request_kwargs) + ) + if hasattr(cached_result, "_hidden_params"): cached_result._hidden_params["cache_key"] = cache_key # type: ignore return CachingHandlerResponse(cached_result=cached_result) elif ( @@ -265,8 +279,6 @@ class LLMCachingHandler: kwargs: Dict[str, Any], args: Optional[Tuple[Any, ...]] = None, ) -> CachingHandlerResponse: - from litellm.utils import CustomStreamWrapper - cached_result: Optional[Any] = None # Check if caching should be performed BEFORE doing expensive kwargs copy @@ -282,6 +294,11 @@ class LLMCachingHandler: args, ) ) + if new_kwargs.get("metadata") is None: + new_kwargs.pop("metadata", None) + if new_kwargs.get("stream") is True and "cache_key" not in new_kwargs: + new_kwargs["cache_key"] = litellm.cache.get_cache_key(**new_kwargs) + self.request_kwargs = new_kwargs print_verbose("Checking Sync Cache") cached_result = litellm.cache.get_cache(**new_kwargs) if cached_result is not None: @@ -322,17 +339,19 @@ class LLMCachingHandler: is_async=False, ) - logging_obj.handle_sync_success_callbacks_for_async_calls( - result=cached_result, - start_time=start_time, - end_time=end_time, - cache_hit=cache_hit, + if not _should_defer_streaming_cache_hit_callbacks(kwargs=kwargs): + logging_obj.handle_sync_success_callbacks_for_async_calls( + result=cached_result, + start_time=start_time, + end_time=end_time, + cache_hit=cache_hit, + ) + cache_key = ( + self.preset_cache_key + or self.request_kwargs.get("cache_key") + or litellm.cache.get_cache_key(**self.request_kwargs) ) - cache_key = litellm.cache.get_cache_key(**kwargs) - if ( - isinstance(cached_result, BaseModel) - or isinstance(cached_result, CustomStreamWrapper) - ) and hasattr(cached_result, "_hidden_params"): + if hasattr(cached_result, "_hidden_params"): cached_result._hidden_params["cache_key"] = cache_key # type: ignore return CachingHandlerResponse(cached_result=cached_result) return CachingHandlerResponse(cached_result=cached_result) @@ -686,6 +705,11 @@ class LLMCachingHandler: args, ) ) + if new_kwargs.get("metadata") is None: + new_kwargs.pop("metadata", None) + if new_kwargs.get("stream") is True and "cache_key" not in new_kwargs: + new_kwargs["cache_key"] = litellm.cache.get_cache_key(**new_kwargs) + self.request_kwargs = new_kwargs cached_result: Optional[Any] = None if call_type == CallTypes.aembedding.value: if isinstance(new_kwargs["input"], str): @@ -710,14 +734,26 @@ class LLMCachingHandler: if all(result is None for result in cached_result): cached_result = None else: + request_kwargs = new_kwargs.copy() + request_cache_key = request_kwargs.pop("cache_key", None) if litellm.cache._supports_async() is True: ## check if dual cache is supported ## + self.preset_cache_key = ( + request_cache_key or litellm.cache.get_cache_key(**request_kwargs) + ) cached_result = await litellm.cache.async_get_cache( - dynamic_cache_object=self.dual_cache, **new_kwargs + dynamic_cache_object=self.dual_cache, + cache_key=self.preset_cache_key, + **request_kwargs, ) else: # fallback for caches that don't support async + self.preset_cache_key = ( + request_cache_key or litellm.cache.get_cache_key(**request_kwargs) + ) cached_result = litellm.cache.get_cache( - dynamic_cache_object=self.dual_cache, **new_kwargs + dynamic_cache_object=self.dual_cache, + cache_key=self.preset_cache_key, + **request_kwargs, ) return cached_result @@ -825,8 +861,27 @@ class LLMCachingHandler: elif (call_type == "aresponses" or call_type == "responses") and isinstance( cached_result, dict ): - # Convert cached dict back to ResponsesAPIResponse object - cached_result = ResponsesAPIResponse(**cached_result) + from litellm.responses.streaming_iterator import ( + CachedResponsesAPIStreamingIterator, + ) + + response_obj = ResponsesAPIResponse(**cached_result) + if ( + hasattr(response_obj, "_hidden_params") + and response_obj._hidden_params is not None + and isinstance(response_obj._hidden_params, dict) + ): + response_obj._hidden_params["cache_hit"] = True + + if kwargs.get("stream", False) is True: + cached_result = CachedResponsesAPIStreamingIterator( + response=response_obj, + logging_obj=logging_obj, + request_data=kwargs, + call_type=call_type, + ) + else: + cached_result = response_obj if ( hasattr(cached_result, "_hidden_params") diff --git a/litellm/caching/dual_cache.py b/litellm/caching/dual_cache.py index 6115a444cee..8060a65b78d 100644 --- a/litellm/caching/dual_cache.py +++ b/litellm/caching/dual_cache.py @@ -92,6 +92,25 @@ class DualCache(BaseCache): if default_redis_ttl is not None: self.default_redis_ttl = default_redis_ttl + def attach_redis_cache( + self, + redis_cache: Optional[RedisCache] = None, + *, + default_redis_ttl: Optional[float] = None, + ) -> None: + """ + Attach a Redis backend if this DualCache does not already have one. + + No-op when ``redis_cache`` is None or when Redis was already set (constructor + or a prior attach). Use this for lazy wiring after a shared Redis client exists. + Does not backfill in-memory-only keys to Redis. + """ + if redis_cache is None or self.redis_cache is not None: + return + self.redis_cache = redis_cache + if default_redis_ttl is not None: + self.default_redis_ttl = default_redis_ttl + def set_cache(self, key, value, local_only: bool = False, **kwargs): # Update both Redis and in-memory cache try: diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index deee4f6ea48..cb9ce475d30 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -551,6 +551,13 @@ class RedisCache(BaseCache): async def async_set_cache(self, key, value, **kwargs): from redis.asyncio import Redis + if key is None: + verbose_logger.debug( + "LiteLLM Redis Caching: async set() skipped — key is None, value=%r", + value, + ) + return None + start_time = time.time() try: _redis_client: Redis = self.init_async_client() # type: ignore @@ -569,8 +576,9 @@ class RedisCache(BaseCache): ) ) verbose_logger.error( - "LiteLLM Redis Caching: async set() - Got exception from REDIS %s, Writing value=%s", + "LiteLLM Redis Caching: async set() - Got exception from REDIS %s, key=%r, value=%r", str(e), + key, value, ) raise e diff --git a/litellm/integrations/arize/arize_phoenix_client.py b/litellm/integrations/arize/arize_phoenix_client.py index 3c83517bb55..8c3c2a5ff0f 100644 --- a/litellm/integrations/arize/arize_phoenix_client.py +++ b/litellm/integrations/arize/arize_phoenix_client.py @@ -2,11 +2,23 @@ Arize Phoenix API client for fetching prompt versions from Arize Phoenix. """ +import urllib.parse from typing import Any, Dict, Optional from litellm.llms.custom_httpx.http_handler import HTTPHandler +def _sanitize_id(identifier: str) -> str: + """Reject path traversal characters and URL-encode the identifier.""" + if any(c in identifier for c in ("/", "\\", "#", "?")): + raise ValueError( + f"Invalid identifier {identifier!r}: contains disallowed characters" + ) + if ".." in identifier: + raise ValueError(f"Invalid identifier {identifier!r}: path traversal detected") + return urllib.parse.quote(identifier, safe="") + + class ArizePhoenixClient: """ Client for interacting with Arize Phoenix API to fetch prompt versions. @@ -53,7 +65,8 @@ class ArizePhoenixClient: Returns: Dictionary containing prompt version data, or None if not found """ - url = f"{self.api_base}/v1/prompt_versions/{prompt_version_id}" + safe_id = _sanitize_id(prompt_version_id) + url = f"{self.api_base}/v1/prompt_versions/{safe_id}" try: # Use the underlying httpx client directly to avoid query param extraction diff --git a/litellm/integrations/bitbucket/bitbucket_client.py b/litellm/integrations/bitbucket/bitbucket_client.py index 0502422cf8b..e742cc14b7d 100644 --- a/litellm/integrations/bitbucket/bitbucket_client.py +++ b/litellm/integrations/bitbucket/bitbucket_client.py @@ -3,11 +3,27 @@ BitBucket API client for fetching .prompt files from BitBucket repositories. """ import base64 +import urllib.parse from typing import Any, Dict, List, Optional from litellm.llms.custom_httpx.http_handler import HTTPHandler +def _sanitize_file_path(file_path: str) -> str: + """Reject path traversal and URL-encode each path segment.""" + if "#" in file_path or "?" in file_path: + raise ValueError( + f"Invalid file path {file_path!r}: contains URL special characters" + ) + parts = file_path.split("/") + for part in parts: + if part == "..": + raise ValueError( + f"Invalid file path {file_path!r}: path traversal detected" + ) + return "/".join(urllib.parse.quote(part, safe="") for part in parts) + + class BitBucketClient: """ Client for interacting with BitBucket API to fetch .prompt files. @@ -72,7 +88,8 @@ class BitBucketClient: Returns: File content as string, or None if file not found """ - url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{file_path}" + safe_path = _sanitize_file_path(file_path) + url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{safe_path}" try: response = self.http_handler.get(url, headers=self.headers) @@ -119,7 +136,8 @@ class BitBucketClient: Returns: List of file paths """ - url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{directory_path}" + safe_dir = _sanitize_file_path(directory_path) if directory_path else "" + url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{safe_dir}" try: response = self.http_handler.get(url, headers=self.headers) @@ -211,7 +229,8 @@ class BitBucketClient: Returns: Dictionary containing file metadata, or None if file not found """ - url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{file_path}" + safe_path = _sanitize_file_path(file_path) + url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{safe_path}" try: # Use GET with Range header to get just the headers (HEAD equivalent) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 723b142dfad..d9e57ee7cee 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -265,6 +265,7 @@ class PrometheusLogger(CustomLogger): ######################################## # LiteLLM Virtual API KEY metrics ######################################## + # Remaining MODEL RPM limit for API Key self.litellm_remaining_api_key_requests_for_model = self._gauge_factory( "litellm_remaining_api_key_requests_for_model", diff --git a/litellm/litellm_core_utils/cli_token_utils.py b/litellm/litellm_core_utils/cli_token_utils.py index e2e304931a4..3776d276912 100644 --- a/litellm/litellm_core_utils/cli_token_utils.py +++ b/litellm/litellm_core_utils/cli_token_utils.py @@ -31,15 +31,23 @@ def load_cli_token() -> Optional[dict]: return None -def get_litellm_gateway_api_key() -> Optional[str]: +def get_litellm_gateway_api_key( + expected_base_url: Optional[str] = None, +) -> Optional[str]: """ Get the stored CLI API key for use with LiteLLM SDK. This function reads the token file created by `litellm-proxy login` and returns the API key for use in Python scripts. + Args: + expected_base_url: When provided, the key is only returned if it was + originally issued for this URL. Pass the target server URL to + prevent credential leakage when the client is pointed at a + different (possibly malicious) server. + Returns: - str: The API key if found, None otherwise + str: The API key if found (and origin matches), None otherwise Example: >>> import litellm @@ -53,6 +61,10 @@ def get_litellm_gateway_api_key() -> Optional[str]: >>> ) """ token_data = load_cli_token() - if token_data and "key" in token_data: - return token_data["key"] - return None + if not token_data or "key" not in token_data: + return None + if expected_base_url is not None: + stored_url = token_data.get("base_url") + if stored_url != expected_base_url.rstrip("/"): + return None + return token_data["key"] diff --git a/litellm/litellm_core_utils/llm_request_utils.py b/litellm/litellm_core_utils/llm_request_utils.py index f5f28822ca1..7be70852978 100644 --- a/litellm/litellm_core_utils/llm_request_utils.py +++ b/litellm/litellm_core_utils/llm_request_utils.py @@ -77,8 +77,8 @@ def get_proxy_server_request_headers(litellm_params: Optional[dict]) -> dict: if litellm_params is None: return {} - proxy_request_headers = ( - litellm_params.get("proxy_server_request", {}).get("headers", {}) or {} - ) + proxy_request_headers = (litellm_params.get("proxy_server_request") or {}).get( + "headers" + ) or {} return proxy_request_headers diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 3a83162fb20..ba840bc3d89 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -4582,6 +4582,11 @@ class BedrockConverseMessagesProcessor: message=cast(ChatCompletionFileObject, element) ) _parts.append(_part) + elif element["type"] == "document": + _part = BedrockConverseMessagesProcessor._process_document_message( + element + ) + _parts.append(_part) _cache_point_block = ( litellm.AmazonConverseConfig()._get_cache_point_block( message_block=cast( @@ -4864,6 +4869,44 @@ class BedrockConverseMessagesProcessor: image_url=cast(str, file_id or file_data), format=format ) + @staticmethod + def _process_document_message(element: dict) -> BedrockContentBlock: + """Convert a document content block to a Bedrock DocumentBlock. + + Handles the Anthropic-style document format: + {"type": "document", "source": {"type": "base64", "media_type": "application/pdf", "data": "..."}} + """ + source = element["source"] + source_type = source.get("type") + if source_type != "base64": + raise ValueError( + f"Bedrock Converse only supports base64-encoded document sources, got '{source_type}'. " + "Please convert the document to base64 before sending to Bedrock." + ) + media_type: str = source["media_type"] + data: str = source["data"] + doc_format = BedrockImageProcessor._validate_format( + mime_type=media_type, image_format=media_type.split("/")[1] + ) + + # Deterministic name using the same hashing pattern as _create_bedrock_block + HASH_SAMPLE_BYTES = 64 * 1024 + normalized = "".join(data.split()).encode("utf-8") + sample = normalized[:HASH_SAMPLE_BYTES] + hasher = hashlib.sha256() + hasher.update(sample) + hasher.update(str(len(normalized)).encode("utf-8")) + content_hash = hasher.hexdigest()[:16] + document_name = f"Document_{content_hash}_{doc_format}" + + return BedrockContentBlock( + document=BedrockDocumentBlock( + source=BedrockSourceBlock(bytes=data), + format=doc_format, + name=document_name, + ) + ) + @staticmethod def add_thinking_blocks_to_assistant_content( thinking_blocks: List[BedrockContentBlock], @@ -4961,6 +5004,11 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 ) ) _parts.append(_part) + elif element["type"] == "document": + _part = BedrockConverseMessagesProcessor._process_document_message( + element + ) + _parts.append(_part) _cache_point_block = ( litellm.AmazonConverseConfig()._get_cache_point_block( message_block=cast( diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index e281b172685..fa7faf3035d 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -2244,7 +2244,7 @@ class CustomStreamWrapper: asyncio.create_task( self.logging_obj.async_failure_handler(e, traceback_exception) ) - raise e + self._handle_stream_fallback_error(e) except Exception as e: traceback_exception = traceback.format_exc() if self.logging_obj is not None: diff --git a/litellm/litellm_core_utils/url_utils.py b/litellm/litellm_core_utils/url_utils.py index a65d0892aa2..224927e5acd 100644 --- a/litellm/litellm_core_utils/url_utils.py +++ b/litellm/litellm_core_utils/url_utils.py @@ -22,7 +22,7 @@ Admins can opt out via two ``litellm`` globals (wired from proxy config): import socket from ipaddress import ip_address, ip_network from typing import Any, List, Set, Tuple -from urllib.parse import urlparse, urlunparse +from urllib.parse import quote, urlparse, urlunparse import httpx @@ -46,6 +46,46 @@ class SSRFError(ValueError): pass +def encode_url_path_segment(value: Any, *, field_name: str = "path parameter") -> str: + """Percent-encode one user-controlled URL path segment. + + ``urllib.parse.quote(..., safe="")`` intentionally leaves RFC 3986 + unreserved characters such as ``.`` unescaped, so reject standalone dot + segments before they can be appended to an upstream URL and normalized by + the HTTP client. + """ + if value is None: + raise ValueError(f"{field_name} is required") + + value_str = str(value) + if value_str == "": + raise ValueError(f"{field_name} is required") + if value_str in {".", ".."}: + raise ValueError(f"{field_name} cannot be a dot path segment") + + return quote(value_str, safe="") + + +def encode_url_path_segments(value: Any, *, field_name: str = "path") -> str: + """Percent-encode a user-controlled URL path made of multiple segments. + + Empty segments are rejected, so leading, trailing, or consecutive slashes + fail closed instead of being normalized by the HTTP client. + """ + if value is None: + raise ValueError(f"{field_name} is required") + + value_str = str(value) + if value_str == "": + raise ValueError(f"{field_name} is required") + + encoded_segments = [] + for segment in value_str.split("/"): + encoded_segments.append(encode_url_path_segment(segment, field_name=field_name)) + + return "/".join(encoded_segments) + + def _is_blocked_ip(addr: str) -> bool: """Return True for any IP not safe to reach from a user-supplied URL. @@ -199,6 +239,47 @@ def validate_url(url: str) -> Tuple[str, str]: return rewritten, host_header +def assert_same_origin(candidate_url: str, expected_url: str) -> None: + """Verify ``candidate_url`` shares scheme, host, and port with ``expected_url``. + + Use when an upstream API returns a URL meant for follow-up requests + (e.g. an async-job polling URL that will be hit with the operator's + API key in the headers). The upstream is trusted because the operator + configured ``api_base``, but the URL it hands back must actually point + back at the same origin or we'd be blindly forwarding credentials + wherever the upstream told us to. + + Hostnames are compared case-insensitively. Default ports are made + explicit (HTTP→80, HTTPS→443) so ``https://api.example.com:443/...`` + and ``https://api.example.com/...`` are treated as the same origin. + + Error messages identify *which* component mismatched but never echo + the operator's ``expected`` host or the candidate's hostname back to + the caller — in the SSRF threat model the caller is the attacker, + and reflecting host info would be a secondary leak of operator + infrastructure details. + """ + candidate = urlparse(candidate_url) + expected = urlparse(expected_url) + + if candidate.scheme not in _ALLOWED_SCHEMES: + raise SSRFError("URL scheme is not allowed") + + if candidate.scheme != expected.scheme: + raise SSRFError("Origin mismatch on scheme") + + candidate_host = _normalize_host(candidate.hostname or "") + expected_host = _normalize_host(expected.hostname or "") + if not candidate_host or candidate_host != expected_host: + raise SSRFError("Origin mismatch on host") + + default_port = 443 if candidate.scheme == "https" else 80 + candidate_port = candidate.port if candidate.port is not None else default_port + expected_port = expected.port if expected.port is not None else default_port + if candidate_port != expected_port: + raise SSRFError("Origin mismatch on port") + + _MAX_REDIRECTS = 10 diff --git a/litellm/llms/anthropic/batches/transformation.py b/litellm/llms/anthropic/batches/transformation.py index 3f03c744efe..fd67a7fbaf1 100644 --- a/litellm/llms/anthropic/batches/transformation.py +++ b/litellm/llms/anthropic/batches/transformation.py @@ -5,6 +5,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cas import httpx from httpx import Headers, Response +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.types.llms.openai import AllMessageValues, CreateBatchRequest @@ -122,7 +123,8 @@ class AnthropicBatchesConfig(BaseBatchesConfig): Complete URL for Anthropic batch retrieval: {api_base}/v1/messages/batches/{batch_id} """ api_base = api_base or self.anthropic_model_info.get_api_base(api_base) - return f"{api_base.rstrip('/')}/v1/messages/batches/{batch_id}" + encoded_batch_id = encode_url_path_segment(batch_id, field_name="batch_id") + return f"{api_base.rstrip('/')}/v1/messages/batches/{encoded_batch_id}" def transform_retrieve_batch_request( self, diff --git a/litellm/llms/anthropic/files/handler.py b/litellm/llms/anthropic/files/handler.py index c56799f30cf..56296df94a1 100644 --- a/litellm/llms/anthropic/files/handler.py +++ b/litellm/llms/anthropic/files/handler.py @@ -9,6 +9,7 @@ import litellm from litellm._logging import verbose_logger from litellm._uuid import uuid from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.types.llms.openai import ( FileContentRequest, @@ -89,7 +90,10 @@ class AnthropicFilesHandler: raise ValueError("Missing Anthropic API Key") # Construct the Anthropic batch results URL - results_url = f"{api_base.rstrip('/')}/v1/messages/batches/{batch_id}/results" + encoded_batch_id = encode_url_path_segment(batch_id, field_name="batch_id") + results_url = ( + f"{api_base.rstrip('/')}/v1/messages/batches/{encoded_batch_id}/results" + ) # Prepare headers headers = { diff --git a/litellm/llms/anthropic/files/transformation.py b/litellm/llms/anthropic/files/transformation.py index aeaab4e57bf..ea9bf00f505 100644 --- a/litellm/llms/anthropic/files/transformation.py +++ b/litellm/llms/anthropic/files/transformation.py @@ -19,6 +19,7 @@ from typing import Any, Dict, List, Optional, Union, cast import httpx from openai.types.file_deleted import FileDeleted +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.files.transformation import ( @@ -185,7 +186,8 @@ class AnthropicFilesConfig(BaseFilesConfig): AnthropicModelInfo.get_api_base(litellm_params.get("api_base")) or ANTHROPIC_FILES_API_BASE ) - return f"{api_base.rstrip('/')}/v1/files/{file_id}", {} + encoded_file_id = encode_url_path_segment(file_id, field_name="file_id") + return f"{api_base.rstrip('/')}/v1/files/{encoded_file_id}", {} def transform_retrieve_file_response( self, @@ -206,7 +208,8 @@ class AnthropicFilesConfig(BaseFilesConfig): AnthropicModelInfo.get_api_base(litellm_params.get("api_base")) or ANTHROPIC_FILES_API_BASE ) - return f"{api_base.rstrip('/')}/v1/files/{file_id}", {} + encoded_file_id = encode_url_path_segment(file_id, field_name="file_id") + return f"{api_base.rstrip('/')}/v1/files/{encoded_file_id}", {} def transform_delete_file_response( self, @@ -268,7 +271,8 @@ class AnthropicFilesConfig(BaseFilesConfig): AnthropicModelInfo.get_api_base(litellm_params.get("api_base")) or ANTHROPIC_FILES_API_BASE ) - return f"{api_base.rstrip('/')}/v1/files/{file_id}/content", {} + encoded_file_id = encode_url_path_segment(file_id, field_name="file_id") + return f"{api_base.rstrip('/')}/v1/files/{encoded_file_id}/content", {} def transform_file_content_response( self, diff --git a/litellm/llms/anthropic/skills/transformation.py b/litellm/llms/anthropic/skills/transformation.py index a992d84d459..4ea768b02af 100644 --- a/litellm/llms/anthropic/skills/transformation.py +++ b/litellm/llms/anthropic/skills/transformation.py @@ -7,6 +7,7 @@ from typing import Any, Dict, Optional, Tuple import httpx from litellm._logging import verbose_logger +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.base_llm.skills.transformation import ( BaseSkillsAPIConfig, LiteLLMLoggingObj, @@ -81,7 +82,8 @@ class AnthropicSkillsConfig(BaseSkillsAPIConfig): api_base = AnthropicModelInfo.get_api_base() if skill_id: - return f"{api_base}/v1/skills/{skill_id}" + encoded_skill_id = encode_url_path_segment(skill_id, field_name="skill_id") + return f"{api_base}/v1/skills/{encoded_skill_id}" return f"{api_base}/v1/{endpoint}" def transform_create_skill_request( diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index ca370366b6f..c0e070b6c1f 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -16,6 +16,7 @@ import litellm from litellm.constants import AZURE_OPERATION_POLLING_TIMEOUT, DEFAULT_MAX_RETRIES from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.logging_utils import track_llm_api_timing +from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, @@ -899,6 +900,17 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): operation_location_url = response.headers["operation-location"] else: raise AzureOpenAIError(status_code=500, message=response.text) + # Reject polling URLs that don't share an origin with ``api_base``. + # Without this an upstream-controlled or attacker-controlled + # value would receive the operator's Azure API key in the + # request headers below. VERIA-51. + try: + assert_same_origin(operation_location_url, api_base) + except SSRFError as ssrf_err: + raise AzureOpenAIError( + status_code=502, + message=f"Rejected polling URL: {ssrf_err}", + ) response = await async_handler.get( url=operation_location_url, headers=headers, @@ -909,8 +921,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): timeout_secs: int = AZURE_OPERATION_POLLING_TIMEOUT start_time = time.time() if "status" not in response.json(): - raise Exception( - "Expected 'status' in response. Got={}".format(response.json()) + # Don't reflect the raw response body — when the polling + # URL points at an internal JSON API (cloud metadata + # service etc.) reflecting it here turns Blind SSRF into + # Full-Read SSRF. VERIA-51. + raise AzureOpenAIError( + status_code=502, + message="Polling response missing 'status' field", ) while response.json()["status"] not in ["succeeded", "failed"]: if time.time() - start_time > timeout_secs: @@ -1010,6 +1027,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): operation_location_url = response.headers["operation-location"] else: raise AzureOpenAIError(status_code=500, message=response.text) + try: + assert_same_origin(operation_location_url, api_base) + except SSRFError as ssrf_err: + raise AzureOpenAIError( + status_code=502, + message=f"Rejected polling URL: {ssrf_err}", + ) response = sync_handler.get( url=operation_location_url, headers=headers, @@ -1020,8 +1044,9 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): timeout_secs: int = AZURE_OPERATION_POLLING_TIMEOUT start_time = time.time() if "status" not in response.json(): - raise Exception( - "Expected 'status' in response. Got={}".format(response.json()) + raise AzureOpenAIError( + status_code=502, + message="Polling response missing 'status' field", ) while response.json()["status"] not in ["succeeded", "failed"]: if time.time() - start_time > timeout_secs: diff --git a/litellm/llms/azure/responses/transformation.py b/litellm/llms/azure/responses/transformation.py index 76a6d485bc4..ca9293325ff 100644 --- a/litellm/llms/azure/responses/transformation.py +++ b/litellm/llms/azure/responses/transformation.py @@ -5,6 +5,7 @@ import httpx from openai.types.responses import ResponseReasoningItem from litellm._logging import verbose_logger +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.azure.common_utils import BaseAzureLLM from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig from litellm.types.llms.openai import * @@ -201,7 +202,10 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): # Insert the response_id at the end of the path component # Remove trailing slash if present to avoid double slashes path = parsed_url.path.rstrip("/") - new_path = f"{path}/{response_id}" + encoded_response_id = encode_url_path_segment( + response_id, field_name="response_id" + ) + new_path = f"{path}/{encoded_response_id}" # Reconstruct the URL with all original components but with the modified path constructed_url = urlunparse( @@ -322,7 +326,10 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): # Insert the response_id and /cancel at the end of the path component # Remove trailing slash if present to avoid double slashes path = parsed_url.path.rstrip("/") - new_path = f"{path}/{response_id}/cancel" + encoded_response_id = encode_url_path_segment( + response_id, field_name="response_id" + ) + new_path = f"{path}/{encoded_response_id}/cancel" # Reconstruct the URL with all original components but with the modified path cancel_url = urlunparse( diff --git a/litellm/llms/azure_ai/agents/handler.py b/litellm/llms/azure_ai/agents/handler.py index c3cd06ab4de..9bae8abce8e 100644 --- a/litellm/llms/azure_ai/agents/handler.py +++ b/litellm/llms/azure_ai/agents/handler.py @@ -36,6 +36,7 @@ from typing import ( import httpx from litellm._logging import verbose_logger +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.azure_ai.agents.transformation import ( AzureAIAgentsConfig, AzureAIAgentsError, @@ -75,20 +76,29 @@ class AzureAIAgentsHandler: def _build_messages_url( self, api_base: str, thread_id: str, api_version: str ) -> str: - return f"{api_base}/threads/{thread_id}/messages?api-version={api_version}" + encoded_thread_id = encode_url_path_segment(thread_id, field_name="thread_id") + return ( + f"{api_base}/threads/{encoded_thread_id}/messages?api-version={api_version}" + ) def _build_runs_url(self, api_base: str, thread_id: str, api_version: str) -> str: - return f"{api_base}/threads/{thread_id}/runs?api-version={api_version}" + encoded_thread_id = encode_url_path_segment(thread_id, field_name="thread_id") + return f"{api_base}/threads/{encoded_thread_id}/runs?api-version={api_version}" def _build_run_status_url( self, api_base: str, thread_id: str, run_id: str, api_version: str ) -> str: - return f"{api_base}/threads/{thread_id}/runs/{run_id}?api-version={api_version}" + encoded_thread_id = encode_url_path_segment(thread_id, field_name="thread_id") + encoded_run_id = encode_url_path_segment(run_id, field_name="run_id") + return f"{api_base}/threads/{encoded_thread_id}/runs/{encoded_run_id}?api-version={api_version}" def _build_list_messages_url( self, api_base: str, thread_id: str, api_version: str ) -> str: - return f"{api_base}/threads/{thread_id}/messages?api-version={api_version}" + encoded_thread_id = encode_url_path_segment(thread_id, field_name="thread_id") + return ( + f"{api_base}/threads/{encoded_thread_id}/messages?api-version={api_version}" + ) def _build_create_thread_and_run_url(self, api_base: str, api_version: str) -> str: """URL for the create-thread-and-run endpoint (supports streaming).""" diff --git a/litellm/llms/azure_ai/ocr/document_intelligence/transformation.py b/litellm/llms/azure_ai/ocr/document_intelligence/transformation.py index 76c247aea81..d4144a75718 100644 --- a/litellm/llms/azure_ai/ocr/document_intelligence/transformation.py +++ b/litellm/llms/azure_ai/ocr/document_intelligence/transformation.py @@ -17,11 +17,13 @@ from urllib.parse import quote import httpx from litellm._logging import verbose_logger +from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin from litellm.constants import ( AZURE_DOCUMENT_INTELLIGENCE_API_VERSION, AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI, AZURE_OPERATION_POLLING_TIMEOUT, ) +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.base_llm.ocr.transformation import ( BaseOCRConfig, DocumentType, @@ -217,11 +219,12 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig): if "/" in model: # Extract the last part after the last slash model_id = model.split("/")[-1] + encoded_model_id = encode_url_path_segment(model_id, field_name="model_id") # Azure Document Intelligence analyze endpoint # Note: API version 2024-11-30+ uses /documentintelligence/ (not /formrecognizer/) url = ( - f"{api_base}/documentintelligence/documentModels/{model_id}:analyze" + f"{api_base}/documentintelligence/documentModels/{encoded_model_id}:analyze" f"?api-version={AZURE_DOCUMENT_INTELLIGENCE_API_VERSION}" ) @@ -599,6 +602,16 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig): "Azure Document Intelligence returned 202 but no Operation-Location header found" ) + # Reject cross-origin polling URLs — the auth headers + # below would otherwise leak to whatever URL the upstream + # (or an attacker-controlled upstream) returns. VERIA-51. + try: + assert_same_origin(operation_url, str(raw_response.request.url)) + except SSRFError as ssrf_err: + raise ValueError( + f"Azure Document Intelligence: rejected polling URL ({ssrf_err})" + ) + # Get headers for polling (need auth) poll_headers = { "Ocp-Apim-Subscription-Key": raw_response.request.headers.get( @@ -711,6 +724,14 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig): "Azure Document Intelligence returned 202 but no Operation-Location header found" ) + # Reject cross-origin polling URLs (see sync path). VERIA-51. + try: + assert_same_origin(operation_url, str(raw_response.request.url)) + except SSRFError as ssrf_err: + raise ValueError( + f"Azure Document Intelligence: rejected polling URL ({ssrf_err})" + ) + # Get headers for polling (need auth) poll_headers = { "Ocp-Apim-Subscription-Key": raw_response.request.headers.get( diff --git a/litellm/llms/bedrock/chat/invoke_agent/transformation.py b/litellm/llms/bedrock/chat/invoke_agent/transformation.py index 2c7135f4d83..4c667b0ce39 100644 --- a/litellm/llms/bedrock/chat/invoke_agent/transformation.py +++ b/litellm/llms/bedrock/chat/invoke_agent/transformation.py @@ -12,6 +12,7 @@ import httpx from litellm._logging import verbose_logger from litellm._uuid import uuid +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.litellm_core_utils.prompt_templates.common_utils import ( convert_content_list_to_str, ) @@ -97,8 +98,15 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM): agent_id, agent_alias_id = self._get_agent_id_and_alias_id(model) session_id = self._get_session_id(optional_params) + encoded_agent_id = encode_url_path_segment(agent_id, field_name="agent_id") + encoded_agent_alias_id = encode_url_path_segment( + agent_alias_id, field_name="agent_alias_id" + ) + encoded_session_id = encode_url_path_segment( + session_id, field_name="session_id" + ) - endpoint_url = f"{endpoint_url}/agents/{agent_id}/agentAliases/{agent_alias_id}/sessions/{session_id}/text" + endpoint_url = f"{endpoint_url}/agents/{encoded_agent_id}/agentAliases/{encoded_agent_alias_id}/sessions/{encoded_session_id}/text" return endpoint_url diff --git a/litellm/llms/bedrock/count_tokens/transformation.py b/litellm/llms/bedrock/count_tokens/transformation.py index a37af131625..c967fd334bc 100644 --- a/litellm/llms/bedrock/count_tokens/transformation.py +++ b/litellm/llms/bedrock/count_tokens/transformation.py @@ -201,13 +201,14 @@ class BedrockCountTokensConfig(BaseAWSLLM): # Remove bedrock/ prefix if present if model_id.startswith("bedrock/"): model_id = model_id[8:] # Remove "bedrock/" prefix + encoded_model_id = self.encode_model_id(model_id=model_id) base_url, _ = self.get_runtime_endpoint( api_base=api_base, aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint, aws_region_name=aws_region_name, ) - endpoint = f"{base_url}/model/{model_id}/count-tokens" + endpoint = f"{base_url}/model/{encoded_model_id}/count-tokens" return endpoint diff --git a/litellm/llms/bedrock/vector_stores/transformation.py b/litellm/llms/bedrock/vector_stores/transformation.py index f028503c6a2..ec20d76102b 100644 --- a/litellm/llms/bedrock/vector_stores/transformation.py +++ b/litellm/llms/bedrock/vector_stores/transformation.py @@ -5,6 +5,7 @@ from urllib.parse import urlparse import httpx from litellm._logging import verbose_logger +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM from litellm.types.integrations.rag.bedrock_knowledgebase import ( @@ -209,7 +210,10 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): if isinstance(query, list): query = " ".join(query) - url = f"{api_base}/{vector_store_id}/retrieve" + encoded_vector_store_id = encode_url_path_segment( + vector_store_id, field_name="vector_store_id" + ) + url = f"{api_base}/{encoded_vector_store_id}/retrieve" request_body: Dict[str, Any] = { "retrievalQuery": BedrockKBRetrievalQuery(text=query), diff --git a/litellm/llms/black_forest_labs/image_edit/handler.py b/litellm/llms/black_forest_labs/image_edit/handler.py index dea2683a049..f5784e08367 100644 --- a/litellm/llms/black_forest_labs/image_edit/handler.py +++ b/litellm/llms/black_forest_labs/image_edit/handler.py @@ -15,6 +15,7 @@ import httpx import litellm from litellm._logging import verbose_logger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, @@ -331,6 +332,17 @@ class BlackForestLabsImageEdit: message="No polling_url in BFL response", ) + # Reject cross-origin polling URLs — the ``x-key`` auth header + # would otherwise leak to whatever URL the upstream returns. + # VERIA-51. + try: + assert_same_origin(polling_url, str(initial_response.request.url)) + except SSRFError as ssrf_err: + raise BlackForestLabsError( + status_code=502, + message=f"Rejected polling URL: {ssrf_err}", + ) + # Get just the auth header for polling polling_headers = {"x-key": headers.get("x-key", "")} @@ -416,6 +428,17 @@ class BlackForestLabsImageEdit: message="No polling_url in BFL response", ) + # Reject cross-origin polling URLs — the ``x-key`` auth header + # would otherwise leak to whatever URL the upstream returns. + # VERIA-51. + try: + assert_same_origin(polling_url, str(initial_response.request.url)) + except SSRFError as ssrf_err: + raise BlackForestLabsError( + status_code=502, + message=f"Rejected polling URL: {ssrf_err}", + ) + # Get just the auth header for polling polling_headers = {"x-key": headers.get("x-key", "")} diff --git a/litellm/llms/black_forest_labs/image_generation/handler.py b/litellm/llms/black_forest_labs/image_generation/handler.py index 5a1d885e527..8af4a236fd4 100644 --- a/litellm/llms/black_forest_labs/image_generation/handler.py +++ b/litellm/llms/black_forest_labs/image_generation/handler.py @@ -15,6 +15,7 @@ import httpx import litellm from litellm._logging import verbose_logger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, @@ -317,6 +318,17 @@ class BlackForestLabsImageGeneration: message="No polling_url in BFL response", ) + # Reject cross-origin polling URLs — the ``x-key`` auth header + # would otherwise leak to whatever URL the upstream returns. + # VERIA-51. + try: + assert_same_origin(polling_url, str(initial_response.request.url)) + except SSRFError as ssrf_err: + raise BlackForestLabsError( + status_code=502, + message=f"Rejected polling URL: {ssrf_err}", + ) + # Get just the auth header for polling polling_headers = {"x-key": headers.get("x-key", "")} @@ -402,6 +414,17 @@ class BlackForestLabsImageGeneration: message="No polling_url in BFL response", ) + # Reject cross-origin polling URLs — the ``x-key`` auth header + # would otherwise leak to whatever URL the upstream returns. + # VERIA-51. + try: + assert_same_origin(polling_url, str(initial_response.request.url)) + except SSRFError as ssrf_err: + raise BlackForestLabsError( + status_code=502, + message=f"Rejected polling URL: {ssrf_err}", + ) + # Get just the auth header for polling polling_headers = {"x-key": headers.get("x-key", "")} diff --git a/litellm/llms/bytez/chat/transformation.py b/litellm/llms/bytez/chat/transformation.py index a72f732a303..5b08670f9f2 100644 --- a/litellm/llms/bytez/chat/transformation.py +++ b/litellm/llms/bytez/chat/transformation.py @@ -5,6 +5,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union import httpx +from litellm.litellm_core_utils.url_utils import encode_url_path_segments from litellm.litellm_core_utils.exception_mapping_utils import exception_type from litellm.litellm_core_utils.logging_utils import track_llm_api_timing from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException @@ -149,7 +150,8 @@ class BytezChatConfig(BaseConfig): litellm_params: dict, stream: Optional[bool] = None, ) -> str: - return f"{API_BASE}/{model}" + encoded_model = encode_url_path_segments(model, field_name="model") + return f"{API_BASE}/{encoded_model}" def transform_request( self, diff --git a/litellm/llms/cloudflare/chat/transformation.py b/litellm/llms/cloudflare/chat/transformation.py index 9e59782bf73..b9e219f5cbc 100644 --- a/litellm/llms/cloudflare/chat/transformation.py +++ b/litellm/llms/cloudflare/chat/transformation.py @@ -5,6 +5,7 @@ from typing import AsyncIterator, Iterator, List, Optional, Union import httpx import litellm +from litellm.litellm_core_utils.url_utils import encode_url_path_segments from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator from litellm.llms.base_llm.chat.transformation import ( BaseConfig, @@ -89,7 +90,8 @@ class CloudflareChatConfig(BaseConfig): api_base = ( f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/run/" ) - return api_base + model + encoded_model = encode_url_path_segments(model, field_name="model") + return api_base + encoded_model def get_supported_openai_params(self, model: str) -> List[str]: return [ diff --git a/litellm/llms/custom_httpx/container_handler.py b/litellm/llms/custom_httpx/container_handler.py index afdd7bc6a8b..599cd705ebf 100644 --- a/litellm/llms/custom_httpx/container_handler.py +++ b/litellm/llms/custom_httpx/container_handler.py @@ -12,6 +12,7 @@ from typing import TYPE_CHECKING, Any, Coroutine, Dict, Optional, Type, Union import httpx import litellm +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, @@ -72,7 +73,8 @@ def _build_url( # Substitute path parameters for param, value in path_params.items(): - path_template = path_template.replace(f"{{{param}}}", value) + encoded_value = encode_url_path_segment(value, field_name=param) + path_template = path_template.replace(f"{{{param}}}", encoded_value) # Parse the api_base to extract existing query params parsed_base = httpx.URL(api_base) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index dc625918b98..0c4816fcda2 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -26,6 +26,7 @@ from litellm._logging import _redact_string, verbose_logger from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.base_llm.anthropic_messages.transformation import ( BaseAnthropicMessagesConfig, ) @@ -8948,7 +8949,10 @@ class BaseLLMHTTPHandler: litellm_params=dict(litellm_params), ) - url = f"{api_base}/{vector_store_id}" + encoded_vector_store_id = encode_url_path_segment( + vector_store_id, field_name="vector_store_id" + ) + url = f"{api_base}/{encoded_vector_store_id}" logging_obj.pre_call( input="", @@ -9015,7 +9019,10 @@ class BaseLLMHTTPHandler: litellm_params=dict(litellm_params), ) - url = f"{api_base}/{vector_store_id}" + encoded_vector_store_id = encode_url_path_segment( + vector_store_id, field_name="vector_store_id" + ) + url = f"{api_base}/{encoded_vector_store_id}" logging_obj.pre_call( input="", @@ -9214,7 +9221,10 @@ class BaseLLMHTTPHandler: litellm_params=dict(litellm_params), ) - url = f"{api_base}/{vector_store_id}" + encoded_vector_store_id = encode_url_path_segment( + vector_store_id, field_name="vector_store_id" + ) + url = f"{api_base}/{encoded_vector_store_id}" request_body: Dict[str, Any] = dict(vector_store_update_optional_params) @@ -9297,7 +9307,10 @@ class BaseLLMHTTPHandler: litellm_params=dict(litellm_params), ) - url = f"{api_base}/{vector_store_id}" + encoded_vector_store_id = encode_url_path_segment( + vector_store_id, field_name="vector_store_id" + ) + url = f"{api_base}/{encoded_vector_store_id}" request_body: Dict[str, Any] = dict(vector_store_update_optional_params) @@ -9363,7 +9376,10 @@ class BaseLLMHTTPHandler: litellm_params=dict(litellm_params), ) - url = f"{api_base}/{vector_store_id}" + encoded_vector_store_id = encode_url_path_segment( + vector_store_id, field_name="vector_store_id" + ) + url = f"{api_base}/{encoded_vector_store_id}" logging_obj.pre_call( input="", @@ -9428,7 +9444,10 @@ class BaseLLMHTTPHandler: litellm_params=dict(litellm_params), ) - url = f"{api_base}/{vector_store_id}" + encoded_vector_store_id = encode_url_path_segment( + vector_store_id, field_name="vector_store_id" + ) + url = f"{api_base}/{encoded_vector_store_id}" logging_obj.pre_call( input="", diff --git a/litellm/llms/elevenlabs/text_to_speech/transformation.py b/litellm/llms/elevenlabs/text_to_speech/transformation.py index 4dac2b8ba92..6a59911701b 100644 --- a/litellm/llms/elevenlabs/text_to_speech/transformation.py +++ b/litellm/llms/elevenlabs/text_to_speech/transformation.py @@ -11,13 +11,14 @@ import httpx from httpx import Headers import litellm -from litellm.types.utils import all_litellm_params +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.text_to_speech.transformation import ( BaseTextToSpeechConfig, TextToSpeechRequestData, ) from litellm.secret_managers.main import get_secret_str +from litellm.types.utils import all_litellm_params from ..common_utils import ElevenLabsException @@ -321,7 +322,8 @@ class ElevenLabsTextToSpeechConfig(BaseTextToSpeechConfig): "ElevenLabs voice_id is required. Pass `voice` when calling `litellm.speech()`." ) - url = f"{base_url}{self.TTS_ENDPOINT_PATH}/{voice_id}" + encoded_voice_id = encode_url_path_segment(voice_id, field_name="voice_id") + url = f"{base_url}{self.TTS_ENDPOINT_PATH}/{encoded_voice_id}" query_params = litellm_params.get(self.ELEVENLABS_QUERY_PARAMS_KEY, {}) if query_params: diff --git a/litellm/llms/gemini/files/transformation.py b/litellm/llms/gemini/files/transformation.py index 401d7bb9f48..63a383ebd3d 100644 --- a/litellm/llms/gemini/files/transformation.py +++ b/litellm/llms/gemini/files/transformation.py @@ -12,6 +12,7 @@ import httpx from openai.types.file_deleted import FileDeleted from litellm._logging import verbose_logger +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data from litellm.llms.base_llm.files.transformation import ( BaseFilesConfig, @@ -258,10 +259,14 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig): normalized_file_id = file_id normalized_file_id = normalized_file_id.strip("/") - if not normalized_file_id.startswith("files/"): - normalized_file_id = f"files/{normalized_file_id}" + if normalized_file_id.startswith("files/"): + normalized_file_id = normalized_file_id.removeprefix("files/") - return normalized_file_id + encoded_file_id = encode_url_path_segment( + normalized_file_id, field_name="file_id" + ) + + return f"files/{encoded_file_id}" def transform_retrieve_file_response( self, @@ -337,13 +342,8 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig): if not api_key: raise ValueError("api_key is required") - # Extract file name from URI if full URI is provided - # file_id could be "files/abc123" or "https://generativelanguage.googleapis.com/v1beta/files/abc123" - if file_id.startswith("http"): - # Extract the file path from full URI - file_name = file_id.split("/v1beta/")[-1] - else: - file_name = file_id if file_id.startswith("files/") else f"files/{file_id}" + # Normalize and encode the file name before interpolating it into the URL. + file_name = self._normalize_gemini_file_id(file_id) # Construct the delete URL url = f"{api_base}/v1beta/{file_name}" diff --git a/litellm/llms/gemini/interactions/transformation.py b/litellm/llms/gemini/interactions/transformation.py index c34da83cb8f..593cbf7c2cf 100644 --- a/litellm/llms/gemini/interactions/transformation.py +++ b/litellm/llms/gemini/interactions/transformation.py @@ -15,6 +15,7 @@ import httpx from litellm._logging import verbose_logger from litellm.litellm_core_utils.core_helpers import process_response_headers +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.base_llm.interactions.transformation import BaseInteractionsAPIConfig from litellm.llms.gemini.common_utils import GeminiError, GeminiModelInfo from litellm.types.interactions import ( @@ -205,8 +206,11 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): resolved_api_base = GeminiModelInfo.get_api_base(api_base) if not GeminiModelInfo.get_api_key(litellm_params.api_key): raise ValueError("Google API key is required") + encoded_interaction_id = encode_url_path_segment( + interaction_id, field_name="interaction_id" + ) return ( - f"{resolved_api_base}/{self.api_version}/interactions/{interaction_id}", + f"{resolved_api_base}/{self.api_version}/interactions/{encoded_interaction_id}", {}, ) @@ -238,8 +242,11 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): resolved_api_base = GeminiModelInfo.get_api_base(api_base) if not GeminiModelInfo.get_api_key(litellm_params.api_key): raise ValueError("Google API key is required") + encoded_interaction_id = encode_url_path_segment( + interaction_id, field_name="interaction_id" + ) return ( - f"{resolved_api_base}/{self.api_version}/interactions/{interaction_id}", + f"{resolved_api_base}/{self.api_version}/interactions/{encoded_interaction_id}", {}, ) @@ -268,8 +275,11 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): resolved_api_base = GeminiModelInfo.get_api_base(api_base) if not GeminiModelInfo.get_api_key(litellm_params.api_key): raise ValueError("Google API key is required") + encoded_interaction_id = encode_url_path_segment( + interaction_id, field_name="interaction_id" + ) return ( - f"{resolved_api_base}/{self.api_version}/interactions/{interaction_id}:cancel", + f"{resolved_api_base}/{self.api_version}/interactions/{encoded_interaction_id}:cancel", {}, ) diff --git a/litellm/llms/hosted_vllm/embedding/README.md b/litellm/llms/hosted_vllm/embedding/README.md index f82b3c77a6e..2c58e16fc23 100644 --- a/litellm/llms/hosted_vllm/embedding/README.md +++ b/litellm/llms/hosted_vllm/embedding/README.md @@ -2,4 +2,15 @@ No transformation is required for hosted_vllm embedding. VLLM is a superset of OpenAI's `embedding` endpoint. -To pass provider-specific parameters, see [this](https://docs.litellm.ai/docs/completion/provider_specific_params) \ No newline at end of file +## `encoding_format` + +For OpenAI-compatible embedding calls (including `openai/...` with a custom `api_base` pointing at vLLM), LiteLLM resolves `encoding_format` when it is not set on the request: + +1. Explicit value on the embedding call (`encoding_format=...`). +2. Model config (`litellm_params.encoding_format` on the proxy `model_list` entry). +3. Environment variable `LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT` (e.g. in `.env` or container env). +4. Default **`float`**. + +That avoids forwarding `encoding_format=None` to the provider/SDK where some servers behave poorly. + +To pass provider-specific parameters, see [provider-specific params](https://docs.litellm.ai/docs/completion/provider_specific_params). \ No newline at end of file diff --git a/litellm/llms/manus/files/transformation.py b/litellm/llms/manus/files/transformation.py index 3381a5327e8..34166161390 100644 --- a/litellm/llms/manus/files/transformation.py +++ b/litellm/llms/manus/files/transformation.py @@ -18,6 +18,7 @@ from openai.types.file_deleted import FileDeleted import litellm from litellm._logging import verbose_logger +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.files.transformation import ( @@ -306,7 +307,8 @@ class ManusFilesConfig(BaseFilesConfig): optional_params=optional_params, litellm_params=litellm_params, ) - return f"{api_base}/{file_id}", {} + encoded_file_id = encode_url_path_segment(file_id, field_name="file_id") + return f"{api_base}/{encoded_file_id}", {} def transform_retrieve_file_response( self, @@ -336,7 +338,8 @@ class ManusFilesConfig(BaseFilesConfig): optional_params=optional_params, litellm_params=litellm_params, ) - return f"{api_base}/{file_id}", {} + encoded_file_id = encode_url_path_segment(file_id, field_name="file_id") + return f"{api_base}/{encoded_file_id}", {} def transform_delete_file_response( self, @@ -422,7 +425,8 @@ class ManusFilesConfig(BaseFilesConfig): optional_params=optional_params, litellm_params=litellm_params, ) - return f"{api_base}/{file_id}/content", {} + encoded_file_id = encode_url_path_segment(file_id, field_name="file_id") + return f"{api_base}/{encoded_file_id}/content", {} def transform_file_content_response( self, diff --git a/litellm/llms/manus/responses/transformation.py b/litellm/llms/manus/responses/transformation.py index 510c41304a8..b3a0073a5c2 100644 --- a/litellm/llms/manus/responses/transformation.py +++ b/litellm/llms/manus/responses/transformation.py @@ -6,6 +6,7 @@ import httpx import litellm from litellm._logging import verbose_logger from litellm.litellm_core_utils.core_helpers import process_response_headers +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( _safe_convert_created_field, ) @@ -270,7 +271,10 @@ class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig): Reference: https://open.manus.im/docs/openai-compatibility """ - url = f"{api_base}/{response_id}" + encoded_response_id = encode_url_path_segment( + response_id, field_name="response_id" + ) + url = f"{api_base}/{encoded_response_id}" data: Dict = {} return url, data diff --git a/litellm/llms/openai/containers/transformation.py b/litellm/llms/openai/containers/transformation.py index 955b9f760d1..7f874ffd3b1 100644 --- a/litellm/llms/openai/containers/transformation.py +++ b/litellm/llms/openai/containers/transformation.py @@ -6,6 +6,7 @@ import litellm from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import ( StandardBuiltInToolCostTracking, ) +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.secret_managers.main import get_secret_str from litellm.types.containers.main import ( ContainerCreateOptionalRequestParams, @@ -198,7 +199,10 @@ class OpenAIContainerConfig(BaseContainerConfig): ) -> Tuple[str, Dict]: """Transform the OpenAI container retrieve request.""" # For container retrieve, we just need to construct the URL - url = join_container_api_base_path(api_base, f"/{container_id}") + encoded_container_id = encode_url_path_segment( + container_id, field_name="container_id" + ) + url = join_container_api_base_path(api_base, f"/{encoded_container_id}") # No additional data needed for GET request data: Dict[str, Any] = {} @@ -230,7 +234,10 @@ class OpenAIContainerConfig(BaseContainerConfig): - DELETE /v1/containers/{container_id} """ # Construct the URL for container delete - url = join_container_api_base_path(api_base, f"/{container_id}") + encoded_container_id = encode_url_path_segment( + container_id, field_name="container_id" + ) + url = join_container_api_base_path(api_base, f"/{encoded_container_id}") # No data needed for DELETE request data: Dict[str, Any] = {} @@ -267,7 +274,10 @@ class OpenAIContainerConfig(BaseContainerConfig): - GET /v1/containers/{container_id}/files """ # Construct the URL for container files - url = join_container_api_base_path(api_base, f"/{container_id}/files") + encoded_container_id = encode_url_path_segment( + container_id, field_name="container_id" + ) + url = join_container_api_base_path(api_base, f"/{encoded_container_id}/files") # Prepare query parameters params: Dict[str, Any] = {} @@ -311,8 +321,12 @@ class OpenAIContainerConfig(BaseContainerConfig): - GET /v1/containers/{container_id}/files/{file_id}/content """ # Construct the URL for container file content + encoded_container_id = encode_url_path_segment( + container_id, field_name="container_id" + ) + encoded_file_id = encode_url_path_segment(file_id, field_name="file_id") url = join_container_api_base_path( - api_base, f"/{container_id}/files/{file_id}/content" + api_base, f"/{encoded_container_id}/files/{encoded_file_id}/content" ) # No query parameters needed diff --git a/litellm/llms/openai/evals/transformation.py b/litellm/llms/openai/evals/transformation.py index c24dbf8637a..66537e56a6f 100644 --- a/litellm/llms/openai/evals/transformation.py +++ b/litellm/llms/openai/evals/transformation.py @@ -7,6 +7,7 @@ from typing import Any, Dict, Optional, Tuple import httpx from litellm._logging import verbose_logger +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.base_llm.evals.transformation import ( BaseEvalsAPIConfig, LiteLLMLoggingObj, @@ -76,7 +77,8 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig): api_base = "https://api.openai.com" if eval_id: - return f"{api_base}/v1/evals/{eval_id}" + encoded_eval_id = encode_url_path_segment(eval_id, field_name="eval_id") + return f"{api_base}/v1/evals/{encoded_eval_id}" return f"{api_base}/v1/{endpoint}" def transform_create_eval_request( @@ -276,7 +278,8 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig): if litellm_params and litellm_params.api_base: api_base = litellm_params.api_base - url = f"{api_base}/v1/evals/{eval_id}/runs" + encoded_eval_id = encode_url_path_segment(eval_id, field_name="eval_id") + url = f"{api_base}/v1/evals/{encoded_eval_id}/runs" # Build request body request_body = {k: v for k, v in create_request.items() if v is not None} @@ -310,7 +313,8 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig): if litellm_params and litellm_params.api_base: api_base = litellm_params.api_base - url = f"{api_base}/v1/evals/{eval_id}/runs" + encoded_eval_id = encode_url_path_segment(eval_id, field_name="eval_id") + url = f"{api_base}/v1/evals/{encoded_eval_id}/runs" # Build query parameters query_params: Dict[str, Any] = {} @@ -350,7 +354,9 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig): headers: dict, ) -> Tuple[str, Dict]: """Transform get run request for OpenAI""" - url = f"{api_base}/v1/evals/{eval_id}/runs/{run_id}" + encoded_eval_id = encode_url_path_segment(eval_id, field_name="eval_id") + encoded_run_id = encode_url_path_segment(run_id, field_name="run_id") + url = f"{api_base}/v1/evals/{encoded_eval_id}/runs/{encoded_run_id}" verbose_logger.debug("Get run request - URL: %s", url) @@ -376,7 +382,9 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig): headers: dict, ) -> Tuple[str, Dict, Dict]: """Transform cancel run request for OpenAI""" - url = f"{api_base}/v1/evals/{eval_id}/runs/{run_id}/cancel" + encoded_eval_id = encode_url_path_segment(eval_id, field_name="eval_id") + encoded_run_id = encode_url_path_segment(run_id, field_name="run_id") + url = f"{api_base}/v1/evals/{encoded_eval_id}/runs/{encoded_run_id}/cancel" # Empty body for cancel request request_body: Dict[str, Any] = {} @@ -405,7 +413,9 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig): headers: dict, ) -> Tuple[str, Dict, Dict]: """Transform delete run request for OpenAI""" - url = f"{api_base}/v1/evals/{eval_id}/runs/{run_id}" + encoded_eval_id = encode_url_path_segment(eval_id, field_name="eval_id") + encoded_run_id = encode_url_path_segment(run_id, field_name="run_id") + url = f"{api_base}/v1/evals/{encoded_eval_id}/runs/{encoded_run_id}" # Empty body for delete request request_body: Dict[str, Any] = {} diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index 87c502032cc..b7d5340d8d4 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -7,6 +7,7 @@ from pydantic import BaseModel, ValidationError import litellm from litellm._logging import verbose_logger from litellm.litellm_core_utils.core_helpers import process_response_headers +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( _safe_convert_created_field, ) @@ -421,7 +422,10 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): OpenAI API expects the following request - DELETE /v1/responses/{response_id} """ - url = f"{api_base}/{response_id}" + encoded_response_id = encode_url_path_segment( + response_id, field_name="response_id" + ) + url = f"{api_base}/{encoded_response_id}" data: Dict = {} return url, data @@ -457,7 +461,10 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): OpenAI API expects the following request - GET /v1/responses/{response_id} """ - url = f"{api_base}/{response_id}" + encoded_response_id = encode_url_path_segment( + response_id, field_name="response_id" + ) + url = f"{api_base}/{encoded_response_id}" data: Dict = {} return url, data @@ -498,7 +505,10 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): limit: int = 20, order: Literal["asc", "desc"] = "desc", ) -> Tuple[str, Dict]: - url = f"{api_base}/{response_id}/input_items" + encoded_response_id = encode_url_path_segment( + response_id, field_name="response_id" + ) + url = f"{api_base}/{encoded_response_id}/input_items" params: Dict[str, Any] = {} if after is not None: params["after"] = after @@ -540,7 +550,10 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): OpenAI API expects the following request - POST /v1/responses/{response_id}/cancel """ - url = f"{api_base}/{response_id}/cancel" + encoded_response_id = encode_url_path_segment( + response_id, field_name="response_id" + ) + url = f"{api_base}/{encoded_response_id}/cancel" data: Dict = {} return url, data diff --git a/litellm/llms/openai/vector_store_files/transformation.py b/litellm/llms/openai/vector_store_files/transformation.py index cd5f10251bb..52202f57fd3 100644 --- a/litellm/llms/openai/vector_store_files/transformation.py +++ b/litellm/llms/openai/vector_store_files/transformation.py @@ -3,6 +3,7 @@ from typing import Any, Dict, Optional, Tuple, cast import httpx import litellm +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.base_llm.vector_store_files.transformation import ( BaseVectorStoreFilesConfig, ) @@ -98,7 +99,10 @@ class OpenAIVectorStoreFilesConfig(BaseVectorStoreFilesConfig): or "https://api.openai.com/v1" ) base_url = base_url.rstrip("/") - return f"{base_url}/vector_stores/{vector_store_id}/files" + encoded_vector_store_id = encode_url_path_segment( + vector_store_id, field_name="vector_store_id" + ) + return f"{base_url}/vector_stores/{encoded_vector_store_id}/files" def transform_create_vector_store_file_request( self, @@ -163,7 +167,8 @@ class OpenAIVectorStoreFilesConfig(BaseVectorStoreFilesConfig): file_id: str, api_base: str, ) -> Tuple[str, Dict[str, Any]]: - return f"{api_base}/{file_id}", {} + encoded_file_id = encode_url_path_segment(file_id, field_name="file_id") + return f"{api_base}/{encoded_file_id}", {} def transform_retrieve_vector_store_file_response( self, @@ -186,7 +191,8 @@ class OpenAIVectorStoreFilesConfig(BaseVectorStoreFilesConfig): file_id: str, api_base: str, ) -> Tuple[str, Dict[str, Any]]: - return f"{api_base}/{file_id}/content", {} + encoded_file_id = encode_url_path_segment(file_id, field_name="file_id") + return f"{api_base}/{encoded_file_id}/content", {} def transform_retrieve_vector_store_file_content_response( self, @@ -218,7 +224,8 @@ class OpenAIVectorStoreFilesConfig(BaseVectorStoreFilesConfig): payload["attributes"] = filtered_attributes else: payload.pop("attributes", None) - return f"{api_base}/{file_id}", payload + encoded_file_id = encode_url_path_segment(file_id, field_name="file_id") + return f"{api_base}/{encoded_file_id}", payload def transform_update_vector_store_file_response( self, @@ -241,7 +248,8 @@ class OpenAIVectorStoreFilesConfig(BaseVectorStoreFilesConfig): file_id: str, api_base: str, ) -> Tuple[str, Dict[str, Any]]: - return f"{api_base}/{file_id}", {} + encoded_file_id = encode_url_path_segment(file_id, field_name="file_id") + return f"{api_base}/{encoded_file_id}", {} def transform_delete_vector_store_file_response( self, diff --git a/litellm/llms/openai/vector_stores/transformation.py b/litellm/llms/openai/vector_stores/transformation.py index 2c11d137480..bd095a0a1b7 100644 --- a/litellm/llms/openai/vector_stores/transformation.py +++ b/litellm/llms/openai/vector_stores/transformation.py @@ -3,6 +3,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast import httpx import litellm +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig from litellm.secret_managers.main import get_secret_str from litellm.types.router import GenericLiteLLMParams @@ -108,7 +109,10 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig): litellm_params: dict, extra_body: Optional[Dict[str, Any]] = None, ) -> Tuple[str, Dict]: - url = f"{api_base}/{vector_store_id}/search" + encoded_vector_store_id = encode_url_path_segment( + vector_store_id, field_name="vector_store_id" + ) + url = f"{api_base}/{encoded_vector_store_id}/search" typed_request_body = VectorStoreSearchRequest( query=query, filters=vector_store_search_optional_params.get("filters", None), diff --git a/litellm/llms/openai/videos/transformation.py b/litellm/llms/openai/videos/transformation.py index 61baa56949c..2d165a7d7df 100644 --- a/litellm/llms/openai/videos/transformation.py +++ b/litellm/llms/openai/videos/transformation.py @@ -1,11 +1,13 @@ import mimetypes from io import BufferedReader, BytesIO from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast +from urllib.parse import quote import httpx from httpx._types import RequestFiles import litellm +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.base_llm.videos.transformation import BaseVideoConfig from litellm.llms.openai.image_edit.transformation import ImageEditRequestUtils from litellm.secret_managers.main import get_secret_str @@ -220,11 +222,18 @@ class OpenAIVideoConfig(BaseVideoConfig): - GET /v1/videos/{video_id}/content?variant=thumbnail """ original_video_id = extract_original_video_id(video_id) + encoded_video_id = encode_url_path_segment( + original_video_id, field_name="video_id" + ) # Construct the URL for video content download - url = f"{api_base.rstrip('/')}/{original_video_id}/content" + url = f"{api_base.rstrip('/')}/{encoded_video_id}/content" if variant is not None: - url = f"{url}?variant={variant}" + # Encode the user-controlled ``variant`` so a value like + # ``thumbnail&extra=1`` cannot inject additional query params + # into the upstream request — same hardening rationale as the + # path-segment encoding above. + url = f"{url}?variant={quote(variant, safe='')}" # No additional data needed for GET content request data: Dict[str, Any] = {} @@ -247,9 +256,12 @@ class OpenAIVideoConfig(BaseVideoConfig): - POST /v1/videos/{video_id}/remix """ original_video_id = extract_original_video_id(video_id) + encoded_video_id = encode_url_path_segment( + original_video_id, field_name="video_id" + ) # Construct the URL for video remix - url = f"{api_base.rstrip('/')}/{original_video_id}/remix" + url = f"{api_base.rstrip('/')}/{encoded_video_id}/remix" # Prepare the request data data = {"prompt": prompt} @@ -391,9 +403,12 @@ class OpenAIVideoConfig(BaseVideoConfig): - DELETE /v1/videos/{video_id} """ original_video_id = extract_original_video_id(video_id) + encoded_video_id = encode_url_path_segment( + original_video_id, field_name="video_id" + ) # Construct the URL for video delete - url = f"{api_base.rstrip('/')}/{original_video_id}" + url = f"{api_base.rstrip('/')}/{encoded_video_id}" # No data needed for DELETE request data: Dict[str, Any] = {} @@ -427,9 +442,12 @@ class OpenAIVideoConfig(BaseVideoConfig): """ # Extract the original video_id (remove provider encoding if present) original_video_id = extract_original_video_id(video_id) + encoded_video_id = encode_url_path_segment( + original_video_id, field_name="video_id" + ) # For video retrieve, we just need to construct the URL - url = f"{api_base.rstrip('/')}/{original_video_id}" + url = f"{api_base.rstrip('/')}/{encoded_video_id}" # No additional data needed for GET request data: Dict[str, Any] = {} @@ -494,7 +512,11 @@ class OpenAIVideoConfig(BaseVideoConfig): litellm_params: GenericLiteLLMParams, headers: dict, ) -> Tuple[str, Dict]: - url = f"{api_base.rstrip('/')}/characters/{character_id}" + original_character_id = extract_original_character_id(character_id) + encoded_character_id = encode_url_path_segment( + original_character_id, field_name="character_id" + ) + url = f"{api_base.rstrip('/')}/characters/{encoded_character_id}" return url, {} def transform_video_get_character_response( diff --git a/litellm/llms/pg_vector/vector_stores/transformation.py b/litellm/llms/pg_vector/vector_stores/transformation.py index 7b22edd8676..fc4cfc7b083 100644 --- a/litellm/llms/pg_vector/vector_stores/transformation.py +++ b/litellm/llms/pg_vector/vector_stores/transformation.py @@ -1,5 +1,6 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.openai.vector_stores.transformation import OpenAIVectorStoreConfig from litellm.secret_managers.main import get_secret_str from litellm.types.router import GenericLiteLLMParams @@ -82,7 +83,10 @@ class PGVectorStoreConfig(OpenAIVectorStoreConfig): litellm_params: dict, extra_body: Optional[Dict[str, Any]] = None, ) -> Tuple[str, Dict]: - url = f"{api_base}/{vector_store_id}/search" + encoded_vector_store_id = encode_url_path_segment( + vector_store_id, field_name="vector_store_id" + ) + url = f"{api_base}/{encoded_vector_store_id}/search" _, request_body = super().transform_search_vector_store_request( vector_store_id=vector_store_id, query=query, diff --git a/litellm/llms/ragflow/chat/transformation.py b/litellm/llms/ragflow/chat/transformation.py index d49a5fd370f..990fc2b2e61 100644 --- a/litellm/llms/ragflow/chat/transformation.py +++ b/litellm/llms/ragflow/chat/transformation.py @@ -13,6 +13,7 @@ Model name format: from typing import List, Optional, Tuple import litellm +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.openai.openai import OpenAIConfig from litellm.secret_managers.main import get_secret, get_secret_str from litellm.types.llms.openai import AllMessageValues @@ -126,10 +127,11 @@ class RAGFlowConfig(OpenAIConfig): api_base = api_base[:-3] # Remove /v1 # Construct the RAGFlow-specific path + encoded_entity_id = encode_url_path_segment(entity_id, field_name="entity_id") if endpoint_type == "chat": - path = f"/api/v1/chats_openai/{entity_id}/chat/completions" + path = f"/api/v1/chats_openai/{encoded_entity_id}/chat/completions" else: # agent - path = f"/api/v1/agents_openai/{entity_id}/chat/completions" + path = f"/api/v1/agents_openai/{encoded_entity_id}/chat/completions" # Ensure path starts with / if not path.startswith("/"): diff --git a/litellm/llms/runwayml/videos/transformation.py b/litellm/llms/runwayml/videos/transformation.py index 8377dea952e..4f84816a2bc 100644 --- a/litellm/llms/runwayml/videos/transformation.py +++ b/litellm/llms/runwayml/videos/transformation.py @@ -6,6 +6,7 @@ from httpx._types import RequestFiles import litellm from litellm.constants import RUNWAYML_DEFAULT_API_VERSION +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.videos.transformation import BaseVideoConfig from litellm.llms.custom_httpx.http_handler import ( @@ -334,9 +335,12 @@ class RunwayMLVideoConfig(BaseVideoConfig): We'll retrieve the task and extract the video URL. """ original_video_id = extract_original_video_id(video_id) + encoded_video_id = encode_url_path_segment( + original_video_id, field_name="video_id" + ) # Get task status to retrieve video URL - url = f"{api_base}/tasks/{original_video_id}" + url = f"{api_base}/tasks/{encoded_video_id}" params: Dict[str, Any] = {} @@ -495,9 +499,12 @@ class RunwayMLVideoConfig(BaseVideoConfig): RunwayML uses task cancellation. """ original_video_id = extract_original_video_id(video_id) + encoded_video_id = encode_url_path_segment( + original_video_id, field_name="video_id" + ) # Construct the URL for task cancellation - url = f"{api_base}/tasks/{original_video_id}/cancel" + url = f"{api_base}/tasks/{encoded_video_id}/cancel" data: Dict[str, Any] = {} @@ -533,9 +540,12 @@ class RunwayMLVideoConfig(BaseVideoConfig): RunwayML uses GET /v1/tasks/{task_id} to retrieve task status. """ original_video_id = extract_original_video_id(video_id) + encoded_video_id = encode_url_path_segment( + original_video_id, field_name="video_id" + ) # Construct the full URL for task status retrieval - url = f"{api_base}/tasks/{original_video_id}" + url = f"{api_base}/tasks/{encoded_video_id}" # Empty dict for GET request (no body) data: Dict[str, Any] = {} diff --git a/litellm/llms/vertex_ai/batches/handler.py b/litellm/llms/vertex_ai/batches/handler.py index 7436bfef58b..c627599da8d 100644 --- a/litellm/llms/vertex_ai/batches/handler.py +++ b/litellm/llms/vertex_ai/batches/handler.py @@ -4,7 +4,11 @@ from typing import Any, Coroutine, Dict, Optional, Union import httpx import litellm -from litellm.litellm_core_utils.url_utils import async_safe_get, safe_get +from litellm.litellm_core_utils.url_utils import ( + async_safe_get, + encode_url_path_segment, + safe_get, +) from litellm.llms.custom_httpx.http_handler import ( _get_httpx_client, get_async_httpx_client, @@ -170,7 +174,8 @@ class VertexAIBatchPrediction(VertexLLM): ) # Append batch_id to the URL - default_api_base = f"{default_api_base}/{batch_id}" + encoded_batch_id = encode_url_path_segment(batch_id, field_name="batch_id") + default_api_base = f"{default_api_base}/{encoded_batch_id}" if len(default_api_base.split(":")) > 1: endpoint = default_api_base.split(":")[-1] @@ -413,7 +418,8 @@ class VertexAIBatchPrediction(VertexLLM): vertex_project=vertex_project or project_id, ) - retrieve_api_base_default = f"{default_api_base}/{batch_id}" + encoded_batch_id = encode_url_path_segment(batch_id, field_name="batch_id") + retrieve_api_base_default = f"{default_api_base}/{encoded_batch_id}" cancel_api_base_default = f"{retrieve_api_base_default}:cancel" _, api_base = self._check_custom_proxy( diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index fae175612b1..c72160f7d0a 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -97,7 +97,7 @@ def get_vertex_ai_model_route( Determine which handler to use for a Vertex AI model based on the model name. Args: - model: The model name (e.g., "llama3-405b", "gemini-pro", "gemma/gemma-3-12b-it", "openai/gpt-oss-120b") + model: The model name (e.g., "llama3-405b", "gemini-pro", "gemma/gemma-3-12b-it", "xai/grok-4.1-fast-non-reasoning") litellm_params: Optional litellm parameters dict that may contain base_model for routing Returns: @@ -113,7 +113,7 @@ def get_vertex_ai_model_route( >>> get_vertex_ai_model_route("gemma/gemma-3-12b-it") VertexAIModelRoute.GEMMA - >>> get_vertex_ai_model_route("openai/gpt-oss-120b") + >>> get_vertex_ai_model_route("xai/grok-4.1-fast-non-reasoning") VertexAIModelRoute.MODEL_GARDEN >>> get_vertex_ai_model_route("1234567890", {"api_base": "http://10.96.32.8"}) @@ -149,8 +149,11 @@ def get_vertex_ai_model_route( if "gemma/" in model: return VertexAIModelRoute.GEMMA - # Check for model garden openai models - if "openai" in model: + # Check for model garden OpenAI-compatible publisher models. + # Examples: + # - openai/gpt-oss-120b-maas + # - xai/grok-4.1-fast-non-reasoning + if "openai" in model or model.startswith("xai/"): return VertexAIModelRoute.MODEL_GARDEN # Check for gemini models @@ -256,8 +259,8 @@ def get_vertex_base_model_name(model: str) -> str: >>> get_vertex_base_model_name("gemma/gemma-3-12b-it") "gemma-3-12b-it" - >>> get_vertex_base_model_name("openai/gpt-oss-120b") - "gpt-oss-120b" + >>> get_vertex_base_model_name("xai/grok-4.1-fast-non-reasoning") + "grok-4.1-fast-non-reasoning" >>> get_vertex_base_model_name("1234567890") "1234567890" diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index f756c7a7c52..9afa5dec465 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -759,16 +759,22 @@ def _transform_request_body( # noqa: PLR0915 ] data = RequestBody(contents=content) - if system_instructions is not None: - data["system_instruction"] = system_instructions - if tools is not None: - data["tools"] = tools - if tool_choice is not None: - data["toolConfig"] = tool_choice - if include_server_side_tool_invocations: - if "toolConfig" not in data: - data["toolConfig"] = {} - data["toolConfig"]["includeServerSideToolInvocations"] = True + # Vertex rejects system_instruction/tools/toolConfig alongside cachedContent. + # Treat dropping these fields as a request mutation guarded by modify_params. + can_send_cache_incompatible_fields = ( + cached_content is None or litellm.modify_params is False + ) + if can_send_cache_incompatible_fields: + if system_instructions is not None: + data["system_instruction"] = system_instructions + if tools is not None: + data["tools"] = tools + if tool_choice is not None: + data["toolConfig"] = tool_choice + if include_server_side_tool_invocations: + if "toolConfig" not in data: + data["toolConfig"] = {} + data["toolConfig"]["includeServerSideToolInvocations"] = True if safety_settings is not None: data["safetySettings"] = safety_settings if generation_config is not None and len(generation_config) > 0: diff --git a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py index 2371bc4865a..99165c37c93 100644 --- a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py +++ b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py @@ -3,7 +3,7 @@ Google AI Studio /batchEmbedContents Embeddings Endpoint """ import json -from typing import Any, Dict, Literal, Optional, Union +from typing import Any, Dict, List, Literal, Optional, Tuple, Union import httpx @@ -13,8 +13,8 @@ from litellm.llms.custom_httpx.http_handler import ( HTTPHandler, get_async_httpx_client, ) -from litellm.types.llms.openai import EmbeddingInput from litellm.types.llms.vertex_ai import ( + GeminiEmbeddingInput, VertexAIBatchEmbeddingsRequestBody, VertexAIBatchEmbeddingsResponseObject, ) @@ -23,7 +23,6 @@ from litellm.types.utils import EmbeddingResponse from ..gemini.vertex_and_google_ai_studio_gemini import VertexLLM from .batch_embed_content_transformation import ( _is_file_reference, - _is_multimodal_input, process_embed_content_response, process_response, transform_openai_input_gemini_content, @@ -32,9 +31,24 @@ from .batch_embed_content_transformation import ( class GoogleBatchEmbeddings(VertexLLM): + @staticmethod + def _flatten_and_detect_file_refs( + input: GeminiEmbeddingInput, + ) -> Tuple[List[str], bool]: + """Flatten nested input lists and detect file references.""" + input_list = [input] if isinstance(input, str) else input + flat_elements = [ + e + for item in input_list + for e in (item if isinstance(item, list) else [item]) + if isinstance(e, str) + ] + has_file_refs = any(_is_file_reference(e) for e in flat_elements) + return flat_elements, has_file_refs + def _resolve_file_references( self, - input: EmbeddingInput, + input: GeminiEmbeddingInput, api_key: str, sync_handler: HTTPHandler, ) -> Dict[str, Dict[str, str]]: @@ -42,7 +56,7 @@ class GoogleBatchEmbeddings(VertexLLM): Resolve Gemini file references (files/...) to get mime_type and uri. Args: - input: EmbeddingInput that may contain file references + input: GeminiEmbeddingInput that may contain file references api_key: Gemini API key sync_handler: HTTP client @@ -73,7 +87,7 @@ class GoogleBatchEmbeddings(VertexLLM): async def _async_resolve_file_references( self, - input: EmbeddingInput, + input: GeminiEmbeddingInput, api_key: str, async_handler: AsyncHTTPHandler, ) -> Dict[str, Dict[str, str]]: @@ -81,7 +95,7 @@ class GoogleBatchEmbeddings(VertexLLM): Async version of _resolve_file_references. Args: - input: EmbeddingInput that may contain file references + input: GeminiEmbeddingInput that may contain file references api_key: Gemini API key async_handler: Async HTTP client @@ -110,10 +124,10 @@ class GoogleBatchEmbeddings(VertexLLM): return resolved_files - def batch_embeddings( + def batch_embeddings( # noqa: PLR0915 self, model: str, - input: EmbeddingInput, + input: GeminiEmbeddingInput, print_verbose, model_response: EmbeddingResponse, custom_llm_provider: Literal["gemini", "vertex_ai"], @@ -151,8 +165,7 @@ class GoogleBatchEmbeddings(VertexLLM): optional_params = optional_params or {} - is_multimodal = _is_multimodal_input(input) - use_embed_content = is_multimodal or (custom_llm_provider == "vertex_ai") + use_embed_content = custom_llm_provider == "vertex_ai" mode: Literal["embedding", "batch_embedding"] if use_embed_content: mode = "embedding" @@ -215,8 +228,22 @@ class GoogleBatchEmbeddings(VertexLLM): resolved_files=resolved_files, ) else: + flat_elements, has_file_refs = self._flatten_and_detect_file_refs(input) + if has_file_refs and not api_key: + raise ValueError( + "An API key is required to resolve Gemini file references (files/...). " + "Pass api_key= or set GEMINI_API_KEY." + ) + resolved_files = {} + if api_key and has_file_refs: + resolved_files = self._resolve_file_references( + input=flat_elements, api_key=api_key, sync_handler=sync_handler + ) request_data = transform_openai_input_gemini_content( - input=input, model=model, optional_params=optional_params + input=input, + model=model, + optional_params=optional_params, + resolved_files=resolved_files, ) ## LOGGING @@ -264,7 +291,7 @@ class GoogleBatchEmbeddings(VertexLLM): url: str, data: Optional[Union[VertexAIBatchEmbeddingsRequestBody, dict]], model_response: EmbeddingResponse, - input: EmbeddingInput, + input: GeminiEmbeddingInput, timeout: Optional[Union[float, httpx.Timeout]], headers={}, client: Optional[AsyncHTTPHandler] = None, @@ -303,8 +330,22 @@ class GoogleBatchEmbeddings(VertexLLM): resolved_files=resolved_files, ) else: + flat_elements, has_file_refs = self._flatten_and_detect_file_refs(input) + if has_file_refs and not api_key: + raise ValueError( + "An API key is required to resolve Gemini file references (files/...). " + "Pass api_key= or set GEMINI_API_KEY." + ) + resolved_files = {} + if api_key and has_file_refs: + resolved_files = await self._async_resolve_file_references( + input=flat_elements, api_key=api_key, async_handler=async_handler + ) data = transform_openai_input_gemini_content( - input=input, model=model, optional_params=optional_params or {} + input=input, + model=model, + optional_params=optional_params or {}, + resolved_files=resolved_files, ) ## LOGGING diff --git a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py index 34fc95e0af7..e1b365c9f42 100644 --- a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py +++ b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py @@ -6,12 +6,12 @@ Why separate file? Make it easy to see how transformation works from typing import Dict, List, Optional, Tuple -from litellm.types.llms.openai import EmbeddingInput from litellm.types.llms.vertex_ai import ( BlobType, ContentType, EmbedContentRequest, FileDataType, + GeminiEmbeddingInput, PartType, VertexAIBatchEmbeddingsRequestBody, VertexAIBatchEmbeddingsResponseObject, @@ -114,33 +114,77 @@ def _parse_data_url(data_url: str) -> Tuple[str, str]: return media_type, base64_data -def _is_multimodal_input(input: EmbeddingInput) -> bool: +def _is_multimodal_input(input: GeminiEmbeddingInput) -> bool: """ - Check if the input contains multimodal data (data URIs, file references, or GCS URLs). + Check if the input contains multimodal data (data URIs, file references, + GCS URLs, or nested lists for combined embeddings). Args: - input: EmbeddingInput (str or List[str]) + input: GeminiEmbeddingInput — str, List[str], or List[List[str]] for combined embeddings Returns: - bool: True if any element is a data URI, file reference, or GCS URL + bool: True if any element is multimodal or a nested list """ if isinstance(input, str): - input_list = [input] - else: - input_list = input + return _is_multimodal_element(input) - for element in input_list: - if isinstance(element, str): - if element.startswith("data:") and ";base64," in element: - return True - if _is_file_reference(element): - return True - if _is_gcs_url(element): + for element in input: + if isinstance(element, list): + if any( + _is_multimodal_element(sub) for sub in element if isinstance(sub, str) + ): return True + elif isinstance(element, str) and _is_multimodal_element(element): + return True return False +def _is_multimodal_element(element: str) -> bool: + """Check if a single string element is multimodal.""" + if element.startswith("data:") and ";base64," in element: + return True + if _is_file_reference(element): + return True + if _is_gcs_url(element): + return True + return False + + +def _build_part_for_input( + element: str, + resolved_files: Optional[Dict[str, Dict[str, str]]] = None, +) -> PartType: + """ + Build a single PartType for an input element, handling text, data URIs, + file references, and GCS URLs. + """ + resolved_files = resolved_files or {} + + if element.startswith("data:") and ";base64," in element: + mime_type, base64_data = _parse_data_url(element) + blob: BlobType = {"mime_type": mime_type, "data": base64_data} + return PartType(inline_data=blob) + elif _is_gcs_url(element): + mime_type = _infer_mime_type_from_gcs_url(element) + file_data: FileDataType = { + "mime_type": mime_type, + "file_uri": element, + } + return PartType(file_data=file_data) + elif _is_file_reference(element): + if element not in resolved_files: + raise ValueError(f"File reference {element} not resolved") + file_info = resolved_files[element] + file_data_ref: FileDataType = { + "mime_type": file_info["mime_type"], + "file_uri": file_info["uri"], + } + return PartType(file_data=file_data_ref) + else: + return PartType(text=element) + + _SUPPORTED_EMBED_PARAMS = {"outputDimensionality", "taskType", "title"} @@ -155,37 +199,60 @@ def _filter_embed_params(optional_params: dict) -> dict: def transform_openai_input_gemini_content( - input: EmbeddingInput, model: str, optional_params: dict + input: GeminiEmbeddingInput, + model: str, + optional_params: dict, + resolved_files: Optional[Dict[str, Dict[str, str]]] = None, ) -> VertexAIBatchEmbeddingsRequestBody: """ - The content to embed. Only the parts.text fields will be counted. + Transform OpenAI embedding input to Gemini batchEmbedContents format. + + Each input element becomes a separate EmbedContentRequest, supporting + text, data URIs, file references, and GCS URLs. + + If an element is a list (nested input), all sub-elements are combined + into a single content with multiple parts, producing one combined + embedding for the group. + + Examples: + input=["text", "image"] → 2 separate embeddings + input=[["text", "image"]] → 1 combined embedding + input=[["text", "image"], "x"] → 2 embeddings (1 combined + 1 separate) """ gemini_model_name = "models/{}".format(model) gemini_params = _filter_embed_params(optional_params) + input_list = [input] if isinstance(input, str) else input requests: List[EmbedContentRequest] = [] - if isinstance(input, str): + + for element in input_list: + if isinstance(element, list): + if not element: + raise ValueError("Nested input list must not be empty") + for sub in element: + if not isinstance(sub, str): + raise ValueError( + f"Elements inside a nested input list must be strings, got {type(sub)}" + ) + parts = [ + _build_part_for_input(sub, resolved_files=resolved_files) + for sub in element + ] + else: + parts = [_build_part_for_input(element, resolved_files=resolved_files)] request = EmbedContentRequest( model=gemini_model_name, - content=ContentType(parts=[PartType(text=input)]), + content=ContentType(parts=parts), **gemini_params, ) requests.append(request) - else: - for i in input: - request = EmbedContentRequest( - model=gemini_model_name, - content=ContentType(parts=[PartType(text=i)]), - **gemini_params, - ) - requests.append(request) return VertexAIBatchEmbeddingsRequestBody(requests=requests) def transform_openai_input_gemini_embed_content( - input: EmbeddingInput, + input: GeminiEmbeddingInput, model: str, optional_params: dict, resolved_files: Optional[Dict[str, Dict[str, str]]] = None, @@ -194,7 +261,7 @@ def transform_openai_input_gemini_embed_content( Transform OpenAI embedding input to Gemini embedContent format (multimodal). Args: - input: EmbeddingInput (str or List[str]) with text, data URIs, or file references + input: GeminiEmbeddingInput with text, data URIs, or file references model: Model name optional_params: Additional parameters (taskType, outputDimensionality, etc.) resolved_files: Dict mapping file names (files/abc) to {mime_type, uri} @@ -210,31 +277,14 @@ def transform_openai_input_gemini_embed_content( parts: List[PartType] = [] for element in input_list: + if isinstance(element, list): + raise ValueError( + "Nested (combined) embeddings are not supported on the embedContent path. " + "Use the batchEmbedContents path or pass a flat list instead." + ) if not isinstance(element, str): raise ValueError(f"Unsupported input type: {type(element)}") - - if element.startswith("data:") and ";base64," in element: - mime_type, base64_data = _parse_data_url(element) - blob: BlobType = {"mime_type": mime_type, "data": base64_data} - parts.append(PartType(inline_data=blob)) - elif _is_gcs_url(element): - mime_type = _infer_mime_type_from_gcs_url(element) - file_data: FileDataType = { - "mime_type": mime_type, - "file_uri": element, - } - parts.append(PartType(file_data=file_data)) - elif _is_file_reference(element): - if element not in resolved_files: - raise ValueError(f"File reference {element} not resolved") - file_info = resolved_files[element] - file_data_ref: FileDataType = { - "mime_type": file_info["mime_type"], - "file_uri": file_info["uri"], - } - parts.append(PartType(file_data=file_data_ref)) - else: - parts.append(PartType(text=element)) + parts.append(_build_part_for_input(element, resolved_files=resolved_files)) request_body: dict = { "content": ContentType(parts=parts), @@ -245,7 +295,7 @@ def transform_openai_input_gemini_embed_content( def process_embed_content_response( - input: EmbeddingInput, + input: GeminiEmbeddingInput, model_response: EmbeddingResponse, model: str, response_json: dict, @@ -291,7 +341,7 @@ def process_embed_content_response( def process_response( - input: EmbeddingInput, + input: GeminiEmbeddingInput, model_response: EmbeddingResponse, model: str, _predictions: VertexAIBatchEmbeddingsResponseObject, @@ -308,8 +358,29 @@ def process_response( model_response.data = openai_embeddings model_response.model = model - input_text = get_formatted_prompt(data={"input": input}, call_type="embedding") - prompt_tokens = token_counter(model=model, text=input_text) + has_nested = isinstance(input, list) and any(isinstance(e, list) for e in input) + if _is_multimodal_input(input) or has_nested: + input_list = input if isinstance(input, list) else [input] + text_elements: List[str] = [] + for e in input_list: + if isinstance(e, list): + text_elements.extend( + sub + for sub in e + if isinstance(sub, str) and not _is_multimodal_element(sub) + ) + elif isinstance(e, str) and not _is_multimodal_element(e): + text_elements.append(e) + if text_elements: + input_text = get_formatted_prompt( + data={"input": text_elements}, call_type="embedding" + ) + prompt_tokens = token_counter(model=model, text=input_text) + else: + prompt_tokens = 0 + else: + input_text = get_formatted_prompt(data={"input": input}, call_type="embedding") + prompt_tokens = token_counter(model=model, text=input_text) model_response.usage = Usage( prompt_tokens=prompt_tokens, total_tokens=prompt_tokens ) diff --git a/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py b/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py index 6cb7a86bea2..61fb848b40a 100644 --- a/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py +++ b/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py @@ -3,6 +3,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union import httpx from litellm import get_model_info +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig from litellm.llms.vertex_ai.vertex_llm_base import VertexBase from litellm.types.router import GenericLiteLLMParams @@ -91,12 +92,18 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): raise ValueError("vector_store_id is required") if api_base: return api_base.rstrip("/") + encoded_collection_id = encode_url_path_segment( + collection_id, field_name="vertex_collection_id" + ) + encoded_datastore_id = encode_url_path_segment( + datastore_id, field_name="vector_store_id" + ) # Vertex AI Search API endpoint for search return ( f"https://discoveryengine.googleapis.com/v1/" f"projects/{vertex_project}/locations/{vertex_location}/" - f"collections/{collection_id}/dataStores/{datastore_id}/servingConfigs/default_config" + f"collections/{encoded_collection_id}/dataStores/{encoded_datastore_id}/servingConfigs/default_config" ) def transform_search_vector_store_request( diff --git a/litellm/llms/vertex_ai/vertex_model_garden/main.py b/litellm/llms/vertex_ai/vertex_model_garden/main.py index c37bb449ecf..7240d9dce57 100644 --- a/litellm/llms/vertex_ai/vertex_model_garden/main.py +++ b/litellm/llms/vertex_ai/vertex_model_garden/main.py @@ -27,6 +27,17 @@ from ..common_utils import VertexAIError, get_vertex_base_model_name from ..vertex_llm_base import VertexBase +def _vertex_model_garden_model_id_in_json_body(model: str) -> bool: + """ + Vertex catalog / publisher models are addressed as publisher/model (e.g. + xai/grok-4.1-fast-reasoning) on the shared OpenAPI URL, with the id in the JSON body. + + Deployed Model Garden endpoints are typically a single segment (often numeric) + and use .../endpoints/{ENDPOINT_ID}/chat/completions with an empty model field. + """ + return "/" in model + + def create_vertex_url( vertex_location: str, vertex_project: str, @@ -34,8 +45,13 @@ def create_vertex_url( model: str, api_base: Optional[str] = None, ) -> str: - """Return the base url for the vertex garden models""" + """Return the api base for vertex model garden (without /chat/completions).""" base_url = get_vertex_base_url(vertex_location) + if _vertex_model_garden_model_id_in_json_body(model): + return ( + f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}" + "/endpoints/openapi" + ) return f"{base_url}/v1beta1/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}" @@ -129,7 +145,10 @@ class VertexAIModelGardenModels(VertexBase): vertex_location=vertex_location or "us-central1", vertex_api_version="v1beta1", ) - model = "" + # Publisher/catalog models: model id must be sent in the JSON body (OpenAPI route). + # Single-segment endpoint ids: model is encoded in the URL path; body model stays empty. + if not _vertex_model_garden_model_id_in_json_body(model): + model = "" return openai_like_chat_completions.completion( model=model, messages=messages, diff --git a/litellm/llms/volcengine/responses/transformation.py b/litellm/llms/volcengine/responses/transformation.py index f6dda4dd25b..99e0a958ef1 100644 --- a/litellm/llms/volcengine/responses/transformation.py +++ b/litellm/llms/volcengine/responses/transformation.py @@ -17,6 +17,7 @@ from pydantic import fields as pyd_fields import litellm from litellm._logging import verbose_logger from litellm.litellm_core_utils.core_helpers import process_response_headers +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( _safe_convert_created_field, ) @@ -300,7 +301,10 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig): litellm_params: GenericLiteLLMParams, headers: dict, ) -> Tuple[str, Dict]: - url = f"{api_base}/{response_id}" + encoded_response_id = encode_url_path_segment( + response_id, field_name="response_id" + ) + url = f"{api_base}/{encoded_response_id}" data: Dict = {} return url, data @@ -333,7 +337,10 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig): litellm_params: GenericLiteLLMParams, headers: dict, ) -> Tuple[str, Dict]: - url = f"{api_base}/{response_id}" + encoded_response_id = encode_url_path_segment( + response_id, field_name="response_id" + ) + url = f"{api_base}/{encoded_response_id}" data: Dict = {} return url, data @@ -372,7 +379,10 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig): limit: int = 20, order: Literal["asc", "desc"] = "desc", ) -> Tuple[str, Dict]: - url = f"{api_base}/{response_id}/input_items" + encoded_response_id = encode_url_path_segment( + response_id, field_name="response_id" + ) + url = f"{api_base}/{encoded_response_id}/input_items" params: Dict[str, Any] = {} if after is not None: params["after"] = after @@ -408,7 +418,10 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig): litellm_params: GenericLiteLLMParams, headers: dict, ) -> Tuple[str, Dict]: - url = f"{api_base}/{response_id}/cancel" + encoded_response_id = encode_url_path_segment( + response_id, field_name="response_id" + ) + url = f"{api_base}/{encoded_response_id}/cancel" data: Dict = {} return url, data diff --git a/litellm/llms/xai/chat/transformation.py b/litellm/llms/xai/chat/transformation.py index bfa55105a6c..64b4a545acb 100644 --- a/litellm/llms/xai/chat/transformation.py +++ b/litellm/llms/xai/chat/transformation.py @@ -43,6 +43,7 @@ class XAIChatConfig(OpenAIGPTConfig): "logprobs", "max_tokens", "n", + "parallel_tool_calls", "presence_penalty", "response_format", "seed", diff --git a/litellm/main.py b/litellm/main.py index 0079bd750cf..0553cf9d422 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -4923,8 +4923,17 @@ def embedding( # noqa: PLR0915 if encoding_format is not None: optional_params["encoding_format"] = encoding_format else: - # Omiting causes openai sdk to add default value of "float" - optional_params["encoding_format"] = None + env_fmt = get_secret_str("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT") + if env_fmt is not None and env_fmt.strip().lower() == "none": + optional_params.pop("encoding_format", None) + else: + _default_fmt = ( + optional_params.get("encoding_format") or env_fmt or "float" + ) + if _default_fmt.strip().lower() == "none": + optional_params.pop("encoding_format", None) + else: + optional_params["encoding_format"] = _default_fmt api_version = None diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index b690941e73a..6078d7e6907 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -33429,6 +33429,72 @@ "source": "https://console.cloud.google.com/vertex-ai/publishers/openai/model-garden/gpt-oss-120b-maas", "supports_reasoning": true }, + "vertex_ai/xai/grok-4.1-fast-non-reasoning": { + "cache_read_input_token_cost": 5e-08, + "input_cost_per_token": 2e-07, + "litellm_provider": "vertex_ai", + "max_input_tokens": 2000000, + "max_output_tokens": 2000000, + "max_tokens": 2000000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "source": "https://docs.x.ai/docs/models (Vertex AI Model Garden)", + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "vertex_ai/xai/grok-4.1-fast-reasoning": { + "cache_read_input_token_cost": 5e-08, + "input_cost_per_token": 2e-07, + "litellm_provider": "vertex_ai", + "max_input_tokens": 2000000, + "max_output_tokens": 2000000, + "max_tokens": 2000000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "source": "https://docs.x.ai/docs/models (Vertex AI Model Garden)", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "vertex_ai/xai/grok-4.20-non-reasoning": { + "cache_read_input_token_cost": 2e-07, + "input_cost_per_token": 2e-06, + "litellm_provider": "vertex_ai", + "max_input_tokens": 2000000, + "max_output_tokens": 2000000, + "max_tokens": 2000000, + "mode": "chat", + "output_cost_per_token": 6e-06, + "source": "https://docs.x.ai/docs/models (Vertex AI Model Garden)", + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "vertex_ai/xai/grok-4.20-reasoning": { + "cache_read_input_token_cost": 2e-07, + "input_cost_per_token": 2e-06, + "litellm_provider": "vertex_ai", + "max_input_tokens": 2000000, + "max_output_tokens": 2000000, + "max_tokens": 2000000, + "mode": "chat", + "output_cost_per_token": 6e-06, + "source": "https://docs.x.ai/docs/models (Vertex AI Model Garden)", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, "vertex_ai/qwen/qwen3-235b-a22b-instruct-2507-maas": { "input_cost_per_token": 2.5e-07, "litellm_provider": "vertex_ai-qwen_models", diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index cebd224a1a7..1794cd14381 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -33,10 +33,12 @@ def get_request_base_url(request: Request) -> str: """ Get the base URL for the request, considering X-Forwarded-* headers. - When behind a proxy (like nginx), the proxy may set: - - X-Forwarded-Proto: The original protocol (http/https) - - X-Forwarded-Host: The original host (may include port) - - X-Forwarded-Port: The original port (if not in Host header) + X-Forwarded-Proto / X-Forwarded-Host / X-Forwarded-Port are only honoured + when the request comes from a configured trusted proxy + (``use_x_forwarded_for`` enabled AND caller in ``mcp_trusted_proxy_ranges``). + Otherwise the request's literal ``base_url`` is returned, so an + untrusted caller cannot poison OAuth-discovery / redirect_uri values + by injecting headers. Args: request: FastAPI Request object @@ -47,34 +49,28 @@ def get_request_base_url(request: Request) -> str: base_url = str(request.base_url).rstrip("/") parsed = urlparse(base_url) - # Get forwarded headers + if not IPAddressUtils.is_request_from_trusted_proxy(request): + return base_url + x_forwarded_proto = request.headers.get("X-Forwarded-Proto") x_forwarded_host = request.headers.get("X-Forwarded-Host") x_forwarded_port = request.headers.get("X-Forwarded-Port") - # Start with the original scheme scheme = x_forwarded_proto if x_forwarded_proto else parsed.scheme - # Handle host and port if x_forwarded_host: # X-Forwarded-Host may already include port (e.g., "example.com:8080") if ":" in x_forwarded_host and not x_forwarded_host.startswith("["): - # Host includes port netloc = x_forwarded_host elif x_forwarded_port: - # Port is separate netloc = f"{x_forwarded_host}:{x_forwarded_port}" else: - # Just host, no explicit port netloc = x_forwarded_host else: - # No X-Forwarded-Host, use original netloc netloc = parsed.netloc if x_forwarded_port and ":" not in netloc: - # Add forwarded port if not already in netloc netloc = f"{netloc}:{x_forwarded_port}" - # Reconstruct the URL return urlunparse((scheme, netloc, parsed.path, "", "", "")) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index f96350500db..9923c3ce4bf 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -169,6 +169,37 @@ def _deserialize_json_dict(data: Any) -> Optional[Dict[str, str]]: class MCPServerManager: _STDIO_ENV_TEMPLATE_PATTERN = re.compile(r"^\$\{(X-[^}]+)\}$") + @staticmethod + def _resolve_oauth2_flow( + *, + auth_type: Optional[MCPAuthType], + oauth2_flow: Optional[str], + token_url: Optional[str], + authorization_url: Optional[str], + client_id: Optional[str], + client_secret: Optional[str], + ) -> Optional[Literal["client_credentials", "authorization_code"]]: + """Infer oauth2_flow for legacy records that omit the field. + + DB rows created before oauth2_flow support may have OAuth2 client + credentials + token_url but a null oauth2_flow. Treat these as M2M, + unless authorization_url is present (interactive OAuth). + """ + if oauth2_flow in ("client_credentials", "authorization_code"): + return cast( + Literal["client_credentials", "authorization_code"], oauth2_flow + ) + if oauth2_flow: + # Ignore unknown/untyped values and continue legacy inference. + return None + if auth_type != MCPAuth.oauth2: + return None + if authorization_url: + return None + if token_url and client_id and client_secret: + return "client_credentials" + return None + def __init__(self): self.registry: Dict[str, MCPServer] = {} self.config_mcp_servers: Dict[str, MCPServer] = {} @@ -342,7 +373,14 @@ class MCPServerManager: # oauth specific fields client_id=server_config.get("client_id", None), client_secret=server_config.get("client_secret", None), - oauth2_flow=server_config.get("oauth2_flow", None), + oauth2_flow=self._resolve_oauth2_flow( + auth_type=auth_type, + oauth2_flow=server_config.get("oauth2_flow", None), + token_url=resolved_token_url, + authorization_url=resolved_authorization_url, + client_id=server_config.get("client_id", None), + client_secret=server_config.get("client_secret", None), + ), scopes=resolved_scopes, authorization_url=resolved_authorization_url, token_url=resolved_token_url, @@ -679,7 +717,17 @@ class MCPServerManager: client_id=client_id_value or getattr(mcp_server, "client_id", None), client_secret=client_secret_value or getattr(mcp_server, "client_secret", None), - oauth2_flow=getattr(mcp_server, "oauth2_flow", None), + oauth2_flow=self._resolve_oauth2_flow( + auth_type=auth_type, + oauth2_flow=getattr(mcp_server, "oauth2_flow", None), + token_url=mcp_server.token_url + or getattr(mcp_oauth_metadata, "token_url", None), + authorization_url=mcp_server.authorization_url + or getattr(mcp_oauth_metadata, "authorization_url", None), + client_id=client_id_value or getattr(mcp_server, "client_id", None), + client_secret=client_secret_value + or getattr(mcp_server, "client_secret", None), + ), scopes=resolved_scopes, authorization_url=mcp_server.authorization_url or getattr(mcp_oauth_metadata, "authorization_url", None), @@ -2426,7 +2474,7 @@ class MCPServerManager: ) ) - async def _call_regular_mcp_tool( + async def _call_regular_mcp_tool( # noqa: PLR0915 self, mcp_server: MCPServer, original_tool_name: str, @@ -2489,7 +2537,11 @@ class MCPServerManager: # oauth2 headers extra_headers: Optional[Dict[str, str]] = None if mcp_server.auth_type == MCPAuth.oauth2: - extra_headers = oauth2_headers + if mcp_server.has_client_credentials: + # For M2M OAuth servers, Authorization must come from token fetch. + extra_headers = None + else: + extra_headers = oauth2_headers if mcp_server.extra_headers and raw_headers: if extra_headers is None: @@ -2501,6 +2553,11 @@ class MCPServerManager: for header in mcp_server.extra_headers: if not isinstance(header, str): continue + if ( + mcp_server.has_client_credentials + and header.lower() == "authorization" + ): + continue header_value = normalized_raw_headers.get(header.lower()) if header_value is None: continue @@ -2536,6 +2593,10 @@ class MCPServerManager: ) extra_headers.update(hook_extra_headers) + # Reset to None if no headers were actually added + if extra_headers is not None and len(extra_headers) == 0: + extra_headers = None + stdio_env = self._build_stdio_env(mcp_server, raw_headers) client = await self._create_mcp_client( diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index ae6055217b8..abb4b5cfa6f 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -153,6 +153,7 @@ if MCP_AVAILABLE: MCPAuthenticatedUser, ) from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, global_mcp_server_manager, ) from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( @@ -900,6 +901,20 @@ if MCP_AVAILABLE: allowed_mcp_server_id ) if mcp_server is not None: + # Apply oauth2_flow resolution for legacy DB rows where it may be NULL + resolved_flow = MCPServerManager._resolve_oauth2_flow( + auth_type=mcp_server.auth_type, + oauth2_flow=mcp_server.oauth2_flow, + token_url=mcp_server.token_url, + authorization_url=mcp_server.authorization_url, + client_id=mcp_server.client_id, + client_secret=mcp_server.client_secret, + ) + if resolved_flow and resolved_flow != mcp_server.oauth2_flow: + # Create a new instance with the resolved flow for this request + mcp_server = mcp_server.model_copy( + update={"oauth2_flow": resolved_flow} + ) allowed_mcp_servers.append(mcp_server) if mcp_servers is not None: @@ -1100,8 +1115,13 @@ if MCP_AVAILABLE: extra_headers: Optional[Dict[str, str]] = None if server.auth_type == MCPAuth.oauth2: - # Copy to avoid mutating the original dict (important for parallel fetching) - extra_headers = oauth2_headers.copy() if oauth2_headers else None + # For OAuth2 M2M servers, upstream Authorization must come from + # client_credentials token fetch, never from caller headers. + if server.has_client_credentials: + extra_headers = None + else: + # Copy to avoid mutating the original dict (important for parallel fetching) + extra_headers = oauth2_headers.copy() if oauth2_headers else None if server.extra_headers and raw_headers: if extra_headers is None: @@ -1114,11 +1134,17 @@ if MCP_AVAILABLE: for header in server.extra_headers: if not isinstance(header, str): continue + if server.has_client_credentials and header.lower() == "authorization": + continue header_value = normalized_raw_headers.get(header.lower()) if header_value is None: continue extra_headers[header] = header_value + # Reset to None if no headers were actually added + if extra_headers is not None and len(extra_headers) == 0: + extra_headers = None + if server_auth_header is None: server_auth_header = mcp_auth_header @@ -1377,11 +1403,19 @@ if MCP_AVAILABLE: spend_meta["per_server_tool_counts"] = per_server_tool_counts end_time = datetime.now() - await litellm_logging_obj.async_success_handler( - result=all_tools, - start_time=list_tools_start_time, - end_time=end_time, - ) + try: + await litellm_logging_obj.async_success_handler( + result=all_tools, + start_time=list_tools_start_time, + end_time=end_time, + ) + except Exception as log_exc: + # list_tools responses must not be dropped due to non-blocking + # observability/serialization failures. + verbose_logger.warning( + "MCP list_tools success logging failed (continuing): %s", + log_exc, + ) verbose_logger.info( f"Successfully fetched {len(all_tools)} tools total from all MCP servers" 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/skills.html b/litellm/proxy/_experimental/out/skills/index.html similarity index 100% rename from litellm/proxy/_experimental/out/skills.html rename to litellm/proxy/_experimental/out/skills/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/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 8331f748c6e..46a514c0870 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -3572,7 +3572,7 @@ "/anthropic/{endpoint}": { "delete": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__put", + "operationId": "anthropic_proxy_route_anthropic__endpoint__delete", "parameters": [ { "in": "path", @@ -3616,7 +3616,7 @@ }, "get": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__put", + "operationId": "anthropic_proxy_route_anthropic__endpoint__get", "parameters": [ { "in": "path", @@ -3660,7 +3660,7 @@ }, "patch": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__put", + "operationId": "anthropic_proxy_route_anthropic__endpoint__patch", "parameters": [ { "in": "path", @@ -3704,7 +3704,7 @@ }, "post": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__put", + "operationId": "anthropic_proxy_route_anthropic__endpoint__post", "parameters": [ { "in": "path", @@ -13260,7 +13260,7 @@ "/langfuse/{endpoint}": { "delete": { "description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)", - "operationId": "langfuse_proxy_route_langfuse__endpoint__put", + "operationId": "langfuse_proxy_route_langfuse__endpoint__delete", "parameters": [ { "in": "path", @@ -13299,7 +13299,7 @@ }, "get": { "description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)", - "operationId": "langfuse_proxy_route_langfuse__endpoint__put", + "operationId": "langfuse_proxy_route_langfuse__endpoint__get", "parameters": [ { "in": "path", @@ -13338,7 +13338,7 @@ }, "patch": { "description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)", - "operationId": "langfuse_proxy_route_langfuse__endpoint__put", + "operationId": "langfuse_proxy_route_langfuse__endpoint__patch", "parameters": [ { "in": "path", @@ -13377,7 +13377,7 @@ }, "post": { "description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)", - "operationId": "langfuse_proxy_route_langfuse__endpoint__put", + "operationId": "langfuse_proxy_route_langfuse__endpoint__post", "parameters": [ { "in": "path", @@ -26883,7 +26883,7 @@ "/toolset/{toolset_name}/mcp": { "delete": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete", "parameters": [ { "in": "path", @@ -26922,7 +26922,7 @@ }, "get": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_get", "parameters": [ { "in": "path", @@ -26961,7 +26961,7 @@ }, "head": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_head", "parameters": [ { "in": "path", @@ -27000,7 +27000,7 @@ }, "options": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options", "parameters": [ { "in": "path", @@ -27039,7 +27039,7 @@ }, "patch": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_patch", "parameters": [ { "in": "path", @@ -27078,7 +27078,7 @@ }, "post": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_post", "parameters": [ { "in": "path", diff --git a/litellm/proxy/_lazy_openapi_snapshot.py b/litellm/proxy/_lazy_openapi_snapshot.py index 315f6a9742a..309a0276aac 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.py +++ b/litellm/proxy/_lazy_openapi_snapshot.py @@ -13,6 +13,7 @@ from pathlib import Path from typing import Dict, Optional SNAPSHOT_FILE = Path(__file__).parent / "_lazy_openapi_snapshot.json" +HTTP_METHODS = {"delete", "get", "head", "options", "patch", "post", "put"} def load_snapshot() -> Optional[Dict[str, Dict]]: @@ -25,6 +26,39 @@ def load_snapshot() -> Optional[Dict[str, Dict]]: return None +def _normalize_operation_ids(paths: Dict[str, Dict]) -> None: + """Make FastAPI-generated operation IDs stable for multi-method routes. + + FastAPI derives the default operation ID suffix from the first item in the + route's methods set. For routes registered with several HTTP methods, that + set iteration order can vary between processes, which makes the snapshot + drift even when no routes changed. + """ + for path_ops in paths.values(): + if not isinstance(path_ops, dict): + continue + + methods = {method for method in path_ops if method in HTTP_METHODS} + if not methods: + continue + + for method, operation in path_ops.items(): + if method not in HTTP_METHODS or not isinstance(operation, dict): + continue + + operation_id = operation.get("operationId") + if not isinstance(operation_id, str): + continue + + for suffix in methods: + suffix_token = f"_{suffix}" + if operation_id.endswith(suffix_token): + operation["operationId"] = ( + operation_id[: -len(suffix_token)] + f"_{method}" + ) + break + + def generate_snapshot() -> Dict[str, Dict]: import importlib @@ -52,13 +86,15 @@ def generate_snapshot() -> Dict[str, Dict]: if not feat_routes: continue full = get_openapi(title=app.title, version=app.version, routes=feat_routes) + paths = full.get("paths", {}) + _normalize_operation_ids(paths) # Group all of a feature's routes under one tag. - for path_ops in full.get("paths", {}).values(): + for path_ops in paths.values(): for op in path_ops.values(): if isinstance(op, dict): op["tags"] = [feat.name] fragments[feat.name] = { - "paths": full.get("paths", {}), + "paths": paths, "components": {"schemas": full.get("components", {}).get("schemas", {})}, } return fragments diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 85320996911..8520e03f834 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -668,6 +668,8 @@ class LiteLLMRoutes(enum.Enum): "/models/{model_id}", "/guardrails/list", "/v2/guardrails/list", + "/project/list", + "/project/info", ] + spend_tracking_routes + key_management_routes @@ -692,6 +694,9 @@ class LiteLLMRoutes(enum.Enum): "/model/{model_id}/update", "/prompt/list", "/prompt/info", + # Project read routes - endpoint scopes results to caller's teams (non-admin) + "/project/list", + "/project/info", # Invitation routes - org/team admins checked in endpoint via _user_has_admin_privileges "/invitation/new", "/invitation/delete", diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 65638ed6c1e..113a8f538c0 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -12,14 +12,13 @@ Run checks for: import asyncio import re import time -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast +from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Type, Union, cast from fastapi import HTTPException, Request, status from pydantic import BaseModel import litellm from litellm._logging import verbose_proxy_logger -from litellm.caching.caching import DualCache from litellm.caching.dual_cache import LimitedSizeOrderedDict from litellm.constants import ( CLI_JWT_EXPIRATION_HOURS, @@ -66,6 +65,8 @@ from litellm.proxy.guardrails.tool_name_extraction import ( TOOL_CAPABLE_CALL_TYPES, extract_request_tool_names, ) +from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.route_llm_request import route_request from litellm.proxy.utils import PrismaClient, ProxyLogging, log_db_metrics from litellm.router import Router @@ -852,7 +853,7 @@ def get_actual_routes(allowed_routes: list) -> list: async def get_default_end_user_budget( prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, ) -> Optional[LiteLLM_BudgetTable]: """ @@ -875,9 +876,12 @@ async def get_default_end_user_budget( cache_key = f"default_end_user_budget:{litellm.max_end_user_budget_id}" # Check cache first - cached_budget = await user_api_key_cache.async_get_cache(key=cache_key) + cached_budget = await user_api_key_cache.async_get_cache( + key=cache_key, + model_type=LiteLLM_BudgetTable, + ) if cached_budget is not None: - return LiteLLM_BudgetTable(**cached_budget) + return cached_budget # Fetch from database try: @@ -891,14 +895,16 @@ async def get_default_end_user_budget( ) return None + _budget_obj = LiteLLM_BudgetTable(**budget_record.dict()) # Cache the budget for 60 seconds await user_api_key_cache.async_set_cache( key=cache_key, - value=budget_record.dict(), + value=_budget_obj, + model_type=LiteLLM_BudgetTable, ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) - return LiteLLM_BudgetTable(**budget_record.dict()) + return _budget_obj except Exception as e: verbose_proxy_logger.error(f"Error fetching default end user budget: {str(e)}") @@ -909,7 +915,7 @@ async def get_default_end_user_budget( async def get_team_member_default_budget( budget_id: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, ) -> Optional[LiteLLM_BudgetTable]: """ Fetches the team-level default per-member budget referenced by team.metadata["team_member_budget_id"]. @@ -966,7 +972,7 @@ async def get_team_member_default_budget( async def _apply_default_budget_to_end_user( end_user_obj: LiteLLM_EndUserTable, prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, ) -> LiteLLM_EndUserTable: """ @@ -1039,7 +1045,7 @@ def _check_end_user_budget( async def get_end_user_object( end_user_id: Optional[str], prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, route: str, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, @@ -1070,10 +1076,12 @@ async def get_end_user_object( _key = "end_user_id:{}".format(end_user_id) # Check cache first - cached_user_obj = await user_api_key_cache.async_get_cache(key=_key) + cached_user_obj = await user_api_key_cache.async_get_cache( + key=_key, + model_type=LiteLLM_EndUserTable, + ) if cached_user_obj is not None: - return_obj = LiteLLM_EndUserTable(**cached_user_obj) - + return_obj = cached_user_obj # Apply default budget if needed return_obj = await _apply_default_budget_to_end_user( end_user_obj=return_obj, @@ -1108,9 +1116,11 @@ async def get_end_user_object( parent_otel_span=parent_otel_span, ) - # Save to cache (always store as dict for consistency) + # Save to cache await user_api_key_cache.async_set_cache( - key="end_user_id:{}".format(end_user_id), value=_response.dict() + key="end_user_id:{}".format(end_user_id), + value=_response, + model_type=LiteLLM_EndUserTable, ) # Check budget limits @@ -1128,7 +1138,7 @@ async def get_end_user_object( async def get_tag_objects_batch( tag_names: List[str], prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> Dict[str, LiteLLM_TagTable]: @@ -1161,12 +1171,12 @@ async def get_tag_objects_batch( # Try to get all tags from cache first for tag_name in tag_names: cache_key = f"tag:{tag_name}" - cached_tag = await user_api_key_cache.async_get_cache(key=cache_key) + cached_tag = await user_api_key_cache.async_get_cache( + key=cache_key, + model_type=LiteLLM_TagTable, + ) if cached_tag is not None: - if isinstance(cached_tag, dict): - tag_objects[tag_name] = LiteLLM_TagTable(**cached_tag) - else: - tag_objects[tag_name] = cached_tag + tag_objects[tag_name] = cached_tag else: uncached_tags.append(tag_name) @@ -1182,11 +1192,13 @@ async def get_tag_objects_batch( for db_tag in db_tags: tag_name = db_tag.tag_name cache_key = f"tag:{tag_name}" - # Cache with default TTL (same as end_user objects) + _tag_obj = LiteLLM_TagTable(**db_tag.dict()) await user_api_key_cache.async_set_cache( - key=cache_key, value=db_tag.dict() + key=cache_key, + value=_tag_obj, + model_type=LiteLLM_TagTable, ) - tag_objects[tag_name] = LiteLLM_TagTable(**db_tag.dict()) + tag_objects[tag_name] = _tag_obj except Exception as e: verbose_proxy_logger.debug(f"Error batch fetching tags from database: {e}") @@ -1197,7 +1209,7 @@ async def get_tag_objects_batch( async def get_tag_object( tag_name: Optional[str], prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> Optional[LiteLLM_TagTable]: @@ -1236,7 +1248,7 @@ async def get_team_membership( user_id: str, team_id: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> Optional["LiteLLM_TeamMembership"]: @@ -1256,9 +1268,12 @@ async def get_team_membership( _key = "team_membership:{}:{}".format(user_id, team_id) # check if in cache - cached_membership_obj = await user_api_key_cache.async_get_cache(key=_key) + cached_membership_obj = await user_api_key_cache.async_get_cache( + key=_key, + model_type=LiteLLM_TeamMembership, + ) if cached_membership_obj is not None: - return LiteLLM_TeamMembership(**cached_membership_obj) + return cached_membership_obj # else, check db try: @@ -1270,10 +1285,12 @@ async def get_team_membership( if response is None: return None - # save the team membership object to cache (store as dict) - await user_api_key_cache.async_set_cache(key=_key, value=response.dict()) - _response = LiteLLM_TeamMembership(**response.dict()) + await user_api_key_cache.async_set_cache( + key=_key, + value=_response, + model_type=LiteLLM_TeamMembership, + ) return _response except Exception: @@ -1441,7 +1458,7 @@ async def _get_fuzzy_user_object( async def get_user_object( user_id: Optional[str], prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, user_id_upsert: bool, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, @@ -1460,12 +1477,12 @@ async def get_user_object( # check if in cache if not check_db_only: - cached_user_obj = await user_api_key_cache.async_get_cache(key=user_id) + cached_user_obj = await user_api_key_cache.async_get_cache( + key=user_id, + model_type=LiteLLM_UserTable, + ) if cached_user_obj is not None: - if isinstance(cached_user_obj, dict): - return LiteLLM_UserTable(**cached_user_obj) - elif isinstance(cached_user_obj, LiteLLM_UserTable): - return cached_user_obj + return cached_user_obj # else, check db if prisma_client is None: raise Exception("No db connected") @@ -1527,7 +1544,8 @@ async def get_user_object( # save the user object to cache await user_api_key_cache.async_set_cache( key=user_id, - value=response_dict, + value=_response, + model_type=LiteLLM_UserTable, ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) @@ -1548,13 +1566,21 @@ async def get_user_object( async def _cache_management_object( key: str, - value: BaseModel, - user_api_key_cache: DualCache, + value: Union[BaseModel, Dict[str, Any]], + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: Optional[ProxyLogging], + *, + model_type: Type[BaseModel], ): + """ + Persist management objects via ``UserApiKeyCache`` (in-memory + optional Redis). + + ``UserApiKeyCache`` serializes with ``model_type`` so Redis and in-memory stay aligned. + """ await user_api_key_cache.async_set_cache( key=key, value=value, + model_type=model_type, ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) @@ -1562,7 +1588,7 @@ async def _cache_management_object( async def _cache_team_object( team_id: str, team_table: LiteLLM_TeamTableCachedObj, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: Optional[ProxyLogging], ): key = "team_id:{}".format(team_id) @@ -1575,13 +1601,14 @@ async def _cache_team_object( value=team_table, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + model_type=LiteLLM_TeamTableCachedObj, ) async def _cache_key_object( hashed_token: str, user_api_key_obj: UserAPIKeyAuth, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: Optional[ProxyLogging], ): key = hashed_token @@ -1594,12 +1621,13 @@ async def _cache_key_object( value=user_api_key_obj, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + model_type=UserAPIKeyAuth, ) async def _delete_cache_key_object( hashed_token: str, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: Optional[ProxyLogging], ): key = hashed_token @@ -1647,7 +1675,7 @@ async def _get_team_object_from_db(team_id: str, prisma_client: PrismaClient): async def _get_team_object_from_user_api_key_cache( team_id: str, prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, last_db_access_time: LimitedSizeOrderedDict, db_cache_expiry: int, proxy_logging_obj: Optional[ProxyLogging], @@ -1708,38 +1736,38 @@ async def _get_team_object_from_user_api_key_cache( async def _get_team_object_from_cache( key: str, proxy_logging_obj: Optional[ProxyLogging], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span], ) -> Optional[LiteLLM_TeamTableCachedObj]: - cached_team_obj: Optional[LiteLLM_TeamTableCachedObj] = None - - ## CHECK REDIS CACHE ## + ## INTERNAL USAGE CACHE (plain DualCache) — checked before UserApiKeyCache stores ## if ( proxy_logging_obj is not None and proxy_logging_obj.internal_usage_cache.dual_cache ): - cached_team_obj = ( + cached_raw = ( await proxy_logging_obj.internal_usage_cache.dual_cache.async_get_cache( key=key, parent_otel_span=parent_otel_span ) ) + if cached_raw is not None: + from_internal = CacheCodec.deserialize( + cached_raw, LiteLLM_TeamTableCachedObj + ) + if from_internal is not None: + return from_internal - if cached_team_obj is None: - cached_team_obj = await user_api_key_cache.async_get_cache(key=key) - - if cached_team_obj is not None: - if isinstance(cached_team_obj, dict): - return LiteLLM_TeamTableCachedObj(**cached_team_obj) - elif isinstance(cached_team_obj, LiteLLM_TeamTableCachedObj): - return cached_team_obj - - return None + decoded = await user_api_key_cache.async_get_cache( + key=key, + parent_otel_span=parent_otel_span, + model_type=LiteLLM_TeamTableCachedObj, + ) + return decoded async def get_team_object( team_id: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, check_cache_only: Optional[bool] = None, @@ -1805,20 +1833,21 @@ async def get_team_object( async def _cache_access_object( access_group_id: str, access_group_table: LiteLLM_AccessGroupTable, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: Optional[ProxyLogging] = None, ): key = "access_group_id:{}".format(access_group_id) await user_api_key_cache.async_set_cache( key=key, value=access_group_table, + model_type=LiteLLM_AccessGroupTable, ttl=DEFAULT_ACCESS_GROUP_CACHE_TTL, ) async def _delete_cache_access_object( access_group_id: str, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: Optional[ProxyLogging] = None, ): key = "access_group_id:{}".format(access_group_id) @@ -1836,7 +1865,7 @@ async def _delete_cache_access_object( async def get_access_object( access_group_id: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> LiteLLM_AccessGroupTable: """ @@ -1858,13 +1887,12 @@ async def get_access_object( key = "access_group_id:{}".format(access_group_id) - # Always check cache first - cached_access_obj = await user_api_key_cache.async_get_cache(key=key) + cached_access_obj = await user_api_key_cache.async_get_cache( + key=key, + model_type=LiteLLM_AccessGroupTable, + ) if cached_access_obj is not None: - if isinstance(cached_access_obj, dict): - return LiteLLM_AccessGroupTable(**cached_access_obj) - elif isinstance(cached_access_obj, LiteLLM_AccessGroupTable): - return cached_access_obj + return cached_access_obj # Not in cache - fetch from DB try: @@ -1910,7 +1938,7 @@ async def get_access_object( async def get_team_object_by_alias( team_alias: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional["Span"] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> LiteLLM_TeamTableCachedObj: @@ -1992,6 +2020,7 @@ async def get_team_object_by_alias( await user_api_key_cache.async_set_cache( key=cache_key, value=team_obj, + model_type=LiteLLM_TeamTableCachedObj, ttl=DEFAULT_IN_MEMORY_TTL, ) # Also cache by team_id for consistency @@ -1999,6 +2028,7 @@ async def get_team_object_by_alias( await user_api_key_cache.async_set_cache( key=team_id_cache_key, value=team_obj, + model_type=LiteLLM_TeamTableCachedObj, ttl=DEFAULT_IN_MEMORY_TTL, ) @@ -2020,7 +2050,7 @@ async def get_team_object_by_alias( async def get_org_object_by_alias( org_alias: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional["Span"] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> Optional[LiteLLM_OrganizationTable]: @@ -2047,12 +2077,12 @@ async def get_org_object_by_alias( # Check cache first (keyed by alias) cache_key = "org_alias:{}".format(org_alias) - cached_org_obj = await user_api_key_cache.async_get_cache(key=cache_key) + cached_org_obj = await user_api_key_cache.async_get_cache( + key=cache_key, + model_type=LiteLLM_OrganizationTable, + ) if cached_org_obj is not None: - if isinstance(cached_org_obj, dict): - return LiteLLM_OrganizationTable(**cached_org_obj) - elif isinstance(cached_org_obj, LiteLLM_OrganizationTable): - return cached_org_obj + return cached_org_obj # Query database by organization_alias try: @@ -2082,13 +2112,15 @@ async def get_org_object_by_alias( # Cache the result await user_api_key_cache.async_set_cache( key=cache_key, - value=org_obj.model_dump(), + value=org_obj, + model_type=LiteLLM_OrganizationTable, ttl=DEFAULT_IN_MEMORY_TTL, ) # Also cache by org_id for consistency await user_api_key_cache.async_set_cache( key="org_id:{}".format(org_obj.organization_id), - value=org_obj.model_dump(), + value=org_obj, + model_type=LiteLLM_OrganizationTable, ttl=DEFAULT_IN_MEMORY_TTL, ) @@ -2291,7 +2323,7 @@ async def get_jwt_key_mapping_object( async def get_key_object( hashed_token: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, check_cache_only: Optional[bool] = None, @@ -2309,15 +2341,14 @@ async def get_key_object( # check if in cache key = hashed_token - cached_key_obj: Optional[UserAPIKeyAuth] = await user_api_key_cache.async_get_cache( - key=key + # Same flow as before: use cache only when we have a hit we can turn into UserAPIKeyAuth + # (dict from Redis / model_dump, or UserAPIKeyAuth from in-memory). Otherwise fall through to DB. + user_api_key_auth = await user_api_key_cache.async_get_cache( + key=key, + model_type=UserAPIKeyAuth, ) - - if cached_key_obj is not None: - if isinstance(cached_key_obj, dict): - return UserAPIKeyAuth(**cached_key_obj) - elif isinstance(cached_key_obj, UserAPIKeyAuth): - return cached_key_obj + if user_api_key_auth is not None: + return user_api_key_auth if check_cache_only: raise Exception( @@ -2374,7 +2405,7 @@ async def get_key_object( async def get_object_permission( object_permission_id: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> Optional[LiteLLM_ObjectPermissionTable]: @@ -2390,12 +2421,12 @@ async def get_object_permission( # check if in cache key = "object_permission_id:{}".format(object_permission_id) - cached_obj_permission = await user_api_key_cache.async_get_cache(key=key) - if cached_obj_permission is not None: - if isinstance(cached_obj_permission, dict): - return LiteLLM_ObjectPermissionTable(**cached_obj_permission) - elif isinstance(cached_obj_permission, LiteLLM_ObjectPermissionTable): - return cached_obj_permission + deserialized_perm = await user_api_key_cache.async_get_cache( + key=key, + model_type=LiteLLM_ObjectPermissionTable, + ) + if deserialized_perm is not None: + return deserialized_perm # else, check db try: @@ -2406,14 +2437,15 @@ async def get_object_permission( if response is None: return None - # save the object permission to cache + _perm_obj = LiteLLM_ObjectPermissionTable(**response.dict()) await user_api_key_cache.async_set_cache( key=key, - value=response.model_dump(), + value=_perm_obj, + model_type=LiteLLM_ObjectPermissionTable, ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) - return LiteLLM_ObjectPermissionTable(**response.dict()) + return _perm_obj except Exception: return None @@ -2422,7 +2454,7 @@ async def get_object_permission( async def get_managed_vector_store_rows_by_uuids( uuids: List[str], prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> List[LiteLLM_ManagedVectorStoresTable]: @@ -2442,14 +2474,12 @@ async def get_managed_vector_store_rows_by_uuids( for uuid in uuids: key = "managed_vector_store_id:{}".format(uuid) - cached = await user_api_key_cache.async_get_cache(key=key) - if cached is not None: - if isinstance(cached, dict): - result.append(LiteLLM_ManagedVectorStoresTable(**cached)) - elif isinstance(cached, LiteLLM_ManagedVectorStoresTable): - result.append(cached) - else: - cache_misses.append(uuid) + deserialized_vs = await user_api_key_cache.async_get_cache( + key=key, + model_type=LiteLLM_ManagedVectorStoresTable, + ) + if deserialized_vs is not None: + result.append(deserialized_vs) else: cache_misses.append(uuid) @@ -2475,7 +2505,8 @@ async def get_managed_vector_store_rows_by_uuids( key = "managed_vector_store_id:{}".format(cached_obj.vector_store_id) await user_api_key_cache.async_set_cache( key=key, - value=row_dict, + value=cached_obj, + model_type=LiteLLM_ManagedVectorStoresTable, ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) result.append(cached_obj) @@ -2487,7 +2518,7 @@ async def get_managed_vector_store_rows_by_uuids( async def get_org_object( org_id: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, include_budget_table: bool = False, @@ -2518,12 +2549,12 @@ async def get_org_object( cache_key = "org_id:{}:with_budget".format(org_id) # check if in cache - cached_org_obj = user_api_key_cache.async_get_cache(key=cache_key) - if cached_org_obj is not None: - if isinstance(cached_org_obj, dict): - return LiteLLM_OrganizationTable(**cached_org_obj) - elif isinstance(cached_org_obj, LiteLLM_OrganizationTable): - return cached_org_obj + deserialized_org = await user_api_key_cache.async_get_cache( + key=cache_key, + model_type=LiteLLM_OrganizationTable, + ) + if deserialized_org is not None: + return deserialized_org # else, check db try: query_kwargs: Dict[str, Any] = {"where": {"organization_id": org_id}} @@ -2537,16 +2568,16 @@ async def get_org_object( if response is None: raise Exception + _org_obj = LiteLLM_OrganizationTable(**response.model_dump()) # Cache the result await user_api_key_cache.async_set_cache( key=cache_key, - value=( - response.model_dump() if hasattr(response, "model_dump") else response - ), + value=_org_obj, + model_type=LiteLLM_OrganizationTable, ttl=DEFAULT_IN_MEMORY_TTL, ) - return response + return _org_obj except Exception: raise Exception( f"Organization doesn't exist in db. Organization={org_id}. Create organization via `/organization/new` call." @@ -2559,7 +2590,7 @@ async def _get_resources_from_access_groups( "access_model_names", "access_mcp_server_ids", "access_agent_ids" ], prisma_client: Optional[PrismaClient] = None, - user_api_key_cache: Optional[DualCache] = None, + user_api_key_cache: Optional[UserApiKeyCache] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> List[str]: """ @@ -2617,7 +2648,7 @@ async def _get_resources_from_access_groups( async def _get_models_from_access_groups( access_group_ids: List[str], prisma_client: Optional[PrismaClient] = None, - user_api_key_cache: Optional[DualCache] = None, + user_api_key_cache: Optional[UserApiKeyCache] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> List[str]: """ @@ -2636,7 +2667,7 @@ async def _get_models_from_access_groups( async def _get_mcp_server_ids_from_access_groups( access_group_ids: List[str], prisma_client: Optional[PrismaClient] = None, - user_api_key_cache: Optional[DualCache] = None, + user_api_key_cache: Optional[UserApiKeyCache] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> List[str]: """ @@ -2655,7 +2686,7 @@ async def _get_mcp_server_ids_from_access_groups( async def _get_agent_ids_from_access_groups( access_group_ids: List[str], prisma_client: Optional[PrismaClient] = None, - user_api_key_cache: Optional[DualCache] = None, + user_api_key_cache: Optional[UserApiKeyCache] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> List[str]: """ @@ -3379,7 +3410,7 @@ async def _check_team_member_budget( user_object: Optional[LiteLLM_UserTable], valid_token: Optional[UserAPIKeyAuth], prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, ): """Check if team member is over their max budget within the team.""" @@ -3447,7 +3478,7 @@ async def _check_team_member_model_access( valid_token: UserAPIKeyAuth, llm_router: Optional[Router], prisma_client: Optional["PrismaClient"], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, ) -> None: """ @@ -3754,7 +3785,7 @@ async def _project_soft_budget_check( async def get_project_object( project_id: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> Optional[LiteLLM_ProjectTableCachedObj]: """ @@ -3769,12 +3800,12 @@ async def get_project_object( # Check cache first cache_key = "project_id:{}".format(project_id) - cached_obj = await user_api_key_cache.async_get_cache(key=cache_key) - if cached_obj is not None: - if isinstance(cached_obj, dict): - return LiteLLM_ProjectTableCachedObj(**cached_obj) - elif isinstance(cached_obj, LiteLLM_ProjectTableCachedObj): - return cached_obj + deserialized_project = await user_api_key_cache.async_get_cache( + key=cache_key, + model_type=LiteLLM_ProjectTableCachedObj, + ) + if deserialized_project is not None: + return deserialized_project # Fetch from DB project_row = await prisma_client.db.litellm_projecttable.find_unique( @@ -3793,6 +3824,7 @@ async def get_project_object( value=project_obj, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + model_type=LiteLLM_ProjectTableCachedObj, ) return project_obj @@ -3802,7 +3834,7 @@ async def _organization_max_budget_check( valid_token: Optional[UserAPIKeyAuth], team_object: Optional[LiteLLM_TeamTable], prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, ): """ @@ -3896,7 +3928,7 @@ async def _organization_max_budget_check( async def _tag_max_budget_check( request_body: dict, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, valid_token: Optional[UserAPIKeyAuth], ): diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 91c8f2dd7c9..cbed34adacf 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -167,6 +167,81 @@ def _allow_model_level_clientside_configurable_parameters( ) +# Config dicts whose entries are spread as ``**dict`` into outbound LLM +# API calls. ``litellm_embedding_config`` is consumed by the Milvus +# vector store transformer; future nested-config keys with the same +# threat shape should be added here. +_NESTED_CONFIG_KEYS: Tuple[str, ...] = ("litellm_embedding_config",) + +# Banned root-level params. Same list applies to every entry in +# ``_NESTED_CONFIG_KEYS`` because those dicts get spread as ``**kwargs`` +# into the same outbound calls. +_BANNED_REQUEST_BODY_PARAMS: Tuple[str, ...] = ( + "api_base", + "base_url", + "user_config", + "aws_sts_endpoint", + "aws_web_identity_token", + "aws_role_name", + "vertex_credentials", + # Endpoint-targeting fields that retarget the outbound request or + # an observability callback. An attacker-controlled value either + # exfiltrates the request payload (incl. messages + admin-set + # tokens) to the attacker's host, or coerces the proxy into + # authenticating against the attacker's host with admin secrets. + "aws_bedrock_runtime_endpoint", + "langsmith_base_url", + "langfuse_host", + "posthog_host", + "braintrust_host", + "slack_webhook_url", + # Provider-specific endpoint overrides that flow into the outbound + # request via ``optional_params``. Same threat as ``api_base``: + # ``s3_endpoint_url`` redirects Bedrock file uploads to attacker + # S3; ``sagemaker_base_url`` redirects all SageMaker traffic; + # ``deployment_url`` redirects SAP deployments. + "s3_endpoint_url", + "sagemaker_base_url", + "deployment_url", +) + + +def _check_banned_params( + body: dict, + general_settings: dict, + llm_router: Optional[Router], + model: str, +) -> None: + """Raise ``ValueError`` if ``body`` carries a banned param without admin opt-in. + + Shared between the root-level check and the nested-config check so a + new banned param only needs to be added in one place. + """ + for param in _BANNED_REQUEST_BODY_PARAMS: + if param not in body: + continue + if general_settings.get("allow_client_side_credentials") is True: + return + if ( + _allow_model_level_clientside_configurable_parameters( + model=model, + param=param, + request_body_value=body[param], + llm_router=llm_router, + ) + is True + ): + return + raise ValueError( + f"Rejected Request: {param} is not allowed in request body. " + "Clientside passthrough requires explicit admin opt-in via " + "either `general_settings.allow_client_side_credentials = true` " + "(proxy-wide) or `configurable_clientside_auth_params` on the " + "deployment in your proxy config.yaml. " + "Relevant Issue: https://huntr.com/bounties/4001e1a2-7b7a-4776-a3ae-e6692ec3d997", + ) + + def is_request_body_safe( request_body: dict, general_settings: dict, llm_router: Optional[Router], model: str ) -> bool: @@ -175,72 +250,31 @@ def is_request_body_safe( A malicious user can set the api_base to their own domain and invoke POST /chat/completions to intercept and steal the OpenAI API key. Relevant issue: https://huntr.com/bounties/4001e1a2-7b7a-4776-a3ae-e6692ec3d997 + + The blocklist is enforced unconditionally. Legitimate clientside + credential / endpoint passthrough goes through one of the two + explicit admin opt-ins (``general_settings.allow_client_side_credentials`` + proxy-wide or ``configurable_clientside_auth_params`` per deployment). + Historically there was a third, *implicit*, *caller-controlled* path: + ``check_complete_credentials`` returned True when the caller supplied + any non-empty ``api_key``, which made the entire blocklist a no-op. + That bypass turned every missing entry on the blocklist into an + exploitable SSRF / credential-exfil hole — see GHSA-jh89-88fc-qrfp, + GHSA-3frq-6r6h-7j64, and the chain of veria-admin findings (Dv_m860l, + b_yRJeQ5, stN90yjP, LBlyOAc8, U2TD78kg). Removed: the blocklist now + has a single, predictable failure mode for missing entries (a 400), + not a credential leak. + + Iterative single-level descent into ``_NESTED_CONFIG_KEYS`` (rather + than recursion) covers nested-config attacks like Milvus's + ``litellm_embedding_config.api_base`` (VERIA-6) without exposing a + recursion-depth DoS surface. """ - banned_params = [ - "api_base", - "base_url", - "user_config", - "aws_sts_endpoint", - "aws_web_identity_token", - "aws_role_name", - "vertex_credentials", - # Endpoint-targeting fields that retarget the outbound request or - # an observability callback. An attacker-controlled value either - # exfiltrates the request payload (incl. messages + admin-set - # tokens) to the attacker's host, or coerces the proxy into - # authenticating against the attacker's host with admin secrets. - "aws_bedrock_runtime_endpoint", - "langsmith_base_url", - "langfuse_host", - "posthog_host", - "braintrust_host", - "slack_webhook_url", - # Provider-specific endpoint overrides that flow into the outbound - # request via ``optional_params``. Same threat as ``api_base``: - # ``s3_endpoint_url`` redirects Bedrock file uploads to attacker - # S3; ``sagemaker_base_url`` redirects all SageMaker traffic; - # ``deployment_url`` redirects SAP deployments. - "s3_endpoint_url", - "sagemaker_base_url", - "deployment_url", - ] - - # The blocklist is enforced unconditionally. Legitimate clientside - # credential / endpoint passthrough goes through one of the two - # explicit admin opt-ins (``general_settings.allow_client_side_credentials`` - # proxy-wide or ``configurable_clientside_auth_params`` per deployment). - # Historically there was a third, *implicit*, *caller-controlled* path: - # ``check_complete_credentials`` returned True when the caller supplied - # any non-empty ``api_key``, which made the entire blocklist a no-op. - # That bypass turned every missing entry on the blocklist into an - # exploitable SSRF / credential-exfil hole — see GHSA-jh89-88fc-qrfp, - # GHSA-3frq-6r6h-7j64, and the chain of veria-admin findings (Dv_m860l, - # b_yRJeQ5, stN90yjP, LBlyOAc8, U2TD78kg). Removed: the blocklist now - # has a single, predictable failure mode for missing entries (a 400), - # not a credential leak. - for param in banned_params: - if param in request_body: - if general_settings.get("allow_client_side_credentials") is True: - return True - elif ( - _allow_model_level_clientside_configurable_parameters( - model=model, - param=param, - request_body_value=request_body[param], - llm_router=llm_router, - ) - is True - ): - return True - raise ValueError( - f"Rejected Request: {param} is not allowed in request body. " - "Clientside passthrough requires explicit admin opt-in via " - "either `general_settings.allow_client_side_credentials = true` " - "(proxy-wide) or `configurable_clientside_auth_params` on the " - "deployment in your proxy config.yaml. " - "Relevant Issue: https://huntr.com/bounties/4001e1a2-7b7a-4776-a3ae-e6692ec3d997", - ) - + _check_banned_params(request_body, general_settings, llm_router, model) + for nested_key in _NESTED_CONFIG_KEYS: + nested = request_body.get(nested_key) + if isinstance(nested, dict): + _check_banned_params(nested, general_settings, llm_router, model) return True diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index f50c950d747..71411bed7fd 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -6,6 +6,8 @@ Currently only supports admin. JWT token must have 'litellm_proxy_admin' in scope. """ +from __future__ import annotations + import fnmatch import hashlib import os @@ -20,7 +22,6 @@ import jwt from jwt.api_jwk import PyJWK from litellm._logging import verbose_proxy_logger -from litellm.caching.caching import DualCache from litellm.constants import DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value from litellm.llms.custom_httpx.httpx_handler import HTTPHandler @@ -46,6 +47,7 @@ from litellm.proxy._types import ( ) from litellm.proxy.auth.auth_checks import can_team_access_model from litellm.proxy.auth.route_checks import RouteChecks +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.utils import PrismaClient, ProxyLogging from .auth_checks import ( @@ -73,7 +75,7 @@ class JWTHandler: """ prisma_client: Optional[PrismaClient] - user_api_key_cache: DualCache + user_api_key_cache: UserApiKeyCache # Supported algos: https://pyjwt.readthedocs.io/en/stable/algorithms.html # "Warning: Make sure not to mix symmetric and asymmetric algorithms that interpret # the key in different ways (e.g. HS* and RS*)." @@ -99,7 +101,7 @@ class JWTHandler: def update_environment( self, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, litellm_jwtauth: LiteLLM_JWTAuth, leeway: int = 0, ) -> None: @@ -952,7 +954,7 @@ class JWTAuthManager: jwt_handler: JWTHandler, jwt_valid_token: dict, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span], proxy_logging_obj: ProxyLogging, ) -> Tuple[Optional[str], Optional[LiteLLM_TeamTable]]: @@ -1045,7 +1047,7 @@ class JWTAuthManager: route: str, jwt_handler: JWTHandler, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span], proxy_logging_obj: ProxyLogging, ) -> Tuple[Optional[str], Optional[LiteLLM_TeamTable]]: @@ -1133,7 +1135,7 @@ class JWTAuthManager: valid_user_email: Optional[bool], jwt_handler: JWTHandler, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span], proxy_logging_obj: ProxyLogging, route: str, @@ -1349,7 +1351,7 @@ class JWTAuthManager: jwt_valid_token: dict, user_object: Optional[LiteLLM_UserTable], prisma_client: Optional[PrismaClient], - user_api_key_cache: Optional[DualCache] = None, + user_api_key_cache: Optional[UserApiKeyCache] = None, ) -> None: """ Sync user role and team memberships with JWT claims @@ -1377,7 +1379,8 @@ class JWTAuthManager: if user_api_key_cache is not None: await user_api_key_cache.async_set_cache( key=user_object.user_id, - value=user_object.model_dump(), + value=user_object, + model_type=LiteLLM_UserTable, ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) @@ -1400,7 +1403,8 @@ class JWTAuthManager: if user_api_key_cache is not None: await user_api_key_cache.async_set_cache( key=user_object.user_id, - value=user_object.model_dump(), + value=user_object, + model_type=LiteLLM_UserTable, ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) return None @@ -1412,7 +1416,7 @@ class JWTAuthManager: request_headers: Optional[dict], jwt_handler: JWTHandler, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span], proxy_logging_obj: ProxyLogging, ) -> None: @@ -1456,7 +1460,7 @@ class JWTAuthManager: user_object: Optional[LiteLLM_UserTable], user_id: Optional[str], prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span], proxy_logging_obj: ProxyLogging, team_id_upsert: Optional[bool], @@ -1514,7 +1518,7 @@ class JWTAuthManager: general_settings: dict, route: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span], proxy_logging_obj: ProxyLogging, request_headers: Optional[dict] = None, diff --git a/litellm/proxy/auth/ip_address_utils.py b/litellm/proxy/auth/ip_address_utils.py index 34fab4849e5..39d3282942f 100644 --- a/litellm/proxy/auth/ip_address_utils.py +++ b/litellm/proxy/auth/ip_address_utils.py @@ -13,6 +13,10 @@ from fastapi import Request from litellm._logging import verbose_proxy_logger from litellm.proxy.auth.auth_utils import _get_request_ip_address +# One-shot warning so operators upgrading from the prior "always trust X-Forwarded-*" +# behaviour see an actionable message in their logs the first time it triggers. +_warned_xff_without_trusted_ranges = False + class IPAddressUtils: """Static utilities for IP-based MCP access control.""" @@ -106,6 +110,61 @@ class IPAddressUtils: return any(addr in network for network in networks) + @staticmethod + def is_request_from_trusted_proxy( + request: Request, + general_settings: Optional[Dict[str, Any]] = None, + ) -> bool: + """ + Return True if X-Forwarded-* headers on this request should be trusted. + + Trusts the headers iff both: + 1. ``use_x_forwarded_for`` is enabled in proxy settings, AND + 2. ``mcp_trusted_proxy_ranges`` is configured AND the direct + connection IP (``request.client.host``) falls inside one of + those CIDRs. + + When ``use_x_forwarded_for`` is enabled but ``mcp_trusted_proxy_ranges`` + is missing, the headers are NOT trusted: there is no way to + distinguish a trusted reverse proxy from a direct attacker, so callers + that build URLs (OAuth issuer / redirect_uri / etc.) must fall back + to the request's literal base URL instead of risking a poisoned host. + """ + if general_settings is None: + try: + from litellm.proxy.proxy_server import ( + general_settings as proxy_general_settings, + ) + + general_settings = proxy_general_settings + except ImportError: + general_settings = {} + + if general_settings is None: + general_settings = {} + + if not general_settings.get("use_x_forwarded_for", False): + return False + + trusted_ranges = general_settings.get("mcp_trusted_proxy_ranges") + if not trusted_ranges: + global _warned_xff_without_trusted_ranges + if not _warned_xff_without_trusted_ranges: + verbose_proxy_logger.warning( + "use_x_forwarded_for is enabled but mcp_trusted_proxy_ranges " + "is not configured. X-Forwarded-* headers will NOT be " + "trusted, so MCP OAuth discovery URLs will use the proxy's " + "literal base URL. Set mcp_trusted_proxy_ranges in " + "general_settings to your reverse-proxy CIDR(s) to allow " + "X-Forwarded-* through." + ) + _warned_xff_without_trusted_ranges = True + return False + + direct_ip = request.client.host if request.client else None + trusted_networks = IPAddressUtils.parse_trusted_proxy_networks(trusted_ranges) + return IPAddressUtils.is_trusted_proxy(direct_ip, trusted_networks) + @staticmethod def get_mcp_client_ip( request: Request, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index b7700feb5bb..bfd1f2e0b3a 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -20,7 +20,6 @@ from fastapi.security.api_key import APIKeyHeader import litellm from litellm._logging import verbose_logger, verbose_proxy_logger from litellm._service_logger import ServiceLogging -from litellm.caching import DualCache from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value @@ -60,6 +59,7 @@ from litellm.proxy.auth.oauth2_check import Oauth2Handler from litellm.proxy.auth.oauth2_proxy_hook import handle_oauth2_proxy_request from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.common_utils.cache_coordinator import EventDrivenCacheCoordinator +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.common_utils.http_parsing_utils import ( _read_request_body, _safe_get_request_headers, @@ -329,7 +329,7 @@ _global_spend_coordinator = EventDrivenCacheCoordinator(log_prefix="[GLOBAL SPEN async def _fetch_global_spend_with_event_coordination( cache_key: str, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, prisma_client: PrismaClient, ) -> Optional[float]: """ @@ -345,14 +345,14 @@ async def _fetch_global_spend_with_event_coordination( return await _global_spend_coordinator.get_or_load( cache_key=cache_key, - cache=user_api_key_cache, + cache=user_api_key_cache, # pyright: ignore[reportArgumentType] load_fn=_load_global_spend, ) async def get_global_proxy_spend( litellm_proxy_admin_name: str, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, prisma_client: Optional[PrismaClient], token: str, proxy_logging_obj: ProxyLogging, @@ -510,7 +510,7 @@ async def _resolve_jwt_to_virtual_key( jwt_claims: dict, jwt_handler: JWTHandler, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span], proxy_logging_obj: ProxyLogging, ) -> Optional[UserAPIKeyAuth]: @@ -1112,9 +1112,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 is_master_key_valid = False ## VALIDATE MASTER KEY ## - try: - assert isinstance(master_key, str) - except Exception: + if not isinstance(master_key, str): raise HTTPException( status_code=500, detail={ @@ -1184,11 +1182,15 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 if len(api_key) > 8 else "****" ) - assert api_key.startswith( - "sk-" - ), "LiteLLM Virtual Key expected. Received={}, expected to start with 'sk-'.".format( - _masked_key - ) # prevent token hashes from being used + if not api_key.startswith("sk-"): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=( + "LiteLLM Virtual Key expected. Received={}, expected to start with 'sk-'.".format( + _masked_key + ) + ), + ) # prevent token hashes from being used else: verbose_logger.warning( "litellm.proxy.proxy_server.user_api_key_auth(): Warning - Key is not a string. Got type={}".format( @@ -1296,7 +1298,8 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 _cache_key = f"{valid_token.team_id}_{valid_token.user_id}" team_member_info = await user_api_key_cache.async_get_cache( - key=_cache_key + key=_cache_key, + model_type=LiteLLM_TeamMembership, ) if team_member_info is None: # read from DB @@ -1304,18 +1307,23 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 _team_id = valid_token.team_id if _user_id is not None and _team_id is not None: - team_member_info = await prisma_client.db.litellm_teammembership.find_first( + _db_member = await prisma_client.db.litellm_teammembership.find_first( where={ "user_id": _user_id, "team_id": _team_id, }, # type: ignore include={"litellm_budget_table": True}, ) - await user_api_key_cache.async_set_cache( - key=_cache_key, - value=team_member_info, - ttl=5, - ) + if _db_member is not None: + team_member_info = LiteLLM_TeamMembership( + **_db_member.dict() + ) + await user_api_key_cache.async_set_cache( + key=_cache_key, + value=team_member_info, + model_type=LiteLLM_TeamMembership, + ttl=5, + ) if ( team_member_info is not None @@ -1462,9 +1470,13 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 else: valid_token.team_object_permission = None - await user_api_key_cache.async_set_cache( - key=valid_token.team_id, value=_team_obj - ) # save team table in cache - used for tpm/rpm limiting - tpm_rpm_limiter.py + # Only cache when the key is a real team_id (non-team keys must not use key=None). + if valid_token.team_id is not None and _team_obj is not None: + await user_api_key_cache.async_set_cache( + key=valid_token.team_id, + value=_team_obj, + model_type=LiteLLM_TeamTableCachedObj, + ) # save team table in cache - used for tpm/rpm limiting - tpm_rpm_limiter.py # Fetch project object if key belongs to a project _project_obj = None diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py index a9ea7a84e18..447837c35e7 100644 --- a/litellm/proxy/client/cli/commands/auth.py +++ b/litellm/proxy/client/cli/commands/auth.py @@ -53,12 +53,16 @@ def clear_token() -> None: os.remove(token_file) -def get_stored_api_key() -> Optional[str]: - """Get the stored API key from token file""" - # Use the SDK-level utility +def get_stored_api_key(expected_base_url: Optional[str] = None) -> Optional[str]: + """Get the stored API key from token file. + + If expected_base_url is provided, the key is only returned when it was + originally issued for that URL. This prevents credential leakage when the + CLI is pointed at a different (possibly malicious) server. + """ from litellm.litellm_core_utils.cli_token_utils import get_litellm_gateway_api_key - return get_litellm_gateway_api_key() + return get_litellm_gateway_api_key(expected_base_url=expected_base_url) # Team selection utilities @@ -572,9 +576,11 @@ def login(ctx: click.Context): api_key = auth_result["api_key"] user_id = auth_result["user_id"] - # Save token data (simplified for CLI - we just need the key) + # Save token data. base_url is stored so we can verify origin + # before reusing the key on a subsequent CLI invocation. save_token( { + "base_url": base_url.rstrip("/"), "key": api_key, "user_id": user_id or "cli-user", "user_email": "unknown", diff --git a/litellm/proxy/client/cli/main.py b/litellm/proxy/client/cli/main.py index 22de5a78614..be55f79c066 100644 --- a/litellm/proxy/client/cli/main.py +++ b/litellm/proxy/client/cli/main.py @@ -74,9 +74,10 @@ def cli(ctx: click.Context, base_url: str, api_key: Optional[str]) -> None: """LiteLLM Proxy CLI - Manage your LiteLLM proxy server""" ctx.ensure_object(dict) - # If no API key provided via flag or environment variable, try to load from saved token + # If no API key provided via flag or environment variable, try to load from saved token. + # Pass base_url so we only use the stored key when it was issued for this server. if api_key is None: - api_key = get_stored_api_key() + api_key = get_stored_api_key(expected_base_url=base_url) ctx.obj["base_url"] = base_url ctx.obj["api_key"] = api_key diff --git a/litellm/proxy/client/client.py b/litellm/proxy/client/client.py index 12b5cd79f79..929ad46a77c 100644 --- a/litellm/proxy/client/client.py +++ b/litellm/proxy/client/client.py @@ -28,12 +28,17 @@ class Client: api_key (Optional[str]): API key for authentication. If provided, it will be sent as a Bearer token. timeout: Request timeout in seconds (default: 30) """ - self._base_url = base_url.rstrip("/") # Remove trailing slash if present - self._api_key = get_litellm_gateway_api_key() or api_key + self._base_url = base_url.rstrip("/") + # Only use the stored CLI key when it was issued for this server. + self._api_key = api_key or get_litellm_gateway_api_key( + expected_base_url=self._base_url + ) # Initialize resource clients - self.http = HTTPClient(base_url=base_url, api_key=api_key, timeout=timeout) + self.http = HTTPClient( + base_url=base_url, api_key=self._api_key, timeout=timeout + ) self.models = ModelsManagementClient( base_url=self._base_url, api_key=self._api_key ) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 76c52f83ee4..438da037527 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -744,6 +744,11 @@ class ProxyBaseLLMRequestProcessing: "aingest", "aretrieve_container", "adelete_container", + "aupload_container_file", + "alist_container_files", + "aretrieve_container_file", + "adelete_container_file", + "aretrieve_container_file_content", "acreate_skill", "alist_skills", "aget_skill", @@ -1001,6 +1006,11 @@ class ProxyBaseLLMRequestProcessing: "aingest", "aretrieve_container", "adelete_container", + "aupload_container_file", + "alist_container_files", + "aretrieve_container_file", + "adelete_container_file", + "aretrieve_container_file_content", "acreate_skill", "alist_skills", "aget_skill", @@ -1594,6 +1604,12 @@ class ProxyBaseLLMRequestProcessing: # here would duplicate the guardrail API call # (e.g. double OpenAI Moderation charges). continue + if "async_post_call_streaming_iterator_hook" in type(cb).__dict__: + # Skip — the guardrail already scanned the assembled + # response via its own streaming iterator hook in the + # streaming pipeline. re running this function async_post_call_success_hook + # here would duplicate the scan and can spuriously block the guardrail that already passed / failed. + continue else: guardrail_result = await cb.async_post_call_success_hook( user_api_key_dict=captured_user_api_key_dict, diff --git a/litellm/proxy/common_utils/cache_coordinator.py b/litellm/proxy/common_utils/cache_coordinator.py index 24da9450ab8..abb0402d3b9 100644 --- a/litellm/proxy/common_utils/cache_coordinator.py +++ b/litellm/proxy/common_utils/cache_coordinator.py @@ -20,11 +20,27 @@ T = TypeVar("T") class AsyncCacheProtocol(Protocol): - """Protocol for cache backends used by EventDrivenCacheCoordinator.""" + """Protocol for cache backends used by EventDrivenCacheCoordinator. - async def async_get_cache(self, key: str, **kwargs: Any) -> Any: ... + Matches ``DualCache`` / ``UserApiKeyCache`` call shapes (explicit optional params + before ``**kwargs``), not only ``(key, **kwargs)``, so overloads validate. + """ - async def async_set_cache(self, key: str, value: Any, **kwargs: Any) -> Any: ... + async def async_get_cache( + self, + key: str, + parent_otel_span: Any = None, + local_only: bool = False, + **kwargs: Any, + ) -> Any: ... + + async def async_set_cache( + self, + key: str, + value: Any, + local_only: bool = False, + **kwargs: Any, + ) -> Any: ... class EventDrivenCacheCoordinator: @@ -36,6 +52,9 @@ class EventDrivenCacheCoordinator: - Other requests: wait for the signal, then read from cache. Create one instance per resource (e.g. one for global spend, one for feature flags). + + Args: + log_prefix: Prefix for debug log messages. """ def __init__(self, log_prefix: str = "[CACHE]"): diff --git a/litellm/proxy/common_utils/cache_pydantic_utils.py b/litellm/proxy/common_utils/cache_pydantic_utils.py new file mode 100644 index 00000000000..80a8d6281a1 --- /dev/null +++ b/litellm/proxy/common_utils/cache_pydantic_utils.py @@ -0,0 +1,93 @@ +""" +DualCache presents a single API for reads and writes, but the two backends behave +differently: the in-memory layer can store arbitrary Python objects (including live +``BaseModel`` instances), while Redis persists strings and therefore needs JSON-safe +payloads (``json.dumps`` on the Redis side). + +Call sites therefore see cache ``value`` / ``cached`` as effectively ``Any``: the same +key may deserialize to a model on one process (memory hit) or to a ``dict`` after a +Redis round-trip. ``CacheCodec`` centralizes encode/decode at that boundary: +``CacheCodec.serialize`` before ``set``, ``CacheCodec.deserialize`` after ``get`` +when you need a typed ``BaseModel``. + +``dataclasses`` are not supported: only ``dict`` and Pydantic ``BaseModel`` inputs +are encoded; pass a Pydantic model or convert with e.g. ``dataclasses.asdict`` first. +""" + +from __future__ import annotations + +from typing import Any, Optional, Type, TypeVar + +from pydantic import BaseModel, ValidationError + +from litellm._logging import verbose_proxy_logger + +T = TypeVar("T", bound=BaseModel) + + +class CacheCodec: + """ + Encode/decode Pydantic models for DualCache (memory vs Redis safe payloads). + + Dataclasses are not supported yet (only ``dict`` and ``BaseModel``). + + Use ``serialize`` with ``model_type`` when writing so the same schema is used + as on read (``deserialize``). Pass ``model_type`` whenever you know it + (validates ``dict`` payloads and normalizes ``BaseModel`` instances). + """ + + @staticmethod + def serialize(value: Any, model_type: Optional[Type[T]] = None) -> Any: + """ + Encode a value for DualCache / Redis (``json.dumps``-safe). + + If ``model_type`` is set, the payload is validated with that model, then + ``model_dump(mode="json", exclude_none=True)`` — symmetric with ``deserialize``. + + If the value is already an instance of ``model_type`` (or a subclass), + ``model_validate`` is skipped to avoid an unnecessary Pydantic copy — the + value is dumped directly. + + If ``model_type`` is omitted, any ``BaseModel`` is dumped as above; other + values (e.g. plain ``dict``) are returned unchanged. + """ + if model_type is not None: + if isinstance(value, model_type): + # Already the right type: dump directly, skip re-validation copy. + return value.model_dump(mode="json", exclude_none=True) + if isinstance(value, (dict, BaseModel)): + return model_type.model_validate(value).model_dump( + mode="json", exclude_none=True + ) + return value + if isinstance(value, BaseModel): + return value.model_dump(mode="json", exclude_none=True) + return value + + @staticmethod + def deserialize(cached: Any, model_type: Type[T]) -> Optional[T]: + """ + Decode a cache entry to ``model_type``. + + - ``None`` → ``None`` + - Already an instance of ``model_type`` (including subclasses) → returned as-is + - ``dict`` → ``model_type.model_validate(...)``; on ``ValidationError``, + logs a warning and returns ``None`` (treat as cache miss; avoids serving + malformed or schema-drifted entries) + - Any other type → ``None`` (caller should treat as cache miss or log) + """ + if cached is None: + return None + if isinstance(cached, model_type): + return cached + if isinstance(cached, dict): + try: + return model_type.model_validate(cached) + except ValidationError as e: + verbose_proxy_logger.warning( + "CacheCodec.deserialize: validation failed for %s (%s)", + model_type.__name__, + e, + ) + return None + return None diff --git a/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py b/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py index c25d8533128..67a24567461 100644 --- a/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py +++ b/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py @@ -8,7 +8,7 @@ from datetime import datetime, timezone from typing import Any, Dict, List, Optional from litellm._logging import verbose_proxy_logger -from litellm.caching import DualCache +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.constants import ( EXPIRED_UI_SESSION_KEY_CLEANUP_JOB_NAME, LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_BATCH_SIZE, @@ -31,7 +31,7 @@ class ExpiredUISessionKeyCleanupManager: def __init__( self, prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, pod_lock_manager=None, ): self.prisma_client = prisma_client diff --git a/litellm/proxy/common_utils/user_api_key_cache.py b/litellm/proxy/common_utils/user_api_key_cache.py new file mode 100644 index 00000000000..914be364579 --- /dev/null +++ b/litellm/proxy/common_utils/user_api_key_cache.py @@ -0,0 +1,162 @@ +from __future__ import annotations + +from typing import Any, Optional, Type, TypeVar, Union, cast, overload + +from pydantic import BaseModel + +from litellm._logging import verbose_proxy_logger +from litellm.caching.dual_cache import DualCache +from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec + +T = TypeVar("T", bound=BaseModel) + + +class UserApiKeyCache(DualCache): + """ + DualCache wrapper for UserAPIKeyAuth-like payloads. + + Stores a Redis-safe JSON payload in BOTH in-memory and Redis to avoid + "memory returns BaseModel, Redis returns dict" format drift. + + When ``model_type`` is provided: + - writes are serialized via ``CacheCodec.serialize(..., model_type=...)`` + - reads are deserialized via ``CacheCodec.deserialize(..., model_type)`` + and return ``Optional[T]``: the model on success, ``None`` on cache miss + **or** if the cached payload fails validation (schema drift). On + validation failure after a cache hit, an error line is emitted via + ``verbose_proxy_logger``. + + When ``model_type`` is omitted, the interface behaves like ``DualCache``: + raw cached payload is returned (dict/str/etc.). + + ``async_set_cache_pipeline`` applies the same untyped Codec pass as omitting + ``model_type`` on ``async_set_cache`` (so ``BaseModel`` rows are dumped before Redis). + + ``get_cache`` / ``async_get_cache`` overloads and implementations must be contiguous + (no other methods in between) so mypy resolves ``@overload`` + implementation correctly. + """ + + @overload + def get_cache( + self, + key: Any, + parent_otel_span: Any = None, + local_only: bool = False, + *, + model_type: Type[T], + **kwargs: Any, + ) -> Optional[T]: ... + + @overload + def get_cache( + self, + key: Any, + parent_otel_span: Any = None, + local_only: bool = False, + **kwargs: Any, + ) -> Any: ... + + def get_cache( # type: ignore[override] + self, + key, + parent_otel_span=None, + local_only: bool = False, + model_type: Optional[Type[BaseModel]] = None, + **kwargs, + ) -> Union[Any, Optional[BaseModel]]: + if model_type is None and "model_type" in kwargs: + model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) + cached = super().get_cache( + key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs + ) + if model_type is None: + return cached + if cached is None: + return None + decoded = CacheCodec.deserialize(cached, model_type=model_type) + if decoded is None: + verbose_proxy_logger.error( + "UserApiKeyCache.get_cache failed to deserialize cached value for " + "key=%r model_type=%s", + key, + getattr(model_type, "__name__", str(model_type)), + ) + return None + return decoded + + @overload + async def async_get_cache( + self, + key: Any, + parent_otel_span: Any = None, + local_only: bool = False, + *, + model_type: Type[T], + **kwargs: Any, + ) -> Optional[T]: ... + + @overload + async def async_get_cache( + self, + key: Any, + parent_otel_span: Any = None, + local_only: bool = False, + **kwargs: Any, + ) -> Any: ... + + async def async_get_cache( # type: ignore[override] + self, + key, + parent_otel_span=None, + local_only: bool = False, + model_type: Optional[Type[BaseModel]] = None, + **kwargs, + ) -> Union[Any, Optional[BaseModel]]: + if model_type is None and "model_type" in kwargs: + model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) + cached = await super().async_get_cache( + key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs + ) + if model_type is None: + return cached + if cached is None: + return None + decoded = CacheCodec.deserialize(cached, model_type=model_type) + if decoded is None: + verbose_proxy_logger.error( + "UserApiKeyCache.async_get_cache failed to deserialize cached value for " + "key=%r model_type=%s", + key, + getattr(model_type, "__name__", str(model_type)), + ) + return None + return decoded + + def set_cache(self, key, value, local_only: bool = False, **kwargs): # type: ignore[override] + model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) + payload = CacheCodec.serialize(value, model_type=model_type) + return super().set_cache( + key=key, value=payload, local_only=local_only, **kwargs + ) + + async def async_set_cache(self, key, value, local_only: bool = False, **kwargs): # type: ignore[override] + model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) + payload = CacheCodec.serialize(value, model_type=model_type) + return await super().async_set_cache( + key=key, value=payload, local_only=local_only, **kwargs + ) + + async def async_set_cache_pipeline( # type: ignore[override] + self, cache_list: list, local_only: bool = False, **kwargs + ) -> None: + """ + Batch writes with the same Codec boundary as ``async_set_cache`` without + ``model_type``: ``BaseModel`` values become JSON-safe dicts; dicts/scalars unchanged. + """ + normalized = [ + (key, CacheCodec.serialize(value, model_type=None)) + for key, value in cache_list + ] + return await super().async_set_cache_pipeline( + cache_list=normalized, local_only=local_only, **kwargs + ) diff --git a/litellm/proxy/container_endpoints/handler_factory.py b/litellm/proxy/container_endpoints/handler_factory.py index fae7f939aed..794051e90f8 100644 --- a/litellm/proxy/container_endpoints/handler_factory.py +++ b/litellm/proxy/container_endpoints/handler_factory.py @@ -19,7 +19,6 @@ from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_headers, get_custom_llm_provider_from_request_query, ) -from litellm.responses.utils import ResponsesAPIRequestUtils def _load_endpoints_config() -> Dict: @@ -64,10 +63,12 @@ def _create_handler_for_path_params( request: Request, container_id: str, file_id: str, + fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): return await _process_binary_request( request=request, + fastapi_response=fastapi_response, container_id=container_id, file_id=file_id, user_api_key_dict=user_api_key_dict, @@ -152,63 +153,61 @@ def _create_handler_for_path_params( async def _process_binary_request( request: Request, + fastapi_response: Response, container_id: str, file_id: str, user_api_key_dict: UserAPIKeyAuth, ): """ - Process binary content requests using the proper transformation pattern. + Process binary content requests through the standard proxy/router pipeline. - This uses the provider config transformations and llm_http_handler - to maintain consistency with the established pattern. + The router owns managed container ID decoding and deployment selection. This + handler only adapts the byte response to FastAPI. """ - from litellm.litellm_core_utils.litellm_logging import Logging - from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler - from litellm.types.router import GenericLiteLLMParams + from litellm.proxy.proxy_server import ( + general_settings, + llm_router, + proxy_config, + proxy_logging_obj, + select_data_generator, + user_api_base, + user_max_tokens, + user_model, + user_request_timeout, + user_temperature, + version, + ) - # Extract custom_llm_provider custom_llm_provider = ( get_custom_llm_provider_from_request_headers(request=request) or get_custom_llm_provider_from_request_query(request=request) or "openai" ) - - # Build litellm_params - credentials are resolved by provider config from env - litellm_params = GenericLiteLLMParams() - - # Decode container ID and extract provider info - decoded = ResponsesAPIRequestUtils._decode_container_id(container_id) - original_container_id = decoded.get("response_id", container_id) - - # If container ID has encoded provider info and user didn't explicitly set provider, use it - decoded_provider = decoded.get("custom_llm_provider") - if decoded_provider and custom_llm_provider == "openai": - custom_llm_provider = decoded_provider - - # Get the provider config - container_provider_config = _get_container_provider_config(custom_llm_provider) - - # Create logging object - logging_obj = Logging( - model="container-file-content", - messages=[], - stream=False, - call_type="container_file_content", - start_time=None, - litellm_call_id="", - function_id="", - ) - - # Use the HTTP handler to make the request - handler = BaseLLMHTTPHandler() + data: Dict[str, Any] = { + "container_id": container_id, + "file_id": file_id, + "custom_llm_provider": custom_llm_provider, + } + processor = ProxyBaseLLMRequestProcessing(data=data) try: - content = await handler.async_container_file_content_handler( - container_id=original_container_id, # Use decoded original ID - file_id=file_id, - container_provider_config=container_provider_config, - litellm_params=litellm_params, - logging_obj=logging_obj, + content = await processor.base_process_llm_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + route_type="aretrieve_container_file_content", + proxy_logging_obj=proxy_logging_obj, + llm_router=llm_router, + general_settings=general_settings, + proxy_config=proxy_config, + select_data_generator=select_data_generator, + model=None, + user_model=user_model, + user_temperature=user_temperature, + user_request_timeout=user_request_timeout, + user_max_tokens=user_max_tokens, + user_api_base=user_api_base, + version=version, ) # Determine content type based on common file extensions in the file_id @@ -229,13 +228,25 @@ async def _process_binary_request( elif ".pdf" in file_id_lower: content_type = "application/pdf" + if not isinstance(content, bytes): + raise TypeError( + "aretrieve_container_file_content expected bytes, got " + f"{type(content).__name__}" + ) + return Response( content=content, + headers=dict(fastapi_response.headers), media_type=content_type, ) except Exception as e: - raise e + raise await processor._handle_llm_api_exception( + e=e, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=proxy_logging_obj, + version=version, + ) async def _process_multipart_upload_request( @@ -284,16 +295,7 @@ async def _process_multipart_upload_request( or "openai" ) - # Decode container ID and extract provider info - decoded = ResponsesAPIRequestUtils._decode_container_id(container_id) - original_container_id = decoded.get("response_id", container_id) - - # If container ID has encoded provider info and user didn't explicitly set provider, use it - decoded_provider = decoded.get("custom_llm_provider") - if decoded_provider and custom_llm_provider == "openai": - custom_llm_provider = decoded_provider - - data["container_id"] = original_container_id # Use decoded original ID + data["container_id"] = container_id data["custom_llm_provider"] = custom_llm_provider processor = ProxyBaseLLMRequestProcessing(data=data) @@ -359,21 +361,6 @@ async def _process_request( or "openai" ) - # Decode container_id if present in path_params - if "container_id" in path_params: - decoded = ResponsesAPIRequestUtils._decode_container_id( - path_params["container_id"] - ) - original_container_id = decoded.get("response_id", path_params["container_id"]) - - # If container ID has encoded provider info and user didn't explicitly set provider, use it - decoded_provider = decoded.get("custom_llm_provider") - if decoded_provider and custom_llm_provider == "openai": - custom_llm_provider = decoded_provider - - # Update path_params with decoded original ID - data["container_id"] = original_container_id - data["custom_llm_provider"] = custom_llm_provider processor = ProxyBaseLLMRequestProcessing(data=data) diff --git a/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py b/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py index 6dd0288cb09..37be832d350 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py +++ b/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py @@ -225,10 +225,10 @@ class ToolPermissionGuardrail(CustomGuardrail): def _parse_tool_call_arguments( self, tool_call: ChatCompletionMessageToolCall - ) -> Dict[str, Any]: + ) -> tuple[Optional[Dict[str, Any]], Optional[str]]: arguments = getattr(tool_call.function, "arguments", None) if not arguments: - return {} + return None, "missing arguments" parsed_arguments: Any = {} try: @@ -236,22 +236,24 @@ class ToolPermissionGuardrail(CustomGuardrail): parsed_arguments = json.loads(arguments) elif isinstance(arguments, dict): parsed_arguments = arguments - except json.JSONDecodeError as exc: + else: + return None, "arguments must be a JSON object" + except (json.JSONDecodeError, TypeError) as exc: verbose_proxy_logger.warning( "Tool Permission Guardrail: Failed to decode arguments for tool %s: %s", tool_call.function.name, exc, ) - return {} + return None, "arguments could not be parsed" if isinstance(parsed_arguments, dict): - return parsed_arguments + return parsed_arguments, None verbose_proxy_logger.debug( - "Tool Permission Guardrail: Ignoring non-dict arguments for tool %s", + "Tool Permission Guardrail: Rejecting non-dict arguments for tool %s", tool_call.function.name, ) - return {} + return None, "arguments must be a JSON object" def _collect_argument_paths( self, @@ -331,10 +333,21 @@ class ToolPermissionGuardrail(CustomGuardrail): continue if rule.allowed_param_patterns and should_check_params: - arguments = self._parse_tool_call_arguments(tool_call) + arguments, parse_error = self._parse_tool_call_arguments(tool_call) + if parse_error: + default_message = f"Tool '{tool_identifier}' {parse_error} required by rule '{rule.id}'" + message = self.render_violation_message( + default=default_message, + context={"tool_name": tool_identifier, "rule_id": rule.id}, + ) + return False, rule.id, message if not arguments: - last_pattern_failure_msg = f"Tool '{tool_identifier}' is missing arguments required by rule '{rule.id}'" - continue + default_message = f"Tool '{tool_identifier}' is missing arguments required by rule '{rule.id}'" + message = self.render_violation_message( + default=default_message, + context={"tool_name": tool_identifier, "rule_id": rule.id}, + ) + return False, rule.id, message patterns_match, failure_message = self._patterns_match_for_rule( arguments=arguments, @@ -365,6 +378,33 @@ class ToolPermissionGuardrail(CustomGuardrail): ) return is_allowed, None, message + @staticmethod + def _get_mapping_value(item: Any, key: str) -> Any: + if isinstance(item, dict): + return item.get(key) + return getattr(item, key, None) + + @staticmethod + def _legacy_function_call_id(choice_index: int) -> str: + return f"legacy_function_call_{choice_index}" + + def _legacy_function_call_to_tool_call( + self, function_call: Any, choice_index: int + ) -> Optional[ChatCompletionMessageToolCall]: + if function_call is None: + return None + + function_name = self._get_mapping_value(function_call, "name") + arguments = self._get_mapping_value(function_call, "arguments") or "" + if not function_name: + return None + + return ChatCompletionMessageToolCall( + id=self._legacy_function_call_id(choice_index), + type="function", + function={"name": function_name, "arguments": arguments}, + ) + def _extract_tool_calls_from_response( self, response: ModelResponse ) -> List[ChatCompletionMessageToolCall]: @@ -379,13 +419,72 @@ class ToolPermissionGuardrail(CustomGuardrail): """ tool_calls = [] - for choice in response.choices: + for choice_index, choice in enumerate(response.choices): if isinstance(choice, Choices): for tool in choice.message.tool_calls or []: tool_calls.append(tool) + legacy_tool_call = self._legacy_function_call_to_tool_call( + getattr(choice.message, "function_call", None), choice_index + ) + if legacy_tool_call is not None: + tool_calls.append(legacy_tool_call) return tool_calls + def _get_request_tool_name(self, tool: Any) -> tuple[Optional[str], Optional[str]]: + tool_type = self._get_mapping_value(tool, "type") + if tool_type != "function": + return None, tool_type + + function = self._get_mapping_value(tool, "function") + tool_name = self._get_mapping_value(function, "name") + return tool_name, tool_type + + def _get_legacy_function_name(self, function: Any) -> Optional[str]: + return self._get_mapping_value(function, "name") + + def _get_named_tool_choice(self, data: dict) -> Optional[str]: + tool_choice = data.get("tool_choice") + if not tool_choice or tool_choice in ("auto", "none", "required"): + return None + if isinstance(tool_choice, str): + return tool_choice + if self._get_mapping_value(tool_choice, "type") != "function": + return None + return self._get_mapping_value( + self._get_mapping_value(tool_choice, "function"), "name" + ) + + def _get_named_function_call(self, data: dict) -> Optional[str]: + function_call = data.get("function_call") + if not function_call or function_call in ("auto", "none"): + return None + if isinstance(function_call, str): + return function_call + return self._get_mapping_value(function_call, "name") + + def _collect_request_tools(self, data: dict) -> List[tuple[str, Optional[str]]]: + request_tools: List[tuple[str, Optional[str]]] = [] + + for tool in data.get("tools") or []: + tool_name, tool_type = self._get_request_tool_name(tool) + if tool_name is not None: + request_tools.append((tool_name, tool_type)) + + for function in data.get("functions") or []: + function_name = self._get_legacy_function_name(function) + if function_name is not None: + request_tools.append((function_name, "function")) + + for forced_tool_name in ( + self._get_named_tool_choice(data), + self._get_named_function_call(data), + ): + if forced_tool_name is not None: + request_tools.append((forced_tool_name, "function")) + + return request_tools + def _modify_request_with_permission_errors( self, data: dict, @@ -410,19 +509,32 @@ class ToolPermissionGuardrail(CustomGuardrail): for tool_use in denied_tool_names: error_tool_names.add(tool_use) - # Modify the tools tools: Optional[List[ChatCompletionToolParam]] = data.get("tools") - if tools is None: - return data - - new_tools = [] - for tool in tools: - if tool["type"] != "function": - continue - tool_name: str = tool["function"]["name"] - if tool_name not in error_tool_names: + if tools is not None: + new_tools = [] + for tool in tools: + tool_name, tool_type = self._get_request_tool_name(tool) + if tool_type == "function" and tool_name in error_tool_names: + continue new_tools.append(tool) - data["tools"] = new_tools + data["tools"] = new_tools + + functions = data.get("functions") + if functions is not None: + data["functions"] = [ + function + for function in functions + if self._get_legacy_function_name(function) not in error_tool_names + ] + + named_tool_choice = self._get_named_tool_choice(data) + if named_tool_choice in error_tool_names: + data["tool_choice"] = "none" + + named_function_call = self._get_named_function_call(data) + if named_function_call in error_tool_names: + data["function_call"] = "none" + return data def _create_permission_error_result( @@ -472,7 +584,7 @@ class ToolPermissionGuardrail(CustomGuardrail): error_results[tool_use.id] = error_result # Modify the response content - for choice in response.choices: + for choice_index, choice in enumerate(response.choices): if isinstance(choice, Choices): filtered_tool_calls = [] error_messages = [] @@ -490,6 +602,15 @@ class ToolPermissionGuardrail(CustomGuardrail): filtered_tool_calls if filtered_tool_calls else None ) + legacy_tool_call = self._legacy_function_call_to_tool_call( + getattr(choice.message, "function_call", None), choice_index + ) + if legacy_tool_call is not None: + legacy_error_result = error_results.get(legacy_tool_call.id) + if legacy_error_result is not None: + choice.message.function_call = None + error_messages.append(legacy_error_result.content) + # Add error messages to content if error_messages: existing_content = choice.message.content @@ -519,21 +640,16 @@ class ToolPermissionGuardrail(CustomGuardrail): if self.should_run_guardrail(data=data, event_type=event_type) is not True: return data - new_tools: Optional[List[ChatCompletionToolParam]] = data.get("tools") - if new_tools is None: + new_tools = self._collect_request_tools(data) + if not new_tools: verbose_proxy_logger.warning( - "Tool Permission Guardrail: not running guardrail. No tools in data" + "Tool Permission Guardrail: not running guardrail. No tools or functions in data" ) return data # Check permissions for each tool denied_tool_names = [] - for tool in new_tools: - if tool["type"] != "function": - continue - tool_name: str = tool["function"]["name"] - tool_type: Optional[str] = tool.get("type") - + for tool_name, tool_type in new_tools: is_allowed, _, message = self._check_tool_permission(tool_name, tool_type) if not is_allowed and message is not None: diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index 7d67750c78f..7c340ff5df6 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -29,6 +29,10 @@ ILLEGAL_DISPLAY_PARAMS = [ "exception", # internal; not JSON-serializable, never for display "litellm_metadata", # internal tracking metadata with auth objects; not for display ] +# Provider routing fields. Allowed for proxy admins so they can see which +# region/version a deployment is checking; gated at the endpoint layer for +# non-admin callers (see _strip_admin_only_fields_from_health_result). +ADMIN_ONLY_HEALTH_DISPLAY_PARAMS = ("api_base", "api_version") MINIMAL_DISPLAY_PARAMS = ["model", "mode_error"] diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index b4b5de1746e..35c9edb937d 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -20,6 +20,7 @@ from litellm.proxy._types import ( CallInfo, EnterpriseLicenseData, Litellm_EntityType, + LitellmUserRoles, ProxyErrorTypes, ProxyException, UserAPIKeyAuth, @@ -28,6 +29,7 @@ from litellm.proxy._types import ( from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.health_check import ( + ADMIN_ONLY_HEALTH_DISPLAY_PARAMS, _clean_endpoint_data, _update_litellm_params_for_health_check, perform_health_check, @@ -723,6 +725,129 @@ async def _save_background_health_checks_to_db( # Continue execution - don't let database save failure break health checks +_PROXY_ADMIN_ROLES = frozenset( + { + LitellmUserRoles.PROXY_ADMIN.value, + # View-only admins are operators (oncall, support); they need the + # routing fields (api_base, api_version) to diagnose health and tell + # which provider region a check is hitting. They cannot mutate config + # so granting them the read-only view is safe. + LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + } +) + + +def _is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool: + """ + Return True if the caller has a proxy-admin role (full or view-only). + + user_role on UserAPIKeyAuth can be either a LitellmUserRoles enum or its + string value depending on how the auth path constructed the object, so we + compare against the raw value rather than the enum identity. + """ + role = user_api_key_dict.user_role + if role is None: + return False + role_value = role.value if hasattr(role, "value") else role + return role_value in _PROXY_ADMIN_ROLES + + +def _strip_admin_only_fields_from_health_result(result: dict) -> dict: + """ + Return a copy of the /health response with provider routing fields + (``api_base``, ``api_version``) removed from each healthy/unhealthy + endpoint entry. Used to hide those fields from non-admin callers while + still showing them which deployments they own and whether each one is + healthy. Proxy admins receive the unmodified result. + """ + out = dict(result) + drop = set(ADMIN_ONLY_HEALTH_DISPLAY_PARAMS) + for key in ("healthy_endpoints", "unhealthy_endpoints"): + eps = out.get(key) + if isinstance(eps, list): + out[key] = [ + ( + {k: v for k, v in ep.items() if k not in drop} + if isinstance(ep, dict) + else ep + ) + for ep in eps + ] + return out + + +def _resolve_targeted_model_ids( + model_list: list, model: Optional[str], model_id: Optional[str] +) -> Optional[set]: + """ + Resolve a ``/health`` ``model`` / ``model_id`` query param to the set of + deployment IDs the response should be scoped to. + + Mirrors the live-path semantics in ``perform_health_check()``: ``model`` + matches either the deployment's ``model_name`` alias or its + ``litellm_params.model`` provider string. ``model_id`` matches + ``model_info.id``. + + Both query params are validated against the supplied ``model_list``. + Callers pass an already-scoped list (filtered to the caller's allowed + models for non-admins, full list for admins), so a ``model_id`` that + isn't present resolves to an empty set rather than a single-element + set — preventing a non-admin from reading another deployment's cached + health entry by guessing its ID. + + Returns ``None`` when no targeting is requested — callers should treat + that as "no filter." + """ + if not model and not model_id: + return None + target_ids: set = set() + for m in model_list: + deployment_id = (m.get("model_info") or {}).get("id") + if not deployment_id: + continue + if model_id and deployment_id == model_id: + target_ids.add(deployment_id) + continue + if model: + litellm_model = (m.get("litellm_params") or {}).get("model") + if m.get("model_name") == model or litellm_model == model: + target_ids.add(deployment_id) + return target_ids + + +def _filter_health_check_results_by_model_ids( + results: dict, allowed_model_ids: set +) -> dict: + """ + Restrict a cached background health-check result dict to endpoints whose + model_id is in ``allowed_model_ids``. + + Endpoints without a model_id (e.g. CLI-model entries that predate the + model_id wiring) are dropped conservatively — we cannot prove they belong + to the caller, so they are excluded rather than leaked. + + Each retained endpoint is shallow-copied before being returned, so any + downstream transform (e.g. _strip_admin_only_fields_from_health_result) + cannot accidentally mutate the shared ``health_check_results`` cache. + """ + healthy = [ + dict(ep) + for ep in (results.get("healthy_endpoints") or []) + if ep.get("model_id") in allowed_model_ids + ] + unhealthy = [ + dict(ep) + for ep in (results.get("unhealthy_endpoints") or []) + if ep.get("model_id") in allowed_model_ids + ] + return { + "healthy_endpoints": healthy, + "unhealthy_endpoints": unhealthy, + "healthy_count": len(healthy), + "unhealthy_count": len(unhealthy), + } + + async def _perform_health_check_and_save( model_list, target_model, @@ -771,6 +896,7 @@ async def _perform_health_check_and_save( @router.get("/health", tags=["health"], dependencies=[Depends(user_api_key_auth)]) async def health_endpoint( + response: Response, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), model: Optional[str] = fastapi.Query( None, description="Specify the model name (optional)" @@ -838,11 +964,33 @@ async def health_endpoint( detail={"error": f"Model with ID {model_id} not found"}, ) + is_admin = _is_proxy_admin(user_api_key_dict) + model_specific_request = bool(model or model_id) + + def _post_process(result: dict) -> dict: + # api_base / api_version reveal which provider/region/internal host the + # deployment talks to; only proxy admins receive them. Non-admin keys + # still see model/model_id and the healthy/unhealthy status. We also + # set a header so non-admin clients that previously parsed those + # fields can detect the change programmatically. + # When a caller asked about a specific model/model_id and zero + # endpoints came back healthy, surface that as a 503 so monitoring + # systems can rely on the HTTP status instead of having to parse the + # body. The body shape is unchanged. + if model_specific_request and result.get("healthy_count", 0) == 0: + response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE + if is_admin: + return result + response.headers["Litellm-Health-Field-Notice"] = ( + "api_base and api_version are admin-only on this endpoint" + ) + return _strip_admin_only_fields_from_health_result(result) + try: if llm_model_list is None: # if no router set, check if user set a model using litellm --model ollama/llama2 if user_model is not None: - return await _perform_health_check_and_save( + cli_result = await _perform_health_check_and_save( model_list=[], target_model=None, cli_model=user_model, @@ -853,20 +1001,81 @@ async def health_endpoint( model_id=None, # CLI model doesn't have model_id max_concurrency=health_check_concurrency, ) + return _post_process(cli_result) raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail={"error": "Model list not initialized"}, ) _llm_model_list = copy.deepcopy(llm_model_list) ### FILTER MODELS FOR ONLY THOSE USER HAS ACCESS TO ### + # Live path: scope by model_name (every deployment has one). + # Cache path: scope by model_id (the cache is keyed on model_id). + # Consequence: a deployment whose model_name the caller can access + # but which lacks model_info.id will appear in the live /health + # response but NOT in the background-cache /health response. This is + # surfaced via the "warnings" field below so operators can fix the + # missing model_info.id rather than guess at the discrepancy. if len(user_api_key_dict.models) > 0: - pass - else: - pass # + allowed_models = set(user_api_key_dict.models) + _llm_model_list = [ + m for m in _llm_model_list if m.get("model_name") in allowed_models + ] if use_background_health_checks: - return health_check_results + # The cached background result covers every model. When the + # caller targets a specific model/model_id we have to narrow the + # cache to that deployment before _post_process evaluates + # healthy_count, otherwise an unhealthy "foo" combined with any + # other healthy model would still report healthy_count > 0 and + # the targeted-503 path would never fire. + targeted_ids = _resolve_targeted_model_ids(_llm_model_list, model, model_id) + if len(user_api_key_dict.models) > 0: + allowed_model_ids = { + (m.get("model_info") or {}).get("id") + for m in _llm_model_list + if (m.get("model_info") or {}).get("id") + } + # _llm_model_list is already scoped to the caller's allowed + # model_names above, so targeted_ids is implicitly the + # intersection of "targeted" and "allowed." + filter_ids = ( + targeted_ids if targeted_ids is not None else allowed_model_ids + ) + filtered = _filter_health_check_results_by_model_ids( + health_check_results, filter_ids + ) + if targeted_ids is None and not allowed_model_ids: + # Caller has accessible model_names but none of the + # matching deployments expose a model_info.id, so the + # cache filter (which keys on model_id) drops every + # entry. Surface this both as a warning log and a + # structured "warnings" field on the response so the + # caller can distinguish "no deployments found" from + # "deployments excluded due to missing model_info.id". + verbose_proxy_logger.warning( + "health_endpoint: scoped key %s has accessible models %s " + "but none of the matching deployments carry a model_info.id; " + "background health-check cache will return an empty result.", + user_api_key_dict.user_id, + list(user_api_key_dict.models), + ) + filtered["warnings"] = [ + "Some accessible deployments are missing model_info.id " + "and were excluded from this response. Ask a proxy admin " + "to populate model_info.id for these models." + ] + return _post_process(filtered) + if targeted_ids is not None: + # Admin caller targeting a specific model: filter the cache + # so the response (and the targeted-503 check) reflects only + # that deployment, not the global aggregate. + return _post_process( + _filter_health_check_results_by_model_ids( + health_check_results, targeted_ids + ) + ) + return _post_process(health_check_results) else: - return await _perform_health_check_and_save( + router_result = await _perform_health_check_and_save( model_list=_llm_model_list, target_model=target_model, cli_model=None, @@ -877,6 +1086,7 @@ async def health_endpoint( model_id=model_id, max_concurrency=health_check_concurrency, ) + return _post_process(router_result) except Exception as e: verbose_proxy_logger.error( "litellm.proxy.proxy_server.py::health_endpoint(): Exception occured - {}".format( @@ -1242,7 +1452,7 @@ def callback_name(callback): tags=["health"], dependencies=[Depends(user_api_key_auth)], ) -async def health_readiness(): +async def health_readiness(response: Response): """ Unprotected endpoint for checking if worker can receive requests """ @@ -1275,8 +1485,8 @@ async def health_readiness(): try: index_info = await litellm.cache.cache._index_info() except Exception as e: - index_info = "index does not exist - error: " + str(e) - cache_type = {"type": cache_type, "index_info": index_info} + index_info = "index does not exist - error: " + str(e) # type: ignore[assignment] + cache_type = {"type": cache_type, "index_info": index_info} # type: ignore[assignment] # check log level log_level_name = logging.getLevelName(verbose_logger.getEffectiveLevel()) @@ -1285,6 +1495,12 @@ async def health_readiness(): # check DB if prisma_client is not None: # if db passed in, check if it's connected db_health_status = await _db_health_readiness_check() + # A configured DB that is not reachable means the worker cannot + # serve requests that depend on persisted state (keys, budgets, + # spend logs). Return 503 so orchestrators take this pod out of + # rotation; "Not connected" (no DB configured at all) stays 200. + if db_health_status["status"] != "connected": + response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE return { "status": "healthy", "db": db_health_status["status"], diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 06b7d857896..2ee0588f19e 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -164,13 +164,13 @@ class _PROXY_BatchRateLimiter(CustomLogger): batch_usage: BatchFileUsage, ) -> None: """ - Check rate limits and increment counters by the batch amounts. + Atomically check + increment rate-limit counters by the batch amounts. - Raises HTTPException if any limit would be exceeded. + Raises HTTPException if any descriptor would exceed its limit; in that + case no counter is modified. Backed by `atomic_check_and_increment_by_n` + which uses a Redis Lua script when available (multi-process atomic) and + falls back to a per-process asyncio.Lock + in-memory operation. """ - from litellm.types.caching import RedisPipelineIncrementOperation - - # Create descriptors and check if batch would exceed limits descriptors = self.parallel_request_limiter._create_rate_limit_descriptors( user_api_key_dict=user_api_key_dict, data=data, @@ -179,73 +179,31 @@ class _PROXY_BatchRateLimiter(CustomLogger): model_has_failures=False, ) - # Check current usage without incrementing - rate_limit_response = await self.parallel_request_limiter.should_rate_limit( - descriptors=descriptors, - parent_otel_span=user_api_key_dict.parent_otel_span, - read_only=True, - ) + increment: Dict[Literal["requests", "tokens"], int] = { + "requests": batch_usage.request_count, + "tokens": batch_usage.total_tokens, + } + increments: List[Dict[Literal["requests", "tokens"], int]] = [ + increment for _ in descriptors + ] - # Verify batch won't exceed any limits - for status in rate_limit_response["statuses"]: - rate_limit_type = status["rate_limit_type"] - limit_remaining = status["limit_remaining"] - - required_capacity = ( - batch_usage.request_count - if rate_limit_type == "requests" - else batch_usage.total_tokens if rate_limit_type == "tokens" else 0 - ) - - if required_capacity > limit_remaining: - self._raise_rate_limit_error( - status, descriptors, batch_usage, rate_limit_type - ) - - # Build pipeline operations for batch increments - # Reuse the same keys that descriptors check - pipeline_operations: List[RedisPipelineIncrementOperation] = [] - - for descriptor in descriptors: - key = descriptor["key"] - value = descriptor["value"] - rate_limit = descriptor.get("rate_limit") - - if rate_limit is None: - continue - - # Add RPM increment if limit is set - if rate_limit.get("requests_per_unit") is not None: - rpm_key = self.parallel_request_limiter.create_rate_limit_keys( - key=key, value=value, rate_limit_type="requests" - ) - pipeline_operations.append( - RedisPipelineIncrementOperation( - key=rpm_key, - increment_value=batch_usage.request_count, - ttl=self.parallel_request_limiter.window_size, - ) - ) - - # Add TPM increment if limit is set - if rate_limit.get("tokens_per_unit") is not None: - tpm_key = self.parallel_request_limiter.create_rate_limit_keys( - key=key, value=value, rate_limit_type="tokens" - ) - pipeline_operations.append( - RedisPipelineIncrementOperation( - key=tpm_key, - increment_value=batch_usage.total_tokens, - ttl=self.parallel_request_limiter.window_size, - ) - ) - - # Execute increments - if pipeline_operations: - await self.parallel_request_limiter.async_increment_tokens_with_ttl_preservation( - pipeline_operations=pipeline_operations, + rate_limit_response = ( + await self.parallel_request_limiter.atomic_check_and_increment_by_n( + descriptors=descriptors, + increments=increments, parent_otel_span=user_api_key_dict.parent_otel_span, ) + ) + + if rate_limit_response["overall_code"] == "OVER_LIMIT": + for status in rate_limit_response["statuses"]: + if status["code"] == "OVER_LIMIT": + self._raise_rate_limit_error( + status, + descriptors, + batch_usage, + status["rate_limit_type"], + ) async def count_input_file_usage( self, diff --git a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py index 72483d29cdc..f7c0592992f 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py @@ -4,7 +4,7 @@ Dynamic rate limiter v3 - Saturation-aware priority-based rate limiting import os from datetime import datetime -from typing import TYPE_CHECKING, Callable, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Callable, Dict, List, Literal, Optional, Union from fastapi import HTTPException @@ -460,92 +460,128 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): if priority_descriptors: descriptors_to_check.extend(priority_descriptors) - # PHASE 1: Read-only check of ALL limits (no increments) - check_response = await self.v3_limiter.should_rate_limit( - descriptors=descriptors_to_check, + # Atomic check-and-increment for the ENFORCED descriptor set: + # - model_saturation_check is always enforced + # - priority_model is enforced only when saturation crosses threshold + # + # Backed by a Redis Lua script (multi-process atomic) with an + # asyncio.Lock + in-memory fallback for single-process deployments. + # All-or-nothing: if any enforced descriptor would exceed its limit, + # no counter is modified and the response carries "OVER_LIMIT". + enforced_descriptors: List[RateLimitDescriptor] = [model_wide_descriptor] + if priority_descriptors and should_enforce_priority: + enforced_descriptors.extend(priority_descriptors) + + per_request_increment: Dict[Literal["requests", "tokens"], int] = { + "requests": 1, + "tokens": 0, + } + atomic_response = await self.v3_limiter.atomic_check_and_increment_by_n( + descriptors=enforced_descriptors, + increments=[per_request_increment for _ in enforced_descriptors], parent_otel_span=user_api_key_dict.parent_otel_span, - read_only=True, # CRITICAL: Don't increment counters yet ) verbose_proxy_logger.debug( - f"Read-only check: {json.dumps(check_response, indent=2)}" + f"Atomic check+increment response: {json.dumps(atomic_response, indent=2)}" ) - # PHASE 2: Decide which limits to enforce - if check_response["overall_code"] == "OVER_LIMIT": - for status in check_response["statuses"]: - if status["code"] == "OVER_LIMIT": - descriptor_key = status["descriptor_key"] + if atomic_response["overall_code"] == "OVER_LIMIT": + for status in atomic_response["statuses"]: + if status["code"] != "OVER_LIMIT": + continue + descriptor_key = status["descriptor_key"] + if descriptor_key == "model_saturation_check": + raise HTTPException( + status_code=429, + detail={ + "error": f"Model capacity reached for {model}. " + f"Priority: {priority}, " + f"Rate limit type: {status['rate_limit_type']}, " + f"Remaining: {status['limit_remaining']}" + }, + headers={ + "retry-after": str(self.v3_limiter.window_size), + "rate_limit_type": str(status["rate_limit_type"]), + "x-litellm-priority": priority or "default", + }, + ) + if descriptor_key == "priority_model": + verbose_proxy_logger.debug( + f"Enforcing priority limits for {model}, saturation: {saturation:.1%}, " + f"priority: {priority}" + ) + raise HTTPException( + status_code=429, + detail={ + "error": f"Priority-based rate limit exceeded. " + f"Priority: {priority}, " + f"Rate limit type: {status['rate_limit_type']}, " + f"Remaining: {status['limit_remaining']}, " + f"Model saturation: {saturation:.1%}" + }, + headers={ + "retry-after": str(self.v3_limiter.window_size), + "rate_limit_type": str(status["rate_limit_type"]), + "x-litellm-priority": priority or "default", + "x-litellm-saturation": f"{saturation:.2%}", + }, + ) - # Model-wide limit exceeded (ALWAYS enforce) - if descriptor_key == "model_saturation_check": - raise HTTPException( - status_code=429, - detail={ - "error": f"Model capacity reached for {model}. " - f"Priority: {priority}, " - f"Rate limit type: {status['rate_limit_type']}, " - f"Remaining: {status['limit_remaining']}" - }, - headers={ - "retry-after": str(self.v3_limiter.window_size), - "rate_limit_type": str(status["rate_limit_type"]), - "x-litellm-priority": priority or "default", - }, - ) + # Fail-closed guard: overall_code says OVER_LIMIT but no status + # matched a descriptor key we know how to translate into a 429. + # Refuse the request rather than silently fall through and let an + # over-limit request proceed to the model. Without this, a future + # caller wiring an unfamiliar descriptor into enforced_descriptors + # would silently bypass the rate limit. + offending = next( + (s for s in atomic_response["statuses"] if s["code"] == "OVER_LIMIT"), + None, + ) + verbose_proxy_logger.error( + f"Dynamic rate limiter: OVER_LIMIT response with unknown " + f"descriptor_key(s) — refusing request. response={atomic_response}" + ) + raise HTTPException( + status_code=429, + detail={ + "error": "Rate limit exceeded", + "descriptor_key": ( + offending["descriptor_key"] if offending else "unknown" + ), + "rate_limit_type": ( + str(offending["rate_limit_type"]) if offending else "unknown" + ), + }, + headers={ + "retry-after": str(self.v3_limiter.window_size), + "x-litellm-priority": priority or "default", + }, + ) - # Priority limit exceeded (ONLY enforce when saturated) - elif descriptor_key == "priority_model" and should_enforce_priority: - verbose_proxy_logger.debug( - f"Enforcing priority limits for {model}, saturation: {saturation:.1%}, " - f"priority: {priority}" - ) - raise HTTPException( - status_code=429, - detail={ - "error": f"Priority-based rate limit exceeded. " - f"Priority: {priority}, " - f"Rate limit type: {status['rate_limit_type']}, " - f"Remaining: {status['limit_remaining']}, " - f"Model saturation: {saturation:.1%}" - }, - headers={ - "retry-after": str(self.v3_limiter.window_size), - "rate_limit_type": str(status["rate_limit_type"]), - "x-litellm-priority": priority or "default", - "x-litellm-saturation": f"{saturation:.2%}", - }, - ) - - # PHASE 3: Increment counters separately to avoid early-exit issues - # Model counter must ALWAYS increment, but priority counter might be over limit - # If we increment them together, v3_limiter's in-memory check will exit early - # and skip incrementing the model counter - - # Step 3a: Increment model-wide counter (always) - model_increment_response = await self.v3_limiter.should_rate_limit( - descriptors=[model_wide_descriptor], - parent_otel_span=user_api_key_dict.parent_otel_span, - read_only=False, - ) - - # Step 3b: Increment priority counter (may be over limit, but we still track it) - if priority_descriptors: - priority_increment_response = await self.v3_limiter.should_rate_limit( + # If priority is NOT enforced (saturation below threshold) but + # priority_descriptors exist, increment them for tracking only — no + # check, no rollback. This matches the prior tracking semantics. + # + # Using the non-atomic should_rate_limit (instead of + # atomic_check_and_increment_by_n) is intentional here: we don't want + # to enforce the limit, we only want to bump the counter so the + # priority allocation has accurate usage when it later becomes + # enforced. The increment-then-check semantics of should_rate_limit + # are fine because we ignore the OVER_LIMIT response. + if priority_descriptors and not should_enforce_priority: + priority_tracking_response = await self.v3_limiter.should_rate_limit( descriptors=priority_descriptors, parent_otel_span=user_api_key_dict.parent_otel_span, read_only=False, ) - - # Combine responses for post-call hook - combined_response = { - "overall_code": model_increment_response["overall_code"], - "statuses": model_increment_response["statuses"] - + priority_increment_response["statuses"], + data["litellm_proxy_rate_limit_response"] = { + "overall_code": atomic_response["overall_code"], + "statuses": atomic_response["statuses"] + + priority_tracking_response["statuses"], } - data["litellm_proxy_rate_limit_response"] = combined_response else: - data["litellm_proxy_rate_limit_response"] = model_increment_response + data["litellm_proxy_rate_limit_response"] = atomic_response async def async_pre_call_hook( self, diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index f29bbd2d9d5..4497e64c17f 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -4,6 +4,7 @@ This is a rate limiter implementation based on a similar one by Envoy proxy. This is currently in development and not yet ready for production. """ +import asyncio import binascii import os from datetime import datetime @@ -80,6 +81,90 @@ end return results """ +CHECK_AND_INCREMENT_BY_N_SCRIPT = """ +-- Atomic check-and-increment-by-N across one or more descriptors. +-- All-or-nothing: if any descriptor would exceed its limit, no counter is +-- modified. +-- +-- Uses Redis server time (`redis.call('TIME')`) instead of a client-supplied +-- timestamp so that window resets are deterministic across replicas with +-- skewed wall-clocks. This prevents a clock-skew-induced reopening of the +-- TOCTOU window across multi-replica deployments. +-- +-- KEYS layout: pairs of (window_key, counter_key), one pair per descriptor. +-- ARGV layout: per-descriptor 4-tuple, starting at ARGV[1]: +-- ARGV[(i-1)*4 + 1] = limit +-- ARGV[(i-1)*4 + 2] = increment +-- ARGV[(i-1)*4 + 3] = ttl_seconds (counter TTL when window resets) +-- ARGV[(i-1)*4 + 4] = window_size_seconds (sliding-window length) +-- +-- Return on success: { 0, new_counter_1, new_counter_2, ... } +-- Return on over-limit: { 1, descriptor_index, current_counter, limit } +local time_reply = redis.call('TIME') +local now = tonumber(time_reply[1]) +local descriptor_count = #KEYS / 2 + +-- Pass 1: read state, validate. Abort without writing if any over limit. +local descriptor_state = {} +for i = 1, descriptor_count do + local window_key = KEYS[(i - 1) * 2 + 1] + local counter_key = KEYS[(i - 1) * 2 + 2] + local arg_base = (i - 1) * 4 + 1 + local limit = tonumber(ARGV[arg_base]) + local increment = tonumber(ARGV[arg_base + 1]) + local window_size = tonumber(ARGV[arg_base + 3]) + + local window_start = redis.call('GET', window_key) + local window_expired = (not window_start) or + ((now - tonumber(window_start)) >= window_size) + + local current_counter + if window_expired then + current_counter = 0 + else + current_counter = tonumber(redis.call('GET', counter_key) or 0) + end + + if current_counter + increment > limit then + return { 1, i, current_counter, limit } + end + + descriptor_state[i] = { window_expired, current_counter } +end + +-- Pass 2: all checks passed. Apply increments. +local results = { 0 } +for i = 1, descriptor_count do + local window_key = KEYS[(i - 1) * 2 + 1] + local counter_key = KEYS[(i - 1) * 2 + 2] + local arg_base = (i - 1) * 4 + 1 + local increment = tonumber(ARGV[arg_base + 1]) + local ttl = tonumber(ARGV[arg_base + 2]) + local window_size = tonumber(ARGV[arg_base + 3]) + + local window_expired = descriptor_state[i][1] + + if window_expired then + redis.call('SET', window_key, tostring(now)) + redis.call('SET', counter_key, increment) + redis.call('EXPIRE', window_key, window_size) + if ttl > 0 then + redis.call('EXPIRE', counter_key, ttl) + end + table.insert(results, increment) + else + local new_counter = redis.call('INCRBY', counter_key, increment) + local current_ttl = redis.call('TTL', counter_key) + if current_ttl == -1 and ttl > 0 then + redis.call('EXPIRE', counter_key, ttl) + end + table.insert(results, new_counter) + end +end + +return results +""" + TOKEN_INCREMENT_SCRIPT = """ local results = {} @@ -162,15 +247,37 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): TOKEN_INCREMENT_SCRIPT ) ) + self.check_and_increment_by_n_script = ( + self.internal_usage_cache.dual_cache.redis_cache.async_register_script( + CHECK_AND_INCREMENT_BY_N_SCRIPT + ) + ) else: self.batch_rate_limiter_script = None self.token_increment_script = None + self.check_and_increment_by_n_script = None self.window_size = int(os.getenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", 60)) # Batch rate limiter (lazy loaded) self._batch_rate_limiter: Optional[Any] = None + # Serializes multi-phase check+increment sequences (batch + dynamic + # limiters) within this process to close the TOCTOU window between + # read-only check and counter increment. Multi-replica deployments + # additionally rely on Redis Lua atomicity for cross-process safety. + # + # Coarse granularity: this single lock serializes ALL atomic check+ + # increment operations across batch and dynamic limiters on this + # instance. A slow batch input-file fetch (which happens upstream of + # the lock) does not block here, but Redis Lua latency does. If + # contention shows up under load (visible as p99 latency spikes + # correlated with batch traffic), shard to a per-descriptor-key lock + # via a `weakref.WeakValueDictionary[str, asyncio.Lock]`. Punted as a + # follow-up because Lua dominates wall-time and the lock is held for + # one round-trip. + self._check_and_increment_lock = asyncio.Lock() + def _get_batch_rate_limiter(self) -> Optional[Any]: """Get or lazy-load the batch rate limiter.""" if self._batch_rate_limiter is None: @@ -588,6 +695,281 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) return rate_limit_response + async def atomic_check_and_increment_by_n( + self, + descriptors: List[RateLimitDescriptor], + increments: List[Dict[Literal["requests", "tokens"], int]], + parent_otel_span: Optional[Span] = None, + ) -> RateLimitResponse: + """ + Atomic check-and-increment-by-N across one or more descriptors. + + All-or-nothing: if any descriptor would exceed its limit, no counter is + modified and the response carries `overall_code = "OVER_LIMIT"` with + the offending descriptor's status. Closes the TOCTOU window between + read and increment in both single-process and multi-process (Redis) + deployments. + + Args: + descriptors: rate-limit descriptors to check + increments: per-descriptor increment amounts, indexed parallel to + `descriptors`. Each entry is `{"requests": int, "tokens": int}` + — values default to 0 when a descriptor has no matching limit. + + Returns: + RateLimitResponse with one status per (descriptor, rate_limit_type) + counter, mirroring `should_rate_limit`'s shape. + """ + if len(descriptors) != len(increments): + raise ValueError( + "atomic_check_and_increment_by_n: descriptors and increments " + "must have the same length" + ) + + keys: List[str] = [] + per_counter_meta: List[Dict[str, Any]] = [] + script_args: List[Any] = [] + + for descriptor, increment_amounts in zip(descriptors, increments): + descriptor_key = descriptor["key"] + descriptor_value = descriptor["value"] + rate_limit: RateLimitDescriptorRateLimitObject = ( + descriptor.get("rate_limit") or RateLimitDescriptorRateLimitObject() + ) + window_size = rate_limit.get("window_size") or self.window_size + window_key = f"{{{descriptor_key}:{descriptor_value}}}:window" + + for rate_limit_type in ("requests", "tokens"): + rlt: Literal["requests", "tokens"] = cast( + Literal["requests", "tokens"], rate_limit_type + ) + if rlt == "requests": + limit_value = rate_limit.get("requests_per_unit") + inc_amount = int(increment_amounts.get("requests", 0) or 0) + else: + limit_value = rate_limit.get("tokens_per_unit") + inc_amount = int(increment_amounts.get("tokens", 0) or 0) + if limit_value is None or inc_amount <= 0: + continue + counter_key = self.create_rate_limit_keys( + descriptor_key, descriptor_value, rlt + ) + # Counter-key TTL and window_size are conceptually distinct + # ("how long the counter Redis key lives" vs "how long the + # sliding window is"). They happen to be equal today because + # we have no descriptor type that needs them apart, but they + # are kept as separate variables here so a future custom-TTL + # descriptor doesn't reintroduce a silent expiry bug. Both + # the Lua script and the in-memory fallback read these from + # their respective ARGV / meta slots. + ttl_seconds = int(window_size) + window_size_seconds = int(window_size) + keys.extend([window_key, counter_key]) + # Per-counter 4-tuple matches the Lua ARGV layout exactly: + # [limit, increment, ttl_seconds, window_size_seconds]. + script_args.extend( + [ + int(limit_value), + inc_amount, + ttl_seconds, + window_size_seconds, + ] + ) + per_counter_meta.append( + { + "descriptor_key": descriptor_key, + "current_limit": int(limit_value), + "rate_limit_type": rlt, + "window_key": window_key, + "counter_key": counter_key, + "increment": inc_amount, + "ttl": ttl_seconds, + "window_size": window_size_seconds, + } + ) + + if not keys: + return RateLimitResponse(overall_code="OK", statuses=[]) + + # Multi-process atomicity via Redis Lua. Single-process atomicity + # falls back to the asyncio.Lock + in-memory sliding window below. + # Note: in-memory state diverges from Redis state — if Lua fails + # mid-write, retrying via in-memory may double-count. See fallback + # warning below. + if self.check_and_increment_by_n_script is not None: + try: + raw = await self.check_and_increment_by_n_script( + keys=keys, + args=script_args, + ) + return self._build_atomic_response(raw, per_counter_meta) + except Exception as e: + # Escalated from warning to error: Lua failures (script timeout, + # Redis OOM, network partition) leave counter state ambiguous. + # The fallback path below uses LOCAL DualCache, which is a + # different store from Redis — counters here will diverge from + # Redis until that key's window expires (TTL bounds divergence). + # Operators should alert on this log line; sustained occurrences + # indicate Redis health degradation that may erode rate-limit + # accuracy. + verbose_proxy_logger.error( + f"atomic_check_and_increment_by_n: Redis Lua execution " + f"failed ({type(e).__name__}: {e}). Falling back to " + f"in-memory enforcement — counters will diverge from Redis " + f"state until window expires (window_size={self.window_size}s)." + ) + + async with self._check_and_increment_lock: + return await self._atomic_check_and_increment_in_memory( + per_counter_meta=per_counter_meta, + parent_otel_span=parent_otel_span, + ) + + def _build_atomic_response( + self, + raw: List[Any], + per_counter_meta: List[Dict[str, Any]], + ) -> RateLimitResponse: + """Convert Lua script return value to RateLimitResponse. + + Indexing invariant: `per_counter_meta` and `KEYS` are parallel-indexed + at the COUNTER level, not the descriptor level. A descriptor with both + RPM and TPM limits emits two `(window_key, counter_key)` pairs and + two meta entries — one per counter. The Lua script's loop variable + `i` therefore enumerates counters, and the over-limit return tuple + `{1, i, ...}` carries a counter index that maps directly to + `per_counter_meta[i - 1]`. Keep these arrays parallel at the counter + level when modifying this code. + """ + if not raw: + return RateLimitResponse(overall_code="OK", statuses=[]) + + status_code = int(raw[0]) + if status_code == 1: + # Over limit: { 1, counter_index (1-based), current_counter, limit } + descriptor_index = int(raw[1]) - 1 + current_counter = int(raw[2]) + limit = int(raw[3]) + meta = per_counter_meta[descriptor_index] + return RateLimitResponse( + overall_code="OVER_LIMIT", + statuses=[ + RateLimitStatus( + code="OVER_LIMIT", + current_limit=limit, + limit_remaining=max(0, limit - current_counter), + rate_limit_type=meta["rate_limit_type"], + descriptor_key=meta["descriptor_key"], + ) + ], + ) + + statuses: List[RateLimitStatus] = [] + for meta, new_counter in zip(per_counter_meta, raw[1:]): + statuses.append( + RateLimitStatus( + code="OK", + current_limit=meta["current_limit"], + limit_remaining=max(0, meta["current_limit"] - int(new_counter)), + rate_limit_type=meta["rate_limit_type"], + descriptor_key=meta["descriptor_key"], + ) + ) + return RateLimitResponse(overall_code="OK", statuses=statuses) + + async def _atomic_check_and_increment_in_memory( + self, + per_counter_meta: List[Dict[str, Any]], + parent_otel_span: Optional[Span] = None, + ) -> RateLimitResponse: + """In-memory all-or-nothing check-and-increment. Caller holds lock. + + Reads/writes the LOCAL DualCache (`local_only=True`) — note this is + a different store from Redis. When this fallback fires after a Lua + failure, in-memory counters will diverge from Redis until each key's + window expires (TTL bounds divergence). + """ + # Use a single 'now' for the duration of this critical section so all + # descriptors evaluate window expiry consistently. + now_int = int(self._get_current_time().timestamp()) + + # Pass 1: read state, validate. + descriptor_state: List[Dict[str, Any]] = [] + for meta in per_counter_meta: + window_size = meta["window_size"] + window_start = await self.internal_usage_cache.async_get_cache( + key=meta["window_key"], + litellm_parent_otel_span=parent_otel_span, + local_only=True, + ) + window_expired = ( + window_start is None or (now_int - int(window_start)) >= window_size + ) + current_counter = ( + 0 + if window_expired + else int( + await self.internal_usage_cache.async_get_cache( + key=meta["counter_key"], + litellm_parent_otel_span=parent_otel_span, + local_only=True, + ) + or 0 + ) + ) + if current_counter + meta["increment"] > meta["current_limit"]: + return RateLimitResponse( + overall_code="OVER_LIMIT", + statuses=[ + RateLimitStatus( + code="OVER_LIMIT", + current_limit=meta["current_limit"], + limit_remaining=max( + 0, meta["current_limit"] - current_counter + ), + rate_limit_type=meta["rate_limit_type"], + descriptor_key=meta["descriptor_key"], + ) + ], + ) + descriptor_state.append( + {"window_expired": window_expired, "current": current_counter} + ) + + # Pass 2: apply increments. + statuses: List[RateLimitStatus] = [] + for meta, state in zip(per_counter_meta, descriptor_state): + new_counter = ( + meta["increment"] + if state["window_expired"] + else state["current"] + meta["increment"] + ) + if state["window_expired"]: + await self.internal_usage_cache.async_set_cache( + key=meta["window_key"], + value=str(now_int), + ttl=meta["window_size"], + litellm_parent_otel_span=parent_otel_span, + local_only=True, + ) + await self.internal_usage_cache.async_set_cache( + key=meta["counter_key"], + value=new_counter, + ttl=meta["ttl"], + litellm_parent_otel_span=parent_otel_span, + local_only=True, + ) + statuses.append( + RateLimitStatus( + code="OK", + current_limit=meta["current_limit"], + limit_remaining=max(0, meta["current_limit"] - new_counter), + rate_limit_type=meta["rate_limit_type"], + descriptor_key=meta["descriptor_key"], + ) + ) + return RateLimitResponse(overall_code="OK", statuses=statuses) + def create_organization_rate_limit_descriptor( self, user_api_key_dict: UserAPIKeyAuth, requested_model: Optional[str] = None ) -> List[RateLimitDescriptor]: diff --git a/litellm/proxy/management_endpoints/access_group_endpoints.py b/litellm/proxy/management_endpoints/access_group_endpoints.py index caaec12f7a3..ceaef20a8d0 100644 --- a/litellm/proxy/management_endpoints/access_group_endpoints.py +++ b/litellm/proxy/management_endpoints/access_group_endpoints.py @@ -236,13 +236,12 @@ async def _patch_key_caches_add_access_group( ) -> None: """Patch cached key objects to include access_group_id.""" for token in key_tokens: - cached_key = await user_api_key_cache.async_get_cache(key=token) + cached_key = await user_api_key_cache.async_get_cache( + key=token, + model_type=UserAPIKeyAuth, + ) if cached_key is None: continue - if isinstance(cached_key, dict): - cached_key = UserAPIKeyAuth(**cached_key) - if not isinstance(cached_key, UserAPIKeyAuth): - continue if cached_key.access_group_ids is None: cached_key.access_group_ids = [access_group_id] elif access_group_id not in cached_key.access_group_ids: @@ -267,12 +266,11 @@ async def _patch_key_caches_remove_access_group( ) -> None: """Patch cached key objects to remove access_group_id.""" for token in key_tokens: - cached_key = await user_api_key_cache.async_get_cache(key=token) - if cached_key is None: - continue - if isinstance(cached_key, dict): - cached_key = UserAPIKeyAuth(**cached_key) - if isinstance(cached_key, UserAPIKeyAuth) and cached_key.access_group_ids: + cached_key = await user_api_key_cache.async_get_cache( + key=token, + model_type=UserAPIKeyAuth, + ) + if cached_key is not None and cached_key.access_group_ids: cached_key.access_group_ids = [ ag for ag in cached_key.access_group_ids if ag != access_group_id ] diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 2485aea14f1..a01f5e63211 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -27,7 +27,7 @@ from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, s import litellm from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid -from litellm.caching import DualCache +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.constants import ( LENGTH_OF_LITELLM_GENERATED_KEY, LITELLM_PROXY_ADMIN_NAME, @@ -1059,7 +1059,7 @@ async def _check_project_key_limits( project_id: str, data: Union[GenerateKeyRequest, UpdateKeyRequest], prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, ) -> None: """ Validate that key's models and budget respect its project's limits. @@ -1834,7 +1834,7 @@ async def _process_single_key_update( user_api_key_dict: UserAPIKeyAuth, litellm_changed_by: Optional[str], prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: Any, llm_router: Optional[Router], user_custom_key_update: Optional[Callable] = None, @@ -3298,7 +3298,7 @@ async def _team_key_deletion_check( user_api_key_dict: UserAPIKeyAuth, key_info: LiteLLM_VerificationToken, prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, ): is_team_key = _is_team_key(data=key_info) @@ -3341,7 +3341,7 @@ async def _team_key_deletion_check( async def can_modify_verification_token( key_info: LiteLLM_VerificationToken, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, ) -> bool: @@ -3415,7 +3415,7 @@ async def can_modify_verification_token( async def delete_verification_tokens( tokens: List, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, user_api_key_dict: UserAPIKeyAuth, litellm_changed_by: Optional[str] = None, ) -> Tuple[Optional[Dict], List[LiteLLM_VerificationToken]]: @@ -3605,7 +3605,7 @@ async def _persist_deleted_verification_tokens( async def delete_key_aliases( key_aliases: List[str], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, litellm_changed_by: Optional[str] = None, @@ -3862,7 +3862,7 @@ async def _execute_virtual_key_regeneration( data: Optional[RegenerateKeyRequest], user_api_key_dict: UserAPIKeyAuth, litellm_changed_by: Optional[str], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, ) -> GenerateKeyResponse: """Generate new token, update DB, invalidate cache, and return response.""" @@ -4152,7 +4152,7 @@ async def _check_proxy_or_team_admin_for_key( key_in_db: LiteLLM_VerificationToken, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, ) -> None: if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value: return @@ -5173,7 +5173,7 @@ async def _check_key_admin_access( user_api_key_dict: UserAPIKeyAuth, hashed_token: str, prisma_client: Any, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, route: str, ) -> None: """ diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index c4564a4eb04..9dfc67370fe 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -39,9 +39,9 @@ from fastapi import APIRouter, Depends, Header, HTTPException, Request, status from fastapi.responses import RedirectResponse import litellm +from litellm.caching.dual_cache import DualCache from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid -from litellm.caching import DualCache from litellm.constants import ( CLI_SSO_SESSION_CACHE_KEY_PREFIX, CLI_SSO_SESSION_TTL_SECONDS, @@ -75,7 +75,11 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken, get_user_object -from litellm.proxy.auth.auth_utils import _get_request_ip_address, _has_user_setup_sso +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache +from litellm.proxy.auth.auth_utils import ( + _get_request_ip_address, + _has_user_setup_sso, +) from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.admin_ui_utils import ( @@ -1301,7 +1305,7 @@ async def get_existing_user_info_from_db( user_id: Optional[str], user_email: Optional[str], prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, ) -> Optional[LiteLLM_UserTable]: try: @@ -1325,7 +1329,7 @@ async def get_existing_user_info_from_db( async def get_user_info_from_db( result: Union[CustomOpenID, OpenID, dict], prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, user_email: Optional[str], user_defined_values: Optional[SSOUserDefinedValues], @@ -1445,7 +1449,7 @@ async def _sync_user_role_from_jwt_role_map( received_response: Optional[dict], user_info: Optional[Union[LiteLLM_UserTable, NewUserResponse]], prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, user_defined_values: Optional[SSOUserDefinedValues], ) -> None: """ @@ -1484,11 +1488,8 @@ async def _sync_user_role_from_jwt_role_map( user_info.user_role = mapped_role.value await user_api_key_cache.async_set_cache( key=user_info.user_id, - value=( - user_info.model_dump() - if hasattr(user_info, "model_dump") - else dict(user_info) - ), + value=user_info, + model_type=LiteLLM_UserTable, ) diff --git a/litellm/proxy/management_helpers/team_member_permission_checks.py b/litellm/proxy/management_helpers/team_member_permission_checks.py index e035168ca00..50339210a6e 100644 --- a/litellm/proxy/management_helpers/team_member_permission_checks.py +++ b/litellm/proxy/management_helpers/team_member_permission_checks.py @@ -1,6 +1,5 @@ from typing import List, Optional -from litellm.caching import DualCache from litellm.proxy._types import ( KeyManagementRoutes, LiteLLM_TeamTableCachedObj, @@ -12,6 +11,7 @@ from litellm.proxy._types import ( ProxyException, UserAPIKeyAuth, ) +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.auth.auth_checks import get_team_object from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.utils import PrismaClient @@ -65,7 +65,7 @@ class TeamMemberPermissionChecks: user_api_key_dict: UserAPIKeyAuth, route: KeyManagementRoutes, prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, existing_key_row: LiteLLM_VerificationToken, ): """ diff --git a/litellm/proxy/middleware/prometheus_auth_middleware.py b/litellm/proxy/middleware/prometheus_auth_middleware.py index 6bdff59da52..3b30fd3d63c 100644 --- a/litellm/proxy/middleware/prometheus_auth_middleware.py +++ b/litellm/proxy/middleware/prometheus_auth_middleware.py @@ -3,6 +3,7 @@ Prometheus Auth Middleware - Pure ASGI implementation """ import json +from typing import Any, List, MutableMapping from fastapi import Request from starlette.types import ASGIApp, Receive, Scope, Send @@ -40,8 +41,17 @@ class PrometheusAuthMiddleware: # Only run auth if configured to do so if litellm.require_auth_for_metrics_endpoint is True: - # Construct Request only when auth is actually needed - request = Request(scope, receive) + # user_api_key_auth reads the request body, which consumes ASGI `receive`. + # Buffer those messages and replay them for the inner app; otherwise a + # successful auth would forward an exhausted receive and /metrics hangs. + buffered_messages: List[MutableMapping[str, Any]] = [] + + async def receive_for_auth() -> MutableMapping[str, Any]: + message = await receive() + buffered_messages.append(message) + return message + + request = Request(scope, receive_for_auth) api_key = request.headers.get(_AUTHORIZATION_HEADER) or "" try: @@ -70,5 +80,18 @@ class PrometheusAuthMiddleware: ) return + replay_idx = 0 + + async def receive_replay() -> MutableMapping[str, Any]: + nonlocal replay_idx + if replay_idx < len(buffered_messages): + msg = buffered_messages[replay_idx] + replay_idx += 1 + return msg + return await receive() + + await self.app(scope, receive_replay, send) + return + # Pass through to the inner application await self.app(scope, receive, send) diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/assembly_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/assembly_passthrough_logging_handler.py index a8c5562d4d6..6277f6b4a75 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/assembly_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/assembly_passthrough_logging_handler.py @@ -1,6 +1,7 @@ import asyncio import json import time +import urllib.parse from datetime import datetime from typing import Literal, Optional from urllib.parse import urlparse @@ -203,8 +204,16 @@ class AssemblyAIPassthroughLoggingHandler: ) if _api_key is None: raise ValueError("AssemblyAI API key not found") + if ( + any(c in transcript_id for c in ("/", "\\", "#", "?")) + or ".." in transcript_id + ): + raise ValueError( + f"Invalid transcript_id {transcript_id!r}: contains disallowed characters" + ) + safe_transcript_id = urllib.parse.quote(transcript_id, safe="") try: - url = f"{_base_url}/v2/transcript/{transcript_id}" + url = f"{_base_url}/v2/transcript/{safe_transcript_id}" headers = { "Authorization": f"Bearer {_api_key}", "Content-Type": "application/json", diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 6cba6a3e96b..a2bdc3ab3cf 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -78,8 +78,11 @@ from litellm.proxy._types import ( InvitationNew, InvitationUpdate, Litellm_EntityType, + LiteLLM_EndUserTable, LiteLLM_JWTAuth, + LiteLLM_TagTable, LiteLLM_TeamTable, + LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, LitellmUserRoles, PassThroughGenericEndpoint, @@ -94,6 +97,7 @@ from litellm.proxy._types import ( UI_TEAM_ID, UserAPIKeyAuth, ) +from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec from litellm.proxy.common_utils.callback_utils import ( normalize_callback_names, process_callback, @@ -206,6 +210,7 @@ from litellm import Router from litellm._logging import verbose_proxy_logger, verbose_router_logger from litellm.caching.caching import DualCache, RedisCache from litellm.caching.redis_cluster_cache import RedisClusterCache +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.constants import ( _REALTIME_BODY_CACHE_SIZE, APSCHEDULER_COALESCE, @@ -1612,7 +1617,7 @@ prisma_client: Optional[PrismaClient] = None shared_aiohttp_session: Optional["ClientSession"] = ( None # Global shared session for connection reuse ) -user_api_key_cache = DualCache( +user_api_key_cache: UserApiKeyCache = UserApiKeyCache( default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value ) spend_counter_cache = DualCache( @@ -2014,14 +2019,16 @@ async def update_cache( # noqa: PLR0915 else: hashed_token = token verbose_proxy_logger.debug("_update_key_cache: hashed_token=%s", hashed_token) - existing_spend_obj: LiteLLM_VerificationTokenView = await user_api_key_cache.async_get_cache(key=hashed_token) # type: ignore + existing_spend_obj = await user_api_key_cache.async_get_cache( + key=hashed_token, model_type=UserAPIKeyAuth + ) verbose_proxy_logger.debug( f"_update_key_cache: existing_spend_obj={existing_spend_obj}" ) if existing_spend_obj is None: return - else: - existing_spend = existing_spend_obj.spend + + existing_spend = existing_spend_obj.spend or 0.0 # Calculate the new cost by adding the existing cost and response_cost new_spend = existing_spend + response_cost @@ -2079,41 +2086,48 @@ async def update_cache( # noqa: PLR0915 existing_team_member_spend + response_cost ) - # Update the cost column for the given token + # Existing spend_obj is mutated; UserApiKeyCache.async_set_cache_pipeline turns + # BaseModel values into dicts for Redis (same Codec path as async_set_cache). existing_spend_obj.spend = new_spend values_to_update_in_cache.append((hashed_token, existing_spend_obj)) ### UPDATE USER SPEND ### async def _update_user_cache(): ## UPDATE CACHE FOR USER ID + GLOBAL PROXY + if response_cost is None: + return user_ids = [user_id] try: for _id in user_ids: # Fetch the existing cost for the given user if _id is None: continue - existing_spend_obj = await user_api_key_cache.async_get_cache(key=_id) - if existing_spend_obj is None: + cached_user = await user_api_key_cache.async_get_cache(key=_id) + if cached_user is None: # do nothing if there is no cache value return + existing_spend_obj = CacheCodec.deserialize( + cached_user, LiteLLM_UserTable + ) + if existing_spend_obj is None: + return verbose_proxy_logger.debug( f"_update_user_db: existing spend: {existing_spend_obj}; response_cost: {response_cost}" ) - if isinstance(existing_spend_obj, dict): - existing_spend = existing_spend_obj["spend"] - else: - existing_spend = existing_spend_obj.spend + existing_spend = existing_spend_obj.spend or 0.0 # Calculate the new cost by adding the existing cost and response_cost new_spend = existing_spend + response_cost - # Update the cost column for the given user - if isinstance(existing_spend_obj, dict): - existing_spend_obj["spend"] = new_spend - values_to_update_in_cache.append((_id, existing_spend_obj)) - else: - existing_spend_obj.spend = new_spend - values_to_update_in_cache.append((_id, existing_spend_obj.json())) + existing_spend_obj.spend = new_spend + values_to_update_in_cache.append( + ( + _id, + CacheCodec.serialize( + existing_spend_obj, model_type=LiteLLM_UserTable + ), + ) + ) ## UPDATE GLOBAL PROXY ## global_proxy_spend = await user_api_key_cache.async_get_cache( key="{}:spend".format(litellm_proxy_admin_name) @@ -2145,31 +2159,33 @@ async def update_cache( # noqa: PLR0915 _id = "end_user_id:{}".format(end_user_id) try: # Fetch the existing cost for the given user - existing_spend_obj = await user_api_key_cache.async_get_cache(key=_id) - if existing_spend_obj is None: + cached_end_user = await user_api_key_cache.async_get_cache(key=_id) + if cached_end_user is None: # if user does not exist in LiteLLM_UserTable, create a new user # do nothing if end-user not in api key cache return + existing_spend_obj = CacheCodec.deserialize( + cached_end_user, LiteLLM_EndUserTable + ) + if existing_spend_obj is None: + return verbose_proxy_logger.debug( f"_update_end_user_db: existing spend: {existing_spend_obj}; response_cost: {response_cost}" ) - if existing_spend_obj is None: - existing_spend = 0 - else: - if isinstance(existing_spend_obj, dict): - existing_spend = existing_spend_obj["spend"] - else: - existing_spend = existing_spend_obj.spend + + existing_spend = existing_spend_obj.spend or 0.0 # Calculate the new cost by adding the existing cost and response_cost new_spend = existing_spend + response_cost - # Update the cost column for the given user - if isinstance(existing_spend_obj, dict): - existing_spend_obj["spend"] = new_spend - values_to_update_in_cache.append((_id, existing_spend_obj)) - else: - existing_spend_obj.spend = new_spend - values_to_update_in_cache.append((_id, existing_spend_obj.json())) + existing_spend_obj.spend = new_spend + values_to_update_in_cache.append( + ( + _id, + CacheCodec.serialize( + existing_spend_obj, model_type=LiteLLM_EndUserTable + ), + ) + ) except Exception as e: verbose_proxy_logger.warning( "Spend tracking - failed to update end user spend in cache. " @@ -2188,36 +2204,32 @@ async def update_cache( # noqa: PLR0915 _id = "team_id:{}".format(team_id) try: - # Fetch the existing cost for the given user - existing_spend_obj: Optional[LiteLLM_TeamTable] = ( - await user_api_key_cache.async_get_cache(key=_id) + cached_team = await user_api_key_cache.async_get_cache(key=_id) + if cached_team is None: + # do nothing if team not in api key cache + return + existing_spend_obj: Optional[LiteLLM_TeamTableCachedObj] = ( + CacheCodec.deserialize(cached_team, LiteLLM_TeamTableCachedObj) ) if existing_spend_obj is None: - # do nothing if team not in api key cache return verbose_proxy_logger.debug( f"_update_team_db: existing spend: {existing_spend_obj}; response_cost: {response_cost}" ) - if existing_spend_obj is None: - existing_spend: Optional[float] = 0.0 - else: - if isinstance(existing_spend_obj, dict): - existing_spend = existing_spend_obj["spend"] - else: - existing_spend = existing_spend_obj.spend - if existing_spend is None: - existing_spend = 0.0 + existing_spend: float = existing_spend_obj.spend or 0.0 # Calculate the new cost by adding the existing cost and response_cost new_spend = existing_spend + response_cost - # Update the cost column for the given user - if isinstance(existing_spend_obj, dict): - existing_spend_obj["spend"] = new_spend - values_to_update_in_cache.append((_id, existing_spend_obj)) - else: - existing_spend_obj.spend = new_spend - values_to_update_in_cache.append((_id, existing_spend_obj)) + existing_spend_obj.spend = new_spend + values_to_update_in_cache.append( + ( + _id, + CacheCodec.serialize( + existing_spend_obj, model_type=LiteLLM_TeamTableCachedObj + ), + ) + ) except Exception as e: verbose_proxy_logger.warning( "Spend tracking - failed to update team spend in cache. " @@ -2244,32 +2256,32 @@ async def update_cache( # noqa: PLR0915 cache_key = f"tag:{tag_name}" # Fetch the existing tag object from cache - existing_tag_obj = await user_api_key_cache.async_get_cache( - key=cache_key - ) - if existing_tag_obj is None: + cached_tag = await user_api_key_cache.async_get_cache(key=cache_key) + if cached_tag is None: # do nothing if tag not in api key cache continue + existing_tag_obj = CacheCodec.deserialize(cached_tag, LiteLLM_TagTable) + if existing_tag_obj is None: + continue + verbose_proxy_logger.debug( f"_update_tag_cache: existing spend for tag={tag_name}: {existing_tag_obj}; response_cost: {response_cost}" ) - if isinstance(existing_tag_obj, dict): - existing_spend = existing_tag_obj.get("spend", 0) or 0 - else: - existing_spend = getattr(existing_tag_obj, "spend", 0) or 0 - + existing_spend = existing_tag_obj.spend or 0.0 # Calculate the new cost by adding the existing cost and response_cost new_spend = existing_spend + response_cost - # Update the spend column for the given tag - if isinstance(existing_tag_obj, dict): - existing_tag_obj["spend"] = new_spend - values_to_update_in_cache.append((cache_key, existing_tag_obj)) - else: - existing_tag_obj.spend = new_spend - values_to_update_in_cache.append((cache_key, existing_tag_obj)) + existing_tag_obj.spend = new_spend + values_to_update_in_cache.append( + ( + cache_key, + CacheCodec.serialize( + existing_tag_obj, model_type=LiteLLM_TagTable + ), + ) + ) except Exception as e: verbose_proxy_logger.warning( "Spend tracking - failed to update tag spend in cache. " @@ -2937,8 +2949,9 @@ class ProxyConfig: def _init_cache( self, cache_params: dict, + enable_redis_auth_cache: bool = False, ): - global redis_usage_cache, llm_router + global redis_usage_cache, llm_router, general_settings from litellm import Cache if "default_in_memory_ttl" in cache_params: @@ -2954,7 +2967,29 @@ class ProxyConfig: ): ## INIT PROXY REDIS USAGE CLIENT ## redis_usage_cache = litellm.cache.cache - spend_counter_cache.redis_cache = redis_usage_cache + spend_counter_cache.attach_redis_cache( + redis_usage_cache, + default_redis_ttl=litellm.default_redis_ttl, + ) + # Note: PKCE verifier storage uses redis_usage_cache directly (not + # user_api_key_cache) to avoid routing all API-key lookups through Redis. + if enable_redis_auth_cache is True: + user_api_key_cache.attach_redis_cache( + redis_usage_cache, + default_redis_ttl=litellm.default_redis_ttl, + ) + verbose_proxy_logger.info( + "enable_redis_auth_cache=True: attached Redis to " + "user_api_key_cache — virtual-key lookups are now " + "shared across all proxy workers." + ) + else: + verbose_proxy_logger.info( + "enable_redis_auth_cache is not set: user_api_key_cache " + "remains in-memory only (per-worker). Set " + "litellm_settings.enable_redis_auth_cache: true to share " + "the auth cache across workers and reduce DB load." + ) litellm_config_cache.redis_cache = redis_usage_cache # Note: PKCE verifier storage uses redis_usage_cache directly (not # user_api_key_cache) to avoid routing all API-key lookups through Redis. @@ -3280,7 +3315,13 @@ class ProxyConfig: cache_params[key] = get_secret(value) ## to pass a complete url, or set ssl=True, etc. just set it as `os.environ[REDIS_URL] = `, _redis.py checks for REDIS specific environment variables - self._init_cache(cache_params=cache_params) + self._init_cache( + cache_params=cache_params, + enable_redis_auth_cache=litellm_settings.get( + "enable_redis_auth_cache", False + ) + is True, + ) if litellm.cache is not None: verbose_proxy_logger.debug( f"{blue_color_code}Set Cache on LiteLLM Proxy{reset_color_code}" @@ -3551,21 +3592,23 @@ class ProxyConfig: verbose_proxy_logger.critical( "LITELLM_MASTER_KEY is not set! All requests will be treated as INTERNAL_USER with no admin access. Set LITELLM_MASTER_KEY for production use." ) - ### USER API KEY CACHE IN-MEMORY TTL ### + ### USER API KEY CACHE TTL (in-memory + Redis when Redis auth sharing is enabled) ### user_api_key_cache_ttl = general_settings.get( "user_api_key_cache_ttl", None ) if user_api_key_cache_ttl is not None: + ttl = float(user_api_key_cache_ttl) + # Mirror TTL on Redis as well when ``litellm_settings.enable_redis_auth_cache`` + # attaches Redis to ``user_api_key_cache``; otherwise DualCache misses in + # memory fall back to a key that outlasts ``user_api_key_cache_ttl``. user_api_key_cache.update_cache_ttl( - default_in_memory_ttl=float(user_api_key_cache_ttl), - default_redis_ttl=None, # user_api_key_cache uses in-memory TTL only; Redis not configured for key lookups + default_in_memory_ttl=ttl, + default_redis_ttl=ttl, ) ### PKCE MULTI-INSTANCE PREREQUISITE CHECK ### # PKCE verifiers are stored in redis_usage_cache when available so they can # be read back by any instance (not just the one that started the auth flow). - # user_api_key_cache is intentionally left in-memory-only to avoid routing - # all API-key lookups through Redis. use_pkce = os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true" if use_pkce and redis_usage_cache is None: global _pkce_no_redis_warning_emitted @@ -6294,7 +6337,7 @@ class ProxyStartupEvent: cls, general_settings: dict, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, ): """Initialize JWT auth on startup""" if general_settings.get("litellm_jwtauth", None) is not None: @@ -6343,7 +6386,7 @@ class ProxyStartupEvent: async def _warm_global_spend_cache( cls, litellm_proxy_admin_name: str, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, prisma_client: PrismaClient, ) -> None: """Warm global spend cache once at startup to reduce impact of first wave of requests.""" @@ -6983,7 +7026,7 @@ class ProxyStartupEvent: cls, database_url: Optional[str], proxy_logging_obj: ProxyLogging, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, ) -> Optional[PrismaClient]: """ - Sets up prisma client @@ -10285,6 +10328,101 @@ def _paginate_models_response( } +def _team_models_resolve_to_names( + team_models: List[str], access_groups: Dict[str, Any] +) -> List[str]: + """Expand team model entries (including access group names) to concrete model names.""" + resolved: List[str] = [] + for name in team_models: + if name in access_groups: + resolved.extend(access_groups[name]) + else: + resolved.append(name) + return resolved + + +async def _load_team_object_for_model_filter( + team_id: str, prisma_client: PrismaClient +) -> Optional[LiteLLM_TeamTable]: + """Load team row from DB; returns None if missing or on error.""" + try: + team_db_object = await prisma_client.db.litellm_teamtable.find_unique( + where={"team_id": team_id} + ) + if team_db_object is None: + verbose_proxy_logger.warning(f"Team {team_id} not found in database") + return None + return LiteLLM_TeamTable(**team_db_object.model_dump()) + except Exception as e: + verbose_proxy_logger.exception(f"Error fetching team {team_id}: {str(e)}") + return None + + +async def _gather_team_accessible_model_ids( + team_object: LiteLLM_TeamTable, + team_id: str, + prisma_client: PrismaClient, + llm_router: Router, +) -> Set[str]: + """Collect model IDs the team can use from router config and DB.""" + team_accessible_model_ids: Set[str] = set() + access_groups = llm_router.get_model_access_groups() if llm_router else {} + + if ( + not team_object.models + or SpecialModelNames.all_proxy_models.value in team_object.models + ): + model_list = llm_router.get_model_list() if llm_router else [] + if model_list is not None: + for model in model_list: + model_id = model.get("model_info", {}).get("id", None) + if model_id is None: + continue + team_model_id = model.get("model_info", {}).get("team_id", None) + if team_model_id is None or team_model_id == team_id: + team_accessible_model_ids.add(model_id) + else: + resolved_model_names: Set[str] = set() + for model_name in team_object.models: + if model_name in access_groups: + resolved_model_names.update(access_groups[model_name]) + else: + resolved_model_names.add(model_name) + + for model_name in resolved_model_names: + _models = ( + llm_router.get_model_list(model_name=model_name, team_id=team_id) + if llm_router + else [] + ) + if _models is not None: + for model in _models: + model_id = model.get("model_info", {}).get("id", None) + if model_id is not None: + team_accessible_model_ids.add(model_id) + + try: + if ( + team_object.models + and SpecialModelNames.all_proxy_models.value not in team_object.models + ): + _resolved_names = _team_models_resolve_to_names( + team_object.models, access_groups + ) + db_models = await prisma_client.db.litellm_proxymodeltable.find_many( + where={"model_name": {"in": _resolved_names}} + ) + for db_model in db_models: + if db_model.model_id: + team_accessible_model_ids.add(db_model.model_id) + except Exception as e: + verbose_proxy_logger.debug( + f"Error querying database models for team {team_id}: {str(e)}" + ) + + return team_accessible_model_ids + + async def _filter_models_by_team_id( all_models: List[Dict[str, Any]], team_id: str, @@ -10307,78 +10445,13 @@ async def _filter_models_by_team_id( Returns: Filtered list of models """ - # Get team from database - try: - team_db_object = await prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": team_id} - ) - if team_db_object is None: - verbose_proxy_logger.warning(f"Team {team_id} not found in database") - # If team doesn't exist, return empty list - return [] - - team_object = LiteLLM_TeamTable(**team_db_object.model_dump()) - except Exception as e: - verbose_proxy_logger.exception(f"Error fetching team {team_id}: {str(e)}") + team_object = await _load_team_object_for_model_filter(team_id, prisma_client) + if team_object is None: return [] - # Get models accessible to this team (similar to _add_team_models_to_all_models) - team_accessible_model_ids: Set[str] = set() - - if ( - not team_object.models # empty list = all model access - or SpecialModelNames.all_proxy_models.value in team_object.models - ): - # Team has access to all models - model_list = llm_router.get_model_list() if llm_router else [] - if model_list is not None: - for model in model_list: - model_id = model.get("model_info", {}).get("id", None) - if model_id is None: - continue - # if team model id set, check if team id matches - team_model_id = model.get("model_info", {}).get("team_id", None) - can_add_model = False - if team_model_id is None: - can_add_model = True - elif team_model_id == team_id: - can_add_model = True - - if can_add_model: - team_accessible_model_ids.add(model_id) - else: - # Team has access to specific models - for model_name in team_object.models: - _models = ( - llm_router.get_model_list(model_name=model_name, team_id=team_id) - if llm_router - else [] - ) - if _models is not None: - for model in _models: - model_id = model.get("model_info", {}).get("id", None) - if model_id is not None: - team_accessible_model_ids.add(model_id) - - # Also search database for models accessible to this team - # This complements the config search done above - try: - if ( - team_object.models - and SpecialModelNames.all_proxy_models.value not in team_object.models - ): - # Team has specific models - check database for those model names - db_models = await prisma_client.db.litellm_proxymodeltable.find_many( - where={"model_name": {"in": team_object.models}} - ) - for db_model in db_models: - model_id = db_model.model_id - if model_id: - team_accessible_model_ids.add(model_id) - except Exception as e: - verbose_proxy_logger.debug( - f"Error querying database models for team {team_id}: {str(e)}" - ) + team_accessible_model_ids = await _gather_team_accessible_model_ids( + team_object, team_id, prisma_client, llm_router + ) # Filter models based on direct_access or access_via_team_ids # Models are already enriched with these fields before this function is called @@ -12818,9 +12891,12 @@ async def update_config( # noqa: PLR0915 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ - For Admin UI - allows admin to update config via UI + For Admin UI - allows admin to update config via UI. - Currently supports modifying General Settings + LiteLLM settings + Writes only the sections present in the request body to LiteLLM_Config rows + (one row per top-level section). Sections the caller did not send are left + untouched — this endpoint never persists pre-existing YAML values to DB as + a side effect of an unrelated update. """ global llm_router, llm_model_list, general_settings, proxy_config, proxy_logging_obj, master_key, prisma_client try: @@ -12828,109 +12904,96 @@ async def update_config( # noqa: PLR0915 raise HTTPException( status_code=403, detail="Only proxy admins can update config" ) - import base64 - """ - - Update the ConfigTable DB - - Run 'add_deployment' - """ if prisma_client is None: raise Exception("No DB Connected") - if store_model_in_db is not True: - raise HTTPException( - status_code=500, - detail={ - "error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature." + async def _read_section(param_name: str) -> dict: + row = await prisma_client.db.litellm_config.find_first( + where={"param_name": param_name} + ) + if row is None or row.param_value is None: + return {} + return dict(row.param_value) + + async def _upsert_section(param_name: str, value: dict) -> None: + serialized = json.dumps(value) + await prisma_client.db.litellm_config.upsert( + where={"param_name": param_name}, + data={ + "create": {"param_name": param_name, "param_value": serialized}, + "update": {"param_value": serialized}, }, ) + # invalidate the DualCache entry so the next reader (this process + # or any other proxy in the cluster) goes to DB. + await invalidate_config_param(param_name) - updated_settings = config_info.json(exclude_none=True) - updated_settings = prisma_client.jsonify_object(updated_settings) - for k, v in updated_settings.items(): - if k == "router_settings": - await prisma_client.db.litellm_config.upsert( - where={"param_name": k}, - data={ - "create": {"param_name": k, "param_value": v}, - "update": {"param_value": v}, - }, - ) - await invalidate_config_param(k) - - ### OLD LOGIC [TODO] MOVE TO DB ### - - # Load existing config - config = await proxy_config.get_config() - verbose_proxy_logger.debug("Loaded config: %s", config) - - # update the general settings + # general_settings: merge per-key, with the alert_to_webhook_url side + # effect of auto-enabling slack alerting. if config_info.general_settings is not None: - config.setdefault("general_settings", {}) - updated_general_settings = config_info.general_settings.dict( - exclude_none=True - ) - - _existing_settings = config["general_settings"] - for k, v in updated_general_settings.items(): - # overwrite existing settings with updated values + existing = await _read_section("general_settings") + updates = config_info.general_settings.dict(exclude_none=True) + for k, v in updates.items(): if k == "alert_to_webhook_url": - # check if slack is already enabled. if not, enable it - if "alerting" not in _existing_settings: - _existing_settings = {"alerting": ["slack"]} - elif isinstance(_existing_settings["alerting"], list): - if "slack" not in _existing_settings["alerting"]: - _existing_settings["alerting"].append("slack") - _existing_settings[k] = v - config["general_settings"] = _existing_settings + if "alerting" not in existing: + existing["alerting"] = ["slack"] + elif ( + isinstance(existing["alerting"], list) + and "slack" not in existing["alerting"] + ): + existing["alerting"].append("slack") + existing[k] = v + await _upsert_section("general_settings", existing) + # environment_variables: encrypt request values, then merge into existing. if config_info.environment_variables is not None: - config.setdefault("environment_variables", {}) - _updated_environment_variables = config_info.environment_variables + existing = await _read_section("environment_variables") + for k, v in config_info.environment_variables.items(): + existing[k] = encrypt_value_helper(value=v) + await _upsert_section("environment_variables", existing) - # encrypt updated_environment_variables # - for k, v in _updated_environment_variables.items(): - encrypted_value = encrypt_value_helper(value=v) - _updated_environment_variables[k] = encrypted_value - - _existing_env_variables = config["environment_variables"] - - for k, v in _updated_environment_variables.items(): - # overwrite existing env variables with updated values - _existing_env_variables[k] = _updated_environment_variables[k] - - # update the litellm settings + # litellm_settings: merge existing + request, request wins (matching + # router_settings semantics — the caller's value for any given key is + # what gets persisted). success_callback is special-cased: it is + # always normalized + deduped, and unioned with any existing list, + # because callbacks are additive (callers send the new entry, not + # the full set). Normalizing on every write — not only when an + # existing entry is present — keeps the DB free of mixed-case + # entries that delete_callback (lowercase lookup) cannot find. if config_info.litellm_settings is not None: - config.setdefault("litellm_settings", {}) - updated_litellm_settings = config_info.litellm_settings - config["litellm_settings"] = { - **updated_litellm_settings, - **config["litellm_settings"], - } + existing = await _read_section("litellm_settings") + updated_litellm_settings = dict(config_info.litellm_settings) - # if litellm.success_callback in updated_litellm_settings and config["litellm_settings"] - if ( - "success_callback" in updated_litellm_settings - and "success_callback" in config["litellm_settings"] - ): - # check both success callback are lists - if isinstance( - config["litellm_settings"]["success_callback"], list - ) and isinstance(updated_litellm_settings["success_callback"], list): - updated_success_callbacks_normalized = normalize_callback_names( - updated_litellm_settings["success_callback"] - ) - combined_success_callback = ( - config["litellm_settings"]["success_callback"] - + updated_success_callbacks_normalized - ) - combined_success_callback = list(set(combined_success_callback)) - config["litellm_settings"][ - "success_callback" - ] = combined_success_callback + incoming_cb = updated_litellm_settings.get("success_callback") + if isinstance(incoming_cb, list): + updated_litellm_settings["success_callback"] = normalize_callback_names( + incoming_cb + ) - # Save the updated config - await proxy_config.save_config(new_config=config) + merged = {**existing, **updated_litellm_settings} + + incoming_cb = updated_litellm_settings.get("success_callback") + existing_cb = existing.get("success_callback") + if isinstance(incoming_cb, list): + if isinstance(existing_cb, list): + # Normalize the existing list too — a row written by a + # different code path may still hold mixed-case names, + # which would otherwise dedup-miss against the lowercase + # incoming entries. + merged["success_callback"] = list( + set(normalize_callback_names(existing_cb) + incoming_cb) + ) + else: + merged["success_callback"] = list(set(incoming_cb)) + + await _upsert_section("litellm_settings", merged) + + # router_settings: merge existing + request, request wins. + if config_info.router_settings is not None: + existing = await _read_section("router_settings") + updates = config_info.router_settings.dict(exclude_none=True) + await _upsert_section("router_settings", {**existing, **updates}) await proxy_config.add_deployment( prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 31dab435b65..69cd7b983e7 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -101,6 +101,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.route_checks import RouteChecks +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.db.create_views import ( create_missing_views, should_create_missing_views, @@ -340,7 +341,7 @@ class ProxyLogging: def __init__( self, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, premium_user: bool = False, ): ## INITIALIZE LITELLM CALLBACKS ## @@ -5718,7 +5719,7 @@ async def get_available_models_for_user( include_model_access_groups: bool = False, only_model_access_groups: bool = False, return_wildcard_routes: bool = False, - user_api_key_cache: Optional["DualCache"] = None, + user_api_key_cache: Optional["UserApiKeyCache"] = None, ) -> List[str]: """ Get the list of models available to a user based on their API key and team permissions. diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 145ec3a641a..da8da1b486f 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -1,9 +1,12 @@ +from __future__ import annotations + import asyncio import json import time import traceback from datetime import datetime -from typing import Any, Dict, List, Optional +from functools import lru_cache +from typing import Any, Dict, List, Literal, Optional import httpx @@ -22,19 +25,26 @@ from litellm.litellm_core_utils.llm_response_utils.response_metadata import ( from litellm.litellm_core_utils.thread_pool_executor import executor from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig from litellm.responses.utils import ResponsesAPIRequestUtils -from litellm.types.llms.openai import ( - OutputTextDeltaEvent, - ResponseAPIUsage, - ResponseCompletedEvent, - ResponsesAPIRequestParams, - ResponsesAPIResponse, - ResponsesAPIStreamEvents, - ResponsesAPIStreamingResponse, -) +from litellm.types.llms.openai import ResponsesAPIStreamEvents from litellm.types.utils import CallTypes from litellm.utils import CustomStreamWrapper, async_post_call_success_deployment_hook +@lru_cache(maxsize=1) +def _get_openai_response_types(): + from litellm.types.llms import openai as openai_types + + return openai_types + + +def _log_background_task_failure(task: "asyncio.Task[Any]", *, task_name: str) -> None: + if task.cancelled(): + return + exception = task.exception() + if exception is not None: + verbose_logger.error("%s failed: %s", task_name, exception) + + class BaseResponsesAPIStreamingIterator: """ Base class for streaming iterators that process responses from the Responses API. @@ -46,7 +56,7 @@ class BaseResponsesAPIStreamingIterator: self, response: httpx.Response, model: str, - responses_api_provider_config: BaseResponsesAPIConfig, + responses_api_provider_config: Optional[BaseResponsesAPIConfig], logging_obj: LiteLLMLoggingObj, litellm_metadata: Optional[Dict[str, Any]] = None, custom_llm_provider: Optional[str] = None, @@ -58,9 +68,13 @@ class BaseResponsesAPIStreamingIterator: self.logging_obj = logging_obj self.finished = False self.responses_api_provider_config = responses_api_provider_config - self.completed_response: Optional[ResponsesAPIStreamingResponse] = None + self.completed_response: Optional[Any] = None self.start_time = getattr(logging_obj, "start_time", datetime.now()) self._failure_handled = False # Track if failure handler has been called + self._completed_response_cached = False + self._completed_response_logged = False + self._completed_response_cache_hit: Optional[bool] = None + self._persist_completed_response_before_logging = True self._stream_created_time: float = time.time() # track request context for hooks @@ -101,7 +115,7 @@ class BaseResponsesAPIStreamingIterator: llm_provider=self.custom_llm_provider or "", ) - def _process_chunk(self, chunk) -> Optional[ResponsesAPIStreamingResponse]: + def _process_chunk(self, chunk) -> Optional[Any]: """Process a single chunk of data from the stream""" if not chunk: return None @@ -122,6 +136,10 @@ class BaseResponsesAPIStreamingIterator: # Format as ResponsesAPIStreamingResponse if isinstance(parsed_chunk, dict): + if self.responses_api_provider_config is None: + raise ValueError( + "responses_api_provider_config is required to process live streaming chunks" + ) openai_responses_api_chunk = ( self.responses_api_provider_config.transform_streaming_response( model=self.model, @@ -195,10 +213,11 @@ class BaseResponsesAPIStreamingIterator: if self.litellm_metadata and self.litellm_metadata.get( "encrypted_content_affinity_enabled" ): + openai_types = _get_openai_response_types() event_type = getattr(openai_responses_api_chunk, "type", None) if event_type in ( - ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, - ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, + openai_types.ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, + openai_types.ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, ): item = getattr(openai_responses_api_chunk, "item", None) if item: @@ -219,10 +238,11 @@ class BaseResponsesAPIStreamingIterator: # Store the completed response (also for incomplete/failed so logging still fires) _chunk_type = getattr(openai_responses_api_chunk, "type", None) + openai_types = _get_openai_response_types() if openai_responses_api_chunk and _chunk_type in ( - ResponsesAPIStreamEvents.RESPONSE_COMPLETED, - ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, - ResponsesAPIStreamEvents.RESPONSE_FAILED, + openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, + openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED, ): self.completed_response = openai_responses_api_chunk # Add cost to usage object if include_cost_in_streaming_usage is True @@ -230,11 +250,11 @@ class BaseResponsesAPIStreamingIterator: litellm.include_cost_in_streaming_usage and self.logging_obj is not None ): - response_obj: Optional[ResponsesAPIResponse] = getattr( + response_obj: Optional[Any] = getattr( openai_responses_api_chunk, "response", None ) if response_obj: - usage_obj: Optional[ResponseAPIUsage] = getattr( + usage_obj: Optional[Any] = getattr( response_obj, "usage", None ) if usage_obj is not None: @@ -247,9 +267,13 @@ class BaseResponsesAPIStreamingIterator: if cost is not None: setattr(usage_obj, "cost", cost) except Exception: + # Best-effort usage cost annotation should not break stream replay. pass - if _chunk_type == ResponsesAPIStreamEvents.RESPONSE_FAILED: + if ( + _chunk_type + == openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED + ): self._handle_logging_failed_response() else: self._handle_logging_completed_response() @@ -266,6 +290,59 @@ class BaseResponsesAPIStreamingIterator: self._handle_failure(e) raise + def _log_completed_response(self, *, is_async: bool) -> None: + if self._completed_response_logged: + return + self._completed_response_logged = True + + if self._persist_completed_response_before_logging: + self._persist_completed_response_to_cache(is_async=is_async) + + # Create a copy for logging to avoid modifying the response object that will be returned to the user + # The logging handlers may transform usage from Responses API format (input_tokens/output_tokens) + # to chat completion format (prompt_tokens/completion_tokens) for internal logging + # Use model_dump + model_validate instead of deepcopy to avoid pickle errors with + # Pydantic ValidatorIterator when response contains tool_choice with allowed_tools (fixes #17192) + logging_response = self.completed_response + if self.completed_response is not None and hasattr( + self.completed_response, "model_dump" + ): + try: + logging_response = type(self.completed_response).model_validate( + self.completed_response.model_dump() + ) + except Exception: + # Fallback to original if serialization fails + pass + + end_time = datetime.now() + if is_async: + asyncio.create_task( + self.logging_obj.async_success_handler( + result=logging_response, + start_time=self.start_time, + end_time=end_time, + cache_hit=self._completed_response_cache_hit, + ) + ) + else: + run_async_function( + async_function=self.logging_obj.async_success_handler, + result=logging_response, + start_time=self.start_time, + end_time=end_time, + cache_hit=self._completed_response_cache_hit, + ) + + executor.submit( + self.logging_obj.success_handler, + result=logging_response, + cache_hit=self._completed_response_cache_hit, + start_time=self.start_time, + end_time=end_time, + ) + self._run_post_success_hooks(end_time=end_time) + def _handle_logging_completed_response(self): """Base implementation - should be overridden by subclasses""" pass @@ -296,6 +373,88 @@ class BaseResponsesAPIStreamingIterator: ) self._handle_failure(exception) + def _get_completed_response_object(self) -> Optional[Any]: + openai_types = _get_openai_response_types() + completed_response = self.completed_response + if isinstance(completed_response, openai_types.ResponsesAPIResponse): + return completed_response + + response_obj = getattr(completed_response, "response", None) + if isinstance(response_obj, openai_types.ResponsesAPIResponse): + return response_obj + + return None + + def _persist_completed_response_to_cache(self, *, is_async: bool) -> None: + if self._completed_response_cached: + return + + completed_response = self.completed_response + openai_types = _get_openai_response_types() + if ( + getattr(completed_response, "type", None) + != openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ): + return + + response_obj = self._get_completed_response_object() + if response_obj is None: + return + + caching_handler = getattr(self.logging_obj, "_llm_caching_handler", None) + if caching_handler is None: + return + + request_kwargs = getattr(caching_handler, "request_kwargs", None) + if ( + not isinstance(request_kwargs, dict) + or request_kwargs.get("stream") is not True + ): + return + request_kwargs = request_kwargs.copy() + preset_cache_key = getattr(caching_handler, "preset_cache_key", None) + request_cache_key = request_kwargs.pop("cache_key", None) + if preset_cache_key is None: + preset_cache_key = request_cache_key + if request_kwargs.get("metadata") is None: + request_kwargs.pop("metadata", None) + request_kwargs.pop("custom_llm_provider", None) + if preset_cache_key is not None: + request_kwargs["cache_key"] = preset_cache_key + + if not caching_handler._should_store_result_in_cache( + original_function=caching_handler.original_function, + kwargs=request_kwargs, + ): + return + + if litellm.cache is None: + return + + cached_response = response_obj.model_dump_json() + if is_async: + cache_write_task = asyncio.create_task( + litellm.cache.async_add_cache( + cached_response, + dynamic_cache_object=getattr(caching_handler, "dual_cache", None), + **request_kwargs, + ) + ) + cache_write_task.add_done_callback( + lambda task: _log_background_task_failure( + task, + task_name="Responses stream cache write", + ) + ) + else: + litellm.cache.add_cache( + cached_response, + dynamic_cache_object=getattr(caching_handler, "dual_cache", None), + **request_kwargs, + ) + + self._completed_response_cached = True + async def _call_post_streaming_deployment_hook(self, chunk): """ Allow callbacks to modify streaming chunks before returning (parity with chat). @@ -480,7 +639,7 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): def __aiter__(self): return self - async def __anext__(self) -> ResponsesAPIStreamingResponse: + async def __anext__(self) -> Any: try: self._check_max_streaming_duration() while True: @@ -520,40 +679,7 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): def _handle_logging_completed_response(self): """Handle logging for completed responses in async context""" - # Create a copy for logging to avoid modifying the response object that will be returned to the user - # The logging handlers may transform usage from Responses API format (input_tokens/output_tokens) - # to chat completion format (prompt_tokens/completion_tokens) for internal logging - # Use model_dump + model_validate instead of deepcopy to avoid pickle errors with - # Pydantic ValidatorIterator when response contains tool_choice with allowed_tools (fixes #17192) - logging_response = self.completed_response - if self.completed_response is not None and hasattr( - self.completed_response, "model_dump" - ): - try: - logging_response = type(self.completed_response).model_validate( - self.completed_response.model_dump() - ) - except Exception: - # Fallback to original if serialization fails - pass - - asyncio.create_task( - self.logging_obj.async_success_handler( - result=logging_response, - start_time=self.start_time, - end_time=datetime.now(), - cache_hit=None, - ) - ) - - executor.submit( - self.logging_obj.success_handler, - result=logging_response, - cache_hit=None, - start_time=self.start_time, - end_time=datetime.now(), - ) - self._run_post_success_hooks(end_time=datetime.now()) + self._log_completed_response(is_async=True) class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): @@ -627,39 +753,7 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): def _handle_logging_completed_response(self): """Handle logging for completed responses in sync context""" - # Create a copy for logging to avoid modifying the response object that will be returned to the user - # The logging handlers may transform usage from Responses API format (input_tokens/output_tokens) - # to chat completion format (prompt_tokens/completion_tokens) for internal logging - # Use model_dump + model_validate instead of deepcopy to avoid pickle errors with - # Pydantic ValidatorIterator when response contains tool_choice with allowed_tools (fixes #17192) - logging_response = self.completed_response - if self.completed_response is not None and hasattr( - self.completed_response, "model_dump" - ): - try: - logging_response = type(self.completed_response).model_validate( - self.completed_response.model_dump() - ) - except Exception: - # Fallback to original if serialization fails - pass - - run_async_function( - async_function=self.logging_obj.async_success_handler, - result=logging_response, - start_time=self.start_time, - end_time=datetime.now(), - cache_hit=None, - ) - - executor.submit( - self.logging_obj.success_handler, - result=logging_response, - cache_hit=None, - start_time=self.start_time, - end_time=datetime.now(), - ) - self._run_post_success_hooks(end_time=datetime.now()) + self._log_completed_response(is_async=False) class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): @@ -683,90 +777,441 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): request_data: Optional[Dict[str, Any]] = None, call_type: Optional[str] = None, ): - super().__init__( - response=response, + transformed = responses_api_provider_config.transform_response_api_response( model=model, - responses_api_provider_config=responses_api_provider_config, + raw_response=response, + logging_obj=logging_obj, + ) + super().__init__( + response=httpx.Response(200), + model=model, + responses_api_provider_config=None, logging_obj=logging_obj, litellm_metadata=litellm_metadata, custom_llm_provider=custom_llm_provider, request_data=request_data, call_type=call_type, ) + self._set_events_from_response(transformed=transformed, logging_obj=logging_obj) - # one-time transform - transformed = ( - self.responses_api_provider_config.transform_response_api_response( - model=self.model, - raw_response=response, - logging_obj=logging_obj, - ) + def _set_events_from_response( + self, + transformed: Any, + logging_obj: LiteLLMLoggingObj, + ) -> None: + self._events = _build_synthetic_response_events( + transformed=transformed, + logging_obj=logging_obj, + chunk_size=self.CHUNK_SIZE, ) - full_text = self._collect_text(transformed) - - # build a list of 5‑char delta events - deltas = [ - OutputTextDeltaEvent( - type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, - delta=full_text[i : i + self.CHUNK_SIZE], - item_id=transformed.id, - output_index=0, - content_index=0, - ) - for i in range(0, len(full_text), self.CHUNK_SIZE) - ] - - # Add cost to usage object if include_cost_in_streaming_usage is True - if litellm.include_cost_in_streaming_usage and logging_obj is not None: - usage_obj: Optional[ResponseAPIUsage] = getattr(transformed, "usage", None) - if usage_obj is not None: - try: - cost: Optional[float] = logging_obj._response_cost_calculator( - result=transformed - ) - if cost is not None: - setattr(usage_obj, "cost", cost) - except Exception: - # If cost calculation fails, continue without cost - pass - - # append the completed event - self._events = deltas + [ - ResponseCompletedEvent( - type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, - response=transformed, - ) - ] self._idx = 0 + self.completed_response = self._events[-1] def __aiter__(self): return self - async def __anext__(self) -> ResponsesAPIStreamingResponse: + async def __anext__(self) -> Any: if self._idx >= len(self._events): raise StopAsyncIteration evt = self._events[self._idx] self._idx += 1 + openai_types = _get_openai_response_types() + if ( + getattr(evt, "type", None) + == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ): + self.completed_response = evt + self._log_completed_response(is_async=True) return evt def __iter__(self): return self - def __next__(self) -> ResponsesAPIStreamingResponse: + def __next__(self) -> Any: if self._idx >= len(self._events): raise StopIteration evt = self._events[self._idx] self._idx += 1 + openai_types = _get_openai_response_types() + if ( + getattr(evt, "type", None) + == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ): + self.completed_response = evt + self._log_completed_response(is_async=False) return evt - def _collect_text(self, resp: ResponsesAPIResponse) -> str: - out = "" - for out_item in resp.output: - item_type = getattr(out_item, "type", None) - if item_type == "message": - for c in getattr(out_item, "content", []): - out += c.text - return out + +class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): + def __init__( + self, + response: Any, + logging_obj: LiteLLMLoggingObj, + request_data: Optional[Dict[str, Any]] = None, + call_type: Optional[str] = None, + ): + BaseResponsesAPIStreamingIterator.__init__( + self, + response=httpx.Response(200), + model=getattr(response, "model", ""), + responses_api_provider_config=None, + logging_obj=logging_obj, + litellm_metadata=None, + custom_llm_provider="cached_response", + request_data=request_data, + call_type=call_type, + ) + self._completed_response_cache_hit = True + self._persist_completed_response_before_logging = False + self._events: List[Any] = [] + self._idx = 0 + self._set_events_from_response(transformed=response, logging_obj=logging_obj) + + def _set_events_from_response( + self, + transformed: Any, + logging_obj: LiteLLMLoggingObj, + ) -> None: + self._events = _build_synthetic_response_events( + transformed=transformed, + logging_obj=logging_obj, + chunk_size=MockResponsesAPIStreamingIterator.CHUNK_SIZE, + ) + self._idx = 0 + self.completed_response = self._events[-1] + + def __aiter__(self): + return self + + async def __anext__(self) -> Any: + if self._idx >= len(self._events): + raise StopAsyncIteration + evt = self._events[self._idx] + self._idx += 1 + openai_types = _get_openai_response_types() + if ( + getattr(evt, "type", None) + == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ): + self.completed_response = evt + self._log_completed_response(is_async=True) + return evt + + def __iter__(self): + return self + + def __next__(self) -> Any: + if self._idx >= len(self._events): + raise StopIteration + evt = self._events[self._idx] + self._idx += 1 + openai_types = _get_openai_response_types() + if ( + getattr(evt, "type", None) + == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ): + self.completed_response = evt + self._log_completed_response(is_async=False) + return evt + + +def _dump_response_object(obj: Any) -> Dict[str, Any]: + if hasattr(obj, "model_dump"): + return obj.model_dump() + if isinstance(obj, dict): + return obj + return {} + + +def _build_response_status_event( + event_type: Literal[ + "response.created", + "response.in_progress", + ], + transformed: Any, +) -> Any: + openai_types = _get_openai_response_types() + in_progress_response = transformed.model_copy( + deep=True, + update={"status": "in_progress", "output": []}, + ) + if event_type == openai_types.ResponsesAPIStreamEvents.RESPONSE_CREATED: + return openai_types.ResponseCreatedEvent( + type=event_type, response=in_progress_response + ) + return openai_types.ResponseInProgressEvent( + type=event_type, response=in_progress_response + ) + + +def _build_content_part_done_event( + *, + item_id: str, + output_index: int, + content_index: int, + part_payload: Dict[str, Any], +) -> Optional[Any]: + openai_types = _get_openai_response_types() + part_type = part_payload.get("type") + part: Any + if part_type == "output_text": + annotations = [ + openai_types.BaseLiteLLMOpenAIResponseObject(**annotation) + for annotation in part_payload.get("annotations", []) or [] + ] + part = openai_types.ContentPartDonePartOutputText( + type="output_text", + text=str(part_payload.get("text") or ""), + annotations=annotations, + logprobs=part_payload.get("logprobs"), + ) + elif part_type == "refusal": + part = openai_types.ContentPartDonePartRefusal( + type="refusal", + refusal=str(part_payload.get("refusal") or ""), + ) + elif part_type == "reasoning_text": + part = openai_types.ContentPartDonePartReasoningText( + type="reasoning_text", + reasoning=str(part_payload.get("reasoning") or ""), + ) + else: + return None + + return openai_types.ContentPartDoneEvent( + type=openai_types.ResponsesAPIStreamEvents.CONTENT_PART_DONE, + item_id=item_id, + output_index=output_index, + content_index=content_index, + part=part, + ) + + +def _add_text_like_part_events( + *, + events: List[Any], + item_id: str, + output_index: int, + content_index: int, + part_payload: Dict[str, Any], + chunk_size: int, +) -> None: + openai_types = _get_openai_response_types() + part_type = part_payload.get("type") + if part_type == "output_text": + text = str(part_payload.get("text") or "") + for i in range(0, len(text), chunk_size): + events.append( + openai_types.OutputTextDeltaEvent( + type=openai_types.ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id=item_id, + output_index=output_index, + content_index=content_index, + delta=text[i : i + chunk_size], + ) + ) + for annotation_index, annotation in enumerate( + part_payload.get("annotations", []) or [] + ): + events.append( + openai_types.OutputTextAnnotationAddedEvent( + type=openai_types.ResponsesAPIStreamEvents.OUTPUT_TEXT_ANNOTATION_ADDED, + item_id=item_id, + output_index=output_index, + content_index=content_index, + annotation_index=annotation_index, + annotation=annotation, + ) + ) + events.append( + openai_types.OutputTextDoneEvent( + type=openai_types.ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE, + item_id=item_id, + output_index=output_index, + content_index=content_index, + text=text, + ) + ) + elif part_type == "refusal": + refusal = str(part_payload.get("refusal") or "") + for i in range(0, len(refusal), chunk_size): + events.append( + openai_types.RefusalDeltaEvent( + type=openai_types.ResponsesAPIStreamEvents.REFUSAL_DELTA, + item_id=item_id, + output_index=output_index, + content_index=content_index, + delta=refusal[i : i + chunk_size], + ) + ) + events.append( + openai_types.RefusalDoneEvent( + type=openai_types.ResponsesAPIStreamEvents.REFUSAL_DONE, + item_id=item_id, + output_index=output_index, + content_index=content_index, + refusal=refusal, + ) + ) + + +def _build_synthetic_response_events( + *, + transformed: Any, + logging_obj: LiteLLMLoggingObj, + chunk_size: int, +) -> List[Any]: + openai_types = _get_openai_response_types() + if litellm.include_cost_in_streaming_usage and logging_obj is not None: + usage_obj: Optional[Any] = getattr(transformed, "usage", None) + if usage_obj is not None: + try: + cost: Optional[float] = logging_obj._response_cost_calculator( + result=transformed + ) + if cost is not None: + setattr(usage_obj, "cost", cost) + except Exception: + pass + + events: List[Any] = [ + _build_response_status_event( + openai_types.ResponsesAPIStreamEvents.RESPONSE_CREATED, transformed + ), + _build_response_status_event( + openai_types.ResponsesAPIStreamEvents.RESPONSE_IN_PROGRESS, transformed + ), + ] + + sequence_number = 0 + for output_index, output_item in enumerate( + getattr(transformed, "output", []) or [] + ): + output_item_payload = _dump_response_object(output_item) + item_id = str(output_item_payload.get("id") or transformed.id) + item_type = output_item_payload.get("type") + + events.append( + openai_types.OutputItemAddedEvent( + type=openai_types.ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, + output_index=output_index, + item=openai_types.BaseLiteLLMOpenAIResponseObject( + **output_item_payload + ), + ) + ) + + if item_type == "message": + for content_index, part in enumerate( + output_item_payload.get("content", []) or [] + ): + part_payload = _dump_response_object(part) + events.append( + openai_types.ContentPartAddedEvent( + type=openai_types.ResponsesAPIStreamEvents.CONTENT_PART_ADDED, + item_id=item_id, + output_index=output_index, + content_index=content_index, + part=openai_types.BaseLiteLLMOpenAIResponseObject( + **part_payload + ), + ) + ) + _add_text_like_part_events( + events=events, + item_id=item_id, + output_index=output_index, + content_index=content_index, + part_payload=part_payload, + chunk_size=chunk_size, + ) + done_event = _build_content_part_done_event( + item_id=item_id, + output_index=output_index, + content_index=content_index, + part_payload=part_payload, + ) + if done_event is not None: + events.append(done_event) + elif item_type == "function_call": + arguments = str(output_item_payload.get("arguments") or "") + for i in range(0, len(arguments), chunk_size): + events.append( + openai_types.FunctionCallArgumentsDeltaEvent( + type=openai_types.ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA, + item_id=item_id, + output_index=output_index, + delta=arguments[i : i + chunk_size], + ) + ) + events.append( + openai_types.FunctionCallArgumentsDoneEvent( + type=openai_types.ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE, + item_id=item_id, + output_index=output_index, + arguments=arguments, + ) + ) + elif item_type == "reasoning": + for summary_index, summary in enumerate( + output_item_payload.get("summary", []) or [] + ): + summary_payload = _dump_response_object(summary) + summary_text = str(summary_payload.get("text") or "") + for i in range(0, len(summary_text), chunk_size): + events.append( + openai_types.ReasoningSummaryTextDeltaEvent( + type=openai_types.ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DELTA, + item_id=item_id, + output_index=output_index, + summary_index=summary_index, + delta=summary_text[i : i + chunk_size], + ) + ) + sequence_number += 1 + events.append( + openai_types.ReasoningSummaryTextDoneEvent( + type=openai_types.ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DONE, + item_id=item_id, + output_index=output_index, + sequence_number=sequence_number, + summary_index=summary_index, + text=summary_text, + ) + ) + sequence_number += 1 + events.append( + openai_types.ReasoningSummaryPartDoneEvent( + type=openai_types.ResponsesAPIStreamEvents.REASONING_SUMMARY_PART_DONE, + item_id=item_id, + output_index=output_index, + sequence_number=sequence_number, + summary_index=summary_index, + part=openai_types.BaseLiteLLMOpenAIResponseObject( + **summary_payload + ), + ) + ) + + sequence_number += 1 + events.append( + openai_types.OutputItemDoneEvent( + type=openai_types.ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, + output_index=output_index, + sequence_number=sequence_number, + item=openai_types.BaseLiteLLMOpenAIResponseObject( + **output_item_payload + ), + ) + ) + + events.append( + openai_types.ResponseCompletedEvent( + type=openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=transformed, + ) + ) + return events # --------------------------------------------------------------------------- @@ -951,8 +1396,8 @@ class ResponsesWebSocketStreaming: # --------------------------------------------------------------------------- _RESPONSE_CREATE_PARAMS: frozenset = ( - ResponsesAPIRequestParams.__required_keys__ - | ResponsesAPIRequestParams.__optional_keys__ + _get_openai_response_types().ResponsesAPIRequestParams.__required_keys__ + | _get_openai_response_types().ResponsesAPIRequestParams.__optional_keys__ ) _MANAGED_WS_SKIP_KWARGS: frozenset = frozenset( @@ -1085,7 +1530,7 @@ class ManagedResponsesWebSocketHandler: @staticmethod def _extract_output_messages( - completed_event: Dict[str, Any] + completed_event: Dict[str, Any], ) -> List[Dict[str, Any]]: """ Convert the output items in a ``response.completed`` event into diff --git a/litellm/router.py b/litellm/router.py index 7448cdd1b47..50fd7eaed0b 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -5261,11 +5261,34 @@ class Router: """ Initialize the Containers API endpoints on the router. - Container operations don't need model-based routing, so we call the - original function directly with the custom_llm_provider. + LiteLLM-managed container IDs (``cntr_...``) encode ``model_id`` and provider + metadata. When present, decode the ID, replace ``container_id`` with the + upstream value, and route through ``_ageneric_api_call_with_fallbacks`` so + deployment credentials (e.g. regional ``api_base`` for Azure) match + :meth:`_init_responses_api_endpoints`. Otherwise call the handler directly. """ if custom_llm_provider and "custom_llm_provider" not in kwargs: kwargs["custom_llm_provider"] = custom_llm_provider + + from litellm.responses.utils import ResponsesAPIRequestUtils + + container_id = kwargs.get("container_id") + if isinstance(container_id, str): + decoded = ResponsesAPIRequestUtils._decode_container_id(container_id) + original_id = decoded.get("response_id", container_id) + if original_id != container_id: + kwargs["container_id"] = original_id + decoded_provider = decoded.get("custom_llm_provider") + if decoded_provider and kwargs.get("custom_llm_provider") == "openai": + kwargs["custom_llm_provider"] = decoded_provider + model_id = decoded.get("model_id") + if model_id: + kwargs["model"] = model_id + return await self._ageneric_api_call_with_fallbacks( + original_function=original_function, + **kwargs, + ) + return await original_function(**kwargs) async def _init_responses_api_endpoints( diff --git a/litellm/router_utils/common_utils.py b/litellm/router_utils/common_utils.py index bef42e23848..f6da26ccd7f 100644 --- a/litellm/router_utils/common_utils.py +++ b/litellm/router_utils/common_utils.py @@ -23,21 +23,28 @@ def add_model_file_id_mappings( healthy_deployments: Union[List[Dict], Dict], responses: List["OpenAIFileObject"] ) -> dict: """ - Create a mapping of model name to file id + Create a mapping of model id to file id { "model_id": "file_id", "model_id": "file_id", } + + `healthy_deployments` may be either a list of deployment dicts (multiple + matched deployments) or a single deployment dict (when the router resolved + a specific deployment, e.g. because the requested model matched a + `model_info.id`). Both shapes must be handled by extracting + `model_info.id` from each deployment. """ - model_file_id_mapping = {} - if isinstance(healthy_deployments, list): - for deployment, response in zip(healthy_deployments, responses): - model_file_id_mapping[deployment.get("model_info", {}).get("id")] = ( - response.id - ) - elif isinstance(healthy_deployments, dict): - for model_id, file_id in healthy_deployments.items(): - model_file_id_mapping[model_id] = file_id + model_file_id_mapping: Dict[str, str] = {} + deployments_list: List[Dict] = ( + healthy_deployments + if isinstance(healthy_deployments, list) + else [healthy_deployments] + ) + for deployment, response in zip(deployments_list, responses): + model_id = deployment.get("model_info", {}).get("id") + if model_id is not None: + model_file_id_mapping[model_id] = response.id return model_file_id_mapping diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 2fd0c4ea970..986ec39f3bb 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -1482,6 +1482,7 @@ class ReasoningSummaryTextDeltaEvent(BaseLiteLLMOpenAIResponseObject): type: Literal[ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DELTA] item_id: str output_index: int + summary_index: int = 0 delta: str @@ -1490,7 +1491,7 @@ class ReasoningSummaryTextDoneEvent(BaseLiteLLMOpenAIResponseObject): item_id: str output_index: int sequence_number: int - summary_index: int + summary_index: int = 0 text: str @@ -1499,7 +1500,7 @@ class ReasoningSummaryPartDoneEvent(BaseLiteLLMOpenAIResponseObject): item_id: str output_index: int sequence_number: int - summary_index: int + summary_index: int = 0 part: BaseLiteLLMOpenAIResponseObject diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index 2e7d57cef25..87bf11a9026 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -6,6 +6,13 @@ from typing_extensions import ( TypedDict, ) +from litellm.types.llms.openai import EmbeddingInput + +# Gemini supports nested-list inputs (e.g. [["text", "image"]]) as an explicit +# opt-in for combined embeddings — a provider-specific extension of the +# OpenAI-faithful EmbeddingInput shape. +GeminiEmbeddingInput = Union[EmbeddingInput, List[List[str]]] + class FunctionResponse(TypedDict): name: str diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 70bec9c6ea7..c38c14a1f75 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -33483,6 +33483,72 @@ "source": "https://console.cloud.google.com/vertex-ai/publishers/openai/model-garden/gpt-oss-120b-maas", "supports_reasoning": true }, + "vertex_ai/xai/grok-4.1-fast-non-reasoning": { + "cache_read_input_token_cost": 5e-08, + "input_cost_per_token": 2e-07, + "litellm_provider": "vertex_ai", + "max_input_tokens": 2000000, + "max_output_tokens": 2000000, + "max_tokens": 2000000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "source": "https://docs.x.ai/docs/models (Vertex AI Model Garden)", + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "vertex_ai/xai/grok-4.1-fast-reasoning": { + "cache_read_input_token_cost": 5e-08, + "input_cost_per_token": 2e-07, + "litellm_provider": "vertex_ai", + "max_input_tokens": 2000000, + "max_output_tokens": 2000000, + "max_tokens": 2000000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "source": "https://docs.x.ai/docs/models (Vertex AI Model Garden)", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "vertex_ai/xai/grok-4.20-non-reasoning": { + "cache_read_input_token_cost": 2e-07, + "input_cost_per_token": 2e-06, + "litellm_provider": "vertex_ai", + "max_input_tokens": 2000000, + "max_output_tokens": 2000000, + "max_tokens": 2000000, + "mode": "chat", + "output_cost_per_token": 6e-06, + "source": "https://docs.x.ai/docs/models (Vertex AI Model Garden)", + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "vertex_ai/xai/grok-4.20-reasoning": { + "cache_read_input_token_cost": 2e-07, + "input_cost_per_token": 2e-06, + "litellm_provider": "vertex_ai", + "max_input_tokens": 2000000, + "max_output_tokens": 2000000, + "max_tokens": 2000000, + "mode": "chat", + "output_cost_per_token": 6e-06, + "source": "https://docs.x.ai/docs/models (Vertex AI Model Garden)", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, "vertex_ai/qwen/qwen3-235b-a22b-instruct-2507-maas": { "input_cost_per_token": 2.5e-07, "litellm_provider": "vertex_ai-qwen_models", @@ -34920,6 +34986,20 @@ "supports_tool_choice": true, "source": "https://aws.amazon.com/bedrock/pricing/" }, + "zai.glm-5": { + "input_cost_per_token": 1e-06, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 200000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 3.2e-06, + "supports_function_calling": true, + "supports_reasoning": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, "zai.glm-4.7-flash": { "input_cost_per_token": 7e-08, "litellm_provider": "bedrock_converse", diff --git a/pyproject.toml b/pyproject.toml index 0ef0a993dd5..65b9bd2c983 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -149,6 +149,8 @@ dev = [ "parameterized==0.9.0", "openapi-core==0.22.0; python_version < '3.14'", "pytest-timeout==2.4.0", + "vcrpy==8.1.1", + "pytest-recording==0.13.4", ] proxy-dev = [ "prisma==0.11.0", diff --git a/tests/_flush_vcr_cache.py b/tests/_flush_vcr_cache.py new file mode 100644 index 00000000000..d236c88fa3c --- /dev/null +++ b/tests/_flush_vcr_cache.py @@ -0,0 +1,44 @@ +from __future__ import annotations + +import os +import sys + +import redis + +from tests._vcr_redis_persister import CASSETTE_REDIS_URL_ENV, _redis_url_from_env + +PREFIX = "litellm:vcr:cassette:" +SCAN_BATCH = 500 + + +def _client() -> redis.Redis: + url = _redis_url_from_env() + if not url: + sys.exit(f"Set {CASSETTE_REDIS_URL_ENV} to flush the VCR cache") + return redis.Redis.from_url( + url, + socket_timeout=5, + socket_connect_timeout=5, + decode_responses=False, + ) + + +def main() -> None: + client = _client() + deleted = 0 + pipeline = client.pipeline(transaction=False) + pending = 0 + for key in client.scan_iter(match=f"{PREFIX}*", count=SCAN_BATCH): + pipeline.delete(key) + pending += 1 + if pending >= SCAN_BATCH: + deleted += sum(pipeline.execute()) + pipeline = client.pipeline(transaction=False) + pending = 0 + if pending: + deleted += sum(pipeline.execute()) + print(f"Deleted {deleted} VCR cassette key(s) under {PREFIX!r}") + + +if __name__ == "__main__": + main() diff --git a/tests/_vcr_redis_persister.py b/tests/_vcr_redis_persister.py new file mode 100644 index 00000000000..4d72a1142bb --- /dev/null +++ b/tests/_vcr_redis_persister.py @@ -0,0 +1,182 @@ +from __future__ import annotations + +import logging +import os +from typing import Any, Optional + +from vcr.persisters.filesystem import CassetteNotFoundError +from vcr.serialize import deserialize, serialize + +CASSETTE_TTL_SECONDS = 24 * 60 * 60 +REDIS_KEY_PREFIX = "litellm:vcr:cassette:" +CASSETTE_REDIS_URL_ENV = "CASSETTE_REDIS_URL" +VCR_VERBOSE_ENV = "LITELLM_VCR_VERBOSE" +MAX_EPISODES_PER_CASSETTE = 50 + +_log = logging.getLogger(__name__) +_passed_by_cassette_key: dict[str, bool] = {} + + +def mark_test_outcome_for_cassette(cassette_path: str, passed: bool) -> None: + _passed_by_cassette_key[redis_key_for(cassette_path)] = passed + + +def redis_key_for(cassette_path: str) -> str: + rel = os.path.relpath(str(cassette_path)) + if rel.endswith(".yaml"): + rel = rel[: -len(".yaml")] + rel = rel.replace("/cassettes/", "/").lstrip("./") + return f"{REDIS_KEY_PREFIX}{rel}" + + +def _redis_url_from_env() -> Optional[str]: + return os.environ.get(CASSETTE_REDIS_URL_ENV) or None + + +def _build_default_client(): + import redis + from redis.backoff import ExponentialBackoff + from redis.exceptions import ConnectionError as RedisConnectionError + from redis.exceptions import TimeoutError as RedisTimeoutError + from redis.retry import Retry + + url = _redis_url_from_env() + if not url: + raise RuntimeError( + f"Set {CASSETTE_REDIS_URL_ENV} to enable the VCR persister. " + "Cassette Redis is intentionally separate from the application " + "Redis (REDIS_URL/REDIS_HOST) to avoid being flushed by tests." + ) + return redis.Redis.from_url( + url, + socket_timeout=5, + socket_connect_timeout=5, + decode_responses=False, + retry=Retry(ExponentialBackoff(cap=2, base=0.1), retries=2), + retry_on_error=[RedisConnectionError, RedisTimeoutError], + ) + + +def make_redis_persister( + client: Optional[Any] = None, + ttl_seconds: int = CASSETTE_TTL_SECONDS, +): + redis_client = client if client is not None else _build_default_client() + + try: + from redis.exceptions import ConnectionError as RedisConnectionError + from redis.exceptions import TimeoutError as RedisTimeoutError + + _transient_errors: tuple = (RedisConnectionError, RedisTimeoutError) + except ImportError: # pragma: no cover - redis is a hard test dep + _transient_errors = () + + class _RedisPersister: + @staticmethod + def load_cassette(cassette_path, serializer): + try: + data = redis_client.get(redis_key_for(cassette_path)) + except _transient_errors as exc: + _log.warning( + "VCR redis load failed for %s; treating as cache miss: %s", + cassette_path, + exc, + ) + raise CassetteNotFoundError() from exc + if data is None: + raise CassetteNotFoundError() + if isinstance(data, bytes): + data = data.decode("utf-8") + return deserialize(data, serializer) + + @staticmethod + def save_cassette(cassette_path, cassette_dict, serializer): + key = redis_key_for(cassette_path) + passed = _passed_by_cassette_key.pop(key, True) + episode_count = len(cassette_dict.get("requests", []) or []) + if episode_count > MAX_EPISODES_PER_CASSETTE: + _log.warning( + "VCR redis save refused for %s; cassette has %d episodes " + "(> MAX_EPISODES_PER_CASSETTE=%d). The test likely produces " + "non-deterministic request bodies (e.g. uuid) and is " + "appending instead of replaying. Opt it out with the " + "no-vcr list in conftest, or stabilize its request body.", + cassette_path, + episode_count, + MAX_EPISODES_PER_CASSETTE, + ) + return + if not passed: + _log.info( + "VCR redis save skipped for %s; test did not pass — " + "leaving any prior cassette intact", + cassette_path, + ) + return + data = serialize(cassette_dict, serializer) + payload = data.encode("utf-8") if isinstance(data, str) else data + try: + redis_client.set(key, payload, ex=ttl_seconds) + except _transient_errors as exc: + _log.warning( + "VCR redis save failed for %s; cassette not persisted: %s", + cassette_path, + exc, + ) + + return _RedisPersister + + +def filter_non_2xx_response(response): + if not isinstance(response, dict): + return response + status = response.get("status") + code = status.get("code") if isinstance(status, dict) else status + if not isinstance(code, int): + return response + return response if 200 <= code < 300 else None + + +_PATCHED_AIOHTTP_RECORD = False + + +def patch_vcrpy_aiohttp_record_path() -> None: + """Re-feed the response body into aiohttp's StreamReader after vcrpy's + record_response drains it, so downstream consumers (e.g. + LiteLLMAiohttpTransport.AiohttpResponseStream) can still read it.""" + global _PATCHED_AIOHTTP_RECORD + if _PATCHED_AIOHTTP_RECORD: + return + import vcr.stubs.aiohttp_stubs as _aiohttp_stubs + + _orig_record_response = _aiohttp_stubs.record_response + + async def _record_response_preserving_body(cassette, vcr_request, response): + await _orig_record_response(cassette, vcr_request, response) + body = getattr(response, "_body", None) or b"" + if body: + response.content.unread_data(body) + + _aiohttp_stubs.record_response = _record_response_preserving_body + _PATCHED_AIOHTTP_RECORD = True + + +def vcr_verbose_enabled() -> bool: + return os.environ.get(VCR_VERBOSE_ENV) == "1" + + +def format_vcr_verdict(cassette: Any) -> str: + if cassette is None: + return "[VCR NOOP]" + played = getattr(cassette, "play_count", 0) or 0 + dirty = getattr(cassette, "dirty", False) + total = len(cassette) if hasattr(cassette, "__len__") else 0 + if played == 0 and not dirty: + return "[VCR NOOP] (no http traffic)" + if played > 0 and not dirty: + return f"[VCR HIT] {played} replayed, 0 new ({total} cassette entries)" + if played == 0 and dirty: + return f"[VCR MISS] 0 replayed, recorded new ({total} cassette entries)" + return ( + f"[VCR PARTIAL] {played} replayed + new recordings ({total} cassette entries)" + ) diff --git a/tests/code_coverage_tests/liccheck.ini b/tests/code_coverage_tests/liccheck.ini index 4ecc0d0bc98..344e38da835 100644 --- a/tests/code_coverage_tests/liccheck.ini +++ b/tests/code_coverage_tests/liccheck.ini @@ -169,3 +169,4 @@ langchain-mcp-adapters: >=0.2.1 # MIT License langgraph: >=1.0.10 # MIT License langgraph-prebuilt: >=1.0.8 # MIT License - https://github.com/langchain-ai/langgraph/blob/main/LICENSE pytest-rerunfailures: >=15.1 # MPL 2.0 license +pytest-recording: >=0.13.4 # MIT license diff --git a/tests/litellm_utils_tests/test_bedrock_token_counter.py b/tests/litellm_utils_tests/test_bedrock_token_counter.py index 62a595baaaf..9fb2463e8b5 100644 --- a/tests/litellm_utils_tests/test_bedrock_token_counter.py +++ b/tests/litellm_utils_tests/test_bedrock_token_counter.py @@ -130,7 +130,7 @@ class TestBedrockCountTokensEndpoint: ) assert ( url - == "https://bedrock-runtime.us-east-1.amazonaws.com/model/amazon.nova-lite-v1:0/count-tokens" + == "https://bedrock-runtime.us-east-1.amazonaws.com/model/amazon.nova-lite-v1%3A0/count-tokens" ) def test_api_base_overrides_default(self): @@ -141,7 +141,7 @@ class TestBedrockCountTokensEndpoint: aws_region_name="us-east-1", api_base=custom_base, ) - assert url == f"{custom_base}/model/amazon.nova-lite-v1:0/count-tokens" + assert url == f"{custom_base}/model/amazon.nova-lite-v1%3A0/count-tokens" def test_aws_bedrock_runtime_endpoint_overrides_default(self): handler = self._make_handler() @@ -153,7 +153,7 @@ class TestBedrockCountTokensEndpoint: aws_region_name="eu-west-1", aws_bedrock_runtime_endpoint=custom_endpoint, ) - assert url == f"{custom_endpoint}/model/amazon.nova-lite-v1:0/count-tokens" + assert url == f"{custom_endpoint}/model/amazon.nova-lite-v1%3A0/count-tokens" def test_api_base_takes_priority_over_aws_bedrock_runtime_endpoint(self): handler = self._make_handler() @@ -165,7 +165,7 @@ class TestBedrockCountTokensEndpoint: api_base=api_base, aws_bedrock_runtime_endpoint=runtime_endpoint, ) - assert url == f"{api_base}/model/amazon.nova-lite-v1:0/count-tokens" + assert url == f"{api_base}/model/amazon.nova-lite-v1%3A0/count-tokens" def test_env_var_overrides_default(self, monkeypatch): monkeypatch.setenv( diff --git a/tests/llm_responses_api_testing/conftest.py b/tests/llm_responses_api_testing/conftest.py index 0b03348190a..80f36e159a9 100644 --- a/tests/llm_responses_api_testing/conftest.py +++ b/tests/llm_responses_api_testing/conftest.py @@ -1,5 +1,6 @@ # conftest.py +import asyncio import importlib import os import sys @@ -9,9 +10,161 @@ import pytest sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path -import litellm -import asyncio +import litellm # noqa: E402 + +from tests._vcr_redis_persister import ( # noqa: E402 + filter_non_2xx_response, + format_vcr_verdict, + make_redis_persister, + mark_test_outcome_for_cassette, + patch_vcrpy_aiohttp_record_path, + vcr_verbose_enabled, +) + + +_controller_pluginmanager = None +_controller_terminal_reporter = None + + +_FILTERED_REQUEST_HEADERS = ( + "authorization", + "x-api-key", + "anthropic-api-key", + "anthropic-version", + "openai-api-key", + "azure-api-key", + "api-key", + "cookie", + "x-amz-security-token", + "x-amz-date", + "x-amz-content-sha256", + "amz-sdk-invocation-id", + "amz-sdk-request", + "x-goog-api-key", + "x-goog-user-project", +) + +_FILTERED_RESPONSE_HEADERS = ( + "set-cookie", + "x-request-id", + "request-id", + "cf-ray", + "anthropic-organization-id", + "openai-organization", + "x-amzn-requestid", + "x-amzn-trace-id", + "date", +) + + +def _scrub_response(response): + if not isinstance(response, dict): + return response + headers = response.get("headers") or {} + if isinstance(headers, dict): + for header in list(headers): + if header.lower() in _FILTERED_RESPONSE_HEADERS: + headers.pop(header, None) + return response + + +def _before_record_response(response): + return filter_non_2xx_response(_scrub_response(response)) + + +@pytest.fixture(scope="module") +def vcr_config(): + return { + "filter_headers": list(_FILTERED_REQUEST_HEADERS), + "decode_compressed_response": True, + "record_mode": "new_episodes", + "allow_playback_repeats": True, + "match_on": ( + "method", + "scheme", + "host", + "port", + "path", + "query", + "body", + ), + "before_record_response": _before_record_response, + } + + +def _vcr_disabled() -> bool: + if os.environ.get("LITELLM_VCR_DISABLE") == "1": + return True + return not os.environ.get("CASSETTE_REDIS_URL") + + +def pytest_recording_configure(config, vcr): + if _vcr_disabled(): + return + vcr.register_persister(make_redis_persister()) + patch_vcrpy_aiohttp_record_path() + + +@pytest.hookimpl(hookwrapper=True) +def pytest_runtest_makereport(item, call): + outcome = yield + rep = outcome.get_result() + setattr(item, f"rep_{rep.when}", rep) + + +@pytest.fixture(autouse=True) +def _vcr_outcome_gate(request, vcr): + yield + cassette = vcr + rep_call = getattr(request.node, "rep_call", None) + test_passed = bool(rep_call and rep_call.passed) + cassette_path = getattr(cassette, "_path", None) if cassette is not None else None + if cassette_path: + mark_test_outcome_for_cassette(cassette_path, test_passed) + + if not vcr_verbose_enabled(): + return + verdict = format_vcr_verdict(cassette) + request.node.user_properties.append(("vcr_verdict", verdict)) + + +def pytest_configure(config): + global _controller_pluginmanager + if os.environ.get("PYTEST_XDIST_WORKER"): + return + _controller_pluginmanager = config.pluginmanager + + +def _resolve_terminal_reporter(): + global _controller_terminal_reporter + if _controller_terminal_reporter is not None: + return _controller_terminal_reporter + if _controller_pluginmanager is None: + return None + _controller_terminal_reporter = _controller_pluginmanager.getplugin( + "terminalreporter" + ) + return _controller_terminal_reporter + + +def pytest_runtest_logreport(report): + if report.when != "teardown": + return + if os.environ.get("PYTEST_XDIST_WORKER"): + return + if not vcr_verbose_enabled(): + return + reporter = _resolve_terminal_reporter() + if reporter is None: + return + verdict = next( + (v for k, v in (report.user_properties or []) if k == "vcr_verdict"), + None, + ) + if not verdict: + return + reporter.write_line(f"{verdict} :: {report.nodeid}") @pytest.fixture(scope="session") @@ -61,15 +214,18 @@ def setup_and_teardown(): def pytest_collection_modifyitems(config, items): - # Separate tests in 'test_amazing_proxy_custom_logger.py' and other tests + if not _vcr_disabled(): + for item in items: + if item.get_closest_marker("vcr") is not None: + continue + item.add_marker(pytest.mark.vcr) + custom_logger_tests = [ item for item in items if "custom_logger" in item.parent.name ] other_tests = [item for item in items if "custom_logger" not in item.parent.name] - # Sort tests based on their names custom_logger_tests.sort(key=lambda x: x.name) other_tests.sort(key=lambda x: x.name) - # Reorder the items list items[:] = custom_logger_tests + other_tests diff --git a/tests/llm_responses_api_testing/test_responses_hooks.py b/tests/llm_responses_api_testing/test_responses_hooks.py index 3227fecdfb2..3799a0b9121 100644 --- a/tests/llm_responses_api_testing/test_responses_hooks.py +++ b/tests/llm_responses_api_testing/test_responses_hooks.py @@ -1,6 +1,9 @@ import asyncio +from contextlib import suppress from datetime import datetime +import json from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock import httpx import pytest @@ -8,8 +11,17 @@ import pytest import litellm from litellm.integrations.custom_logger import CustomLogger from litellm.responses import streaming_iterator as streaming_module -from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator -from litellm.types.llms.openai import ResponsesAPIStreamEvents +from litellm.responses.streaming_iterator import ( + CachedResponsesAPIStreamingIterator, + MockResponsesAPIStreamingIterator, + ResponsesAPIStreamingIterator, + SyncResponsesAPIStreamingIterator, +) +from litellm.types.llms.openai import ( + ResponseCompletedEvent, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, +) from litellm.types.utils import CallTypes @@ -19,15 +31,19 @@ class _FakeLoggingObj: self.async_success_calls = 0 self.failure_calls = 0 self.async_failure_calls = 0 + self.last_success_kwargs = None + self.last_async_success_kwargs = None self.start_time = datetime.now() self.model_call_details = {"litellm_params": {}} # Signature alignment with Logging handlers def success_handler(self, *args, **kwargs): self.success_calls += 1 + self.last_success_kwargs = kwargs async def async_success_handler(self, *args, **kwargs): self.async_success_calls += 1 + self.last_async_success_kwargs = kwargs def failure_handler(self, *args, **kwargs): self.failure_calls += 1 @@ -36,6 +52,115 @@ class _FakeLoggingObj: self.async_failure_calls += 1 +def _make_completed_response(response_id: str = "resp_test") -> ResponseCompletedEvent: + return ResponseCompletedEvent( + type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=ResponsesAPIResponse( + id=response_id, + created_at=int(datetime.now().timestamp()), + status="completed", + model="test-model", + object="response", + output=[ + { + "type": "message", + "id": f"msg_{response_id}", + "status": "completed", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "cached streamed response", + "annotations": [], + } + ], + } + ], + ), + ) + + +@pytest.mark.asyncio +async def test_log_background_task_failure_logs_task_exceptions(monkeypatch): + error_logger = MagicMock() + monkeypatch.setattr(streaming_module.verbose_logger, "error", error_logger) + + async def _boom(): + raise RuntimeError("boom") + + task = asyncio.create_task(_boom()) + with suppress(RuntimeError): + await task + + streaming_module._log_background_task_failure(task, task_name="cache write") + + error_logger.assert_called_once() + assert error_logger.call_args.args == ( + "%s failed: %s", + "cache write", + task.exception(), + ) + + +@pytest.mark.asyncio +async def test_log_background_task_failure_ignores_cancelled_tasks(monkeypatch): + error_logger = MagicMock() + monkeypatch.setattr(streaming_module.verbose_logger, "error", error_logger) + + task = asyncio.create_task(asyncio.sleep(1)) + task.cancel() + with suppress(asyncio.CancelledError): + await task + + streaming_module._log_background_task_failure(task, task_name="cache write") + + error_logger.assert_not_called() + + +def test_content_part_done_event_supports_refusal_and_reasoning_text(): + refusal_event = streaming_module._build_content_part_done_event( + item_id="msg_1", + output_index=0, + content_index=0, + part_payload={"type": "refusal", "refusal": "no"}, + ) + reasoning_event = streaming_module._build_content_part_done_event( + item_id="msg_1", + output_index=0, + content_index=1, + part_payload={"type": "reasoning_text", "reasoning": "because"}, + ) + unsupported_event = streaming_module._build_content_part_done_event( + item_id="msg_1", + output_index=0, + content_index=2, + part_payload={"type": "image"}, + ) + + assert refusal_event.part.type == "refusal" + assert refusal_event.part.refusal == "no" + assert reasoning_event.part.type == "reasoning_text" + assert reasoning_event.part.reasoning == "because" + assert unsupported_event is None + + +def test_dump_response_object_handles_model_and_unknown_values(): + response = ResponsesAPIResponse( + id="resp_dump", + created_at=int(datetime.now().timestamp()), + status="completed", + model="gpt-4.1-mini", + object="response", + output=[], + ) + + assert streaming_module._dump_response_object(response)["id"] == "resp_dump" + assert streaming_module._dump_response_object({"type": "message"}) == { + "type": "message" + } + assert streaming_module._dump_response_object(object()) == {} + + @pytest.mark.asyncio async def test_responses_streaming_triggers_hooks(monkeypatch): """ @@ -167,3 +292,768 @@ async def test_responses_streaming_failure_triggers_failure_handlers(): await asyncio.sleep(0.2) assert logging_obj.failure_calls >= 1 assert logging_obj.async_failure_calls >= 1 + + +def test_process_chunk_requires_provider_config(): + iterator = ResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=None, + logging_obj=_FakeLoggingObj(), + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + + with pytest.raises(ValueError, match="responses_api_provider_config is required"): + iterator._process_chunk(json.dumps({"type": "response.completed"})) + + +def test_process_chunk_wraps_encrypted_content_with_model_id(): + openai_types = streaming_module._get_openai_response_types() + + class _EncryptedConfig: + def transform_streaming_response(self, **kwargs): + return openai_types.OutputItemAddedEvent( + type=openai_types.ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, + output_index=0, + item=openai_types.BaseLiteLLMOpenAIResponseObject( + id="rs_123", + type="reasoning", + encrypted_content="ciphertext", + ), + ) + + iterator = ResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=_EncryptedConfig(), + logging_obj=_FakeLoggingObj(), + litellm_metadata={ + "encrypted_content_affinity_enabled": True, + "model_info": {"id": "model-123"}, + }, + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + + event = iterator._process_chunk(json.dumps({"type": "response.output_item.added"})) + + assert event.item.encrypted_content.startswith("litellm_enc:") + assert event.item.encrypted_content.endswith(";ciphertext") + + +def test_process_chunk_completed_response_updates_id_and_usage_cost(monkeypatch): + original_include_cost = litellm.include_cost_in_streaming_usage + litellm.include_cost_in_streaming_usage = True + openai_types = streaming_module._get_openai_response_types() + + class _CompletedConfig: + def transform_streaming_response(self, **kwargs): + return openai_types.ResponseCompletedEvent( + type=openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=ResponsesAPIResponse( + id="resp_live", + created_at=int(datetime.now().timestamp()), + status="completed", + model="test-model", + object="response", + output=[], + usage=openai_types.ResponseAPIUsage( + input_tokens=1, + output_tokens=2, + total_tokens=3, + ), + ), + ) + + logging_obj = _FakeLoggingObj() + logging_obj._response_cost_calculator = MagicMock(return_value=1.23) + iterator = ResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=_CompletedConfig(), + logging_obj=logging_obj, + litellm_metadata={"model_info": {"id": "model-123"}}, + custom_llm_provider="openai", + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + completion_handler = MagicMock() + monkeypatch.setattr( + iterator, "_handle_logging_completed_response", completion_handler + ) + + try: + # Chunk must include a top-level "response" key so BaseResponsesAPIStreamingIterator + # runs _update_responses_api_response_id_with_model_id (see streaming_iterator.py). + event = iterator._process_chunk( + json.dumps( + {"type": "response.completed", "response": {"id": "resp_live"}} + ) + ) + finally: + litellm.include_cost_in_streaming_usage = original_include_cost + + assert iterator.completed_response is event + assert event.response.id != "resp_live" + assert event.response.id.startswith("resp_") + assert event.response.usage.cost == 1.23 + completion_handler.assert_called_once() + + +def test_process_chunk_failed_response_triggers_failure_logging(monkeypatch): + openai_types = streaming_module._get_openai_response_types() + + class _FailedConfig: + def transform_streaming_response(self, **kwargs): + return openai_types.ResponseFailedEvent( + type=openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED, + response=ResponsesAPIResponse( + id="resp_failed", + created_at=int(datetime.now().timestamp()), + status="failed", + model="test-model", + object="response", + output=[], + error={"message": "provider failed"}, + ), + ) + + iterator = ResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=_FailedConfig(), + logging_obj=_FakeLoggingObj(), + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + failure_handler = MagicMock() + monkeypatch.setattr(iterator, "_handle_logging_failed_response", failure_handler) + + event = iterator._process_chunk(json.dumps({"type": "response.failed"})) + + assert iterator.completed_response is event + failure_handler.assert_called_once() + + +@pytest.mark.asyncio +async def test_handle_logging_failed_response_uses_response_error_message(): + openai_types = streaming_module._get_openai_response_types() + logging_obj = _FakeLoggingObj() + iterator = ResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=SimpleNamespace(), + logging_obj=logging_obj, + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + iterator.completed_response = openai_types.ResponseFailedEvent( + type=openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED, + response=ResponsesAPIResponse( + id="resp_failed_real", + created_at=int(datetime.now().timestamp()), + status="failed", + model="test-model", + object="response", + output=[], + error={"message": "provider failed"}, + ), + ) + + iterator._handle_logging_failed_response() + await asyncio.sleep(0.2) + + assert logging_obj.failure_calls == 1 + assert logging_obj.async_failure_calls == 1 + + +def test_process_chunk_returns_none_for_invalid_json_and_non_dict_payload(): + class _NoopConfig: + def transform_streaming_response(self, **kwargs): + raise AssertionError("should not be called") + + iterator = ResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=_NoopConfig(), + logging_obj=_FakeLoggingObj(), + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + + assert iterator._process_chunk("not-json") is None + assert iterator._process_chunk(json.dumps(["not", "a", "dict"])) is None + + +def test_process_chunk_cost_annotation_failure_is_nonfatal(monkeypatch): + original_include_cost = litellm.include_cost_in_streaming_usage + litellm.include_cost_in_streaming_usage = True + openai_types = streaming_module._get_openai_response_types() + + class _CompletedConfig: + def transform_streaming_response(self, **kwargs): + return openai_types.ResponseCompletedEvent( + type=openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=ResponsesAPIResponse( + id="resp_cost_failure", + created_at=int(datetime.now().timestamp()), + status="completed", + model="test-model", + object="response", + output=[], + usage=openai_types.ResponseAPIUsage( + input_tokens=1, + output_tokens=2, + total_tokens=3, + ), + ), + ) + + logging_obj = _FakeLoggingObj() + logging_obj._response_cost_calculator = MagicMock(side_effect=RuntimeError("boom")) + iterator = ResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=_CompletedConfig(), + logging_obj=logging_obj, + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + completion_handler = MagicMock() + monkeypatch.setattr( + iterator, "_handle_logging_completed_response", completion_handler + ) + + try: + event = iterator._process_chunk(json.dumps({"type": "response.completed"})) + finally: + litellm.include_cost_in_streaming_usage = original_include_cost + + assert iterator.completed_response is event + assert event.response.usage.cost is None + completion_handler.assert_called_once() + + +def test_get_completed_response_object_accepts_direct_response(): + logging_obj = _FakeLoggingObj() + iterator = SyncResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=SimpleNamespace(), + logging_obj=logging_obj, + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + direct_response = _make_completed_response("resp_direct").response + iterator.completed_response = direct_response + + assert iterator._get_completed_response_object() is direct_response + + +@pytest.mark.asyncio +async def test_responses_streaming_completed_event_persists_async_cache(): + logging_obj = _FakeLoggingObj() + original_cache = litellm.cache + litellm.cache = SimpleNamespace( + async_add_cache=AsyncMock(), + add_cache=MagicMock(), + ) + caching_handler = SimpleNamespace( + request_kwargs={ + "model": "test-model", + "input": "hello", + "stream": True, + "caching": True, + "cache_key": "stale-request-cache-key", + "metadata": None, + "custom_llm_provider": "openai", + }, + preset_cache_key="responses-stream-cache-key", + original_function=litellm.aresponses, + async_set_cache=AsyncMock(), + _should_store_result_in_cache=lambda original_function, kwargs: True, + ) + logging_obj._llm_caching_handler = caching_handler + + iterator = ResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=SimpleNamespace(), + logging_obj=logging_obj, + request_data=caching_handler.request_kwargs, + call_type=CallTypes.aresponses.value, + ) + iterator.completed_response = _make_completed_response() + + iterator._handle_logging_completed_response() + await asyncio.sleep(0.2) + + litellm.cache.async_add_cache.assert_called_once() + assert litellm.cache.async_add_cache.call_args.kwargs["stream"] is True + assert ( + litellm.cache.async_add_cache.call_args.kwargs["cache_key"] + == "responses-stream-cache-key" + ) + assert "metadata" not in litellm.cache.async_add_cache.call_args.kwargs + assert "custom_llm_provider" not in litellm.cache.async_add_cache.call_args.kwargs + assert ( + json.loads(litellm.cache.async_add_cache.call_args.args[0])["id"] + == iterator.completed_response.response.id + ) + litellm.cache = original_cache + + +def test_responses_streaming_completed_event_persists_sync_cache(): + logging_obj = _FakeLoggingObj() + original_cache = litellm.cache + litellm.cache = SimpleNamespace( + async_add_cache=AsyncMock(), + add_cache=MagicMock(), + ) + caching_handler = SimpleNamespace( + request_kwargs={ + "model": "test-model", + "input": "hello", + "stream": True, + "caching": True, + "cache_key": "stale-request-cache-key", + "metadata": None, + "custom_llm_provider": "openai", + }, + preset_cache_key="responses-stream-cache-key", + original_function=litellm.responses, + sync_set_cache=MagicMock(), + _should_store_result_in_cache=lambda original_function, kwargs: True, + ) + logging_obj._llm_caching_handler = caching_handler + + iterator = SyncResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=SimpleNamespace(), + logging_obj=logging_obj, + request_data=caching_handler.request_kwargs, + call_type=CallTypes.responses.value, + ) + iterator.completed_response = _make_completed_response("resp_sync") + + iterator._handle_logging_completed_response() + + litellm.cache.add_cache.assert_called_once() + assert litellm.cache.add_cache.call_args.kwargs["stream"] is True + assert ( + litellm.cache.add_cache.call_args.kwargs["cache_key"] + == "responses-stream-cache-key" + ) + assert "metadata" not in litellm.cache.add_cache.call_args.kwargs + assert "custom_llm_provider" not in litellm.cache.add_cache.call_args.kwargs + assert ( + json.loads(litellm.cache.add_cache.call_args.args[0])["id"] + == iterator.completed_response.response.id + ) + litellm.cache = original_cache + + +def test_log_completed_response_sync_direct_path(monkeypatch): + hook_calls = {"post_call": 0, "metadata": 0} + + async def fake_post_call(request_data, response, call_type): + hook_calls["post_call"] += 1 + + def fake_update_metadata(**kwargs): + hook_calls["metadata"] += 1 + + monkeypatch.setattr( + streaming_module, + "async_post_call_success_deployment_hook", + fake_post_call, + ) + monkeypatch.setattr( + streaming_module, + "update_response_metadata", + fake_update_metadata, + ) + + logging_obj = _FakeLoggingObj() + iterator = SyncResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=SimpleNamespace(), + logging_obj=logging_obj, + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + iterator._persist_completed_response_before_logging = False + iterator.completed_response = _make_completed_response("resp_log_sync") + + iterator._log_completed_response(is_async=False) + asyncio.run(asyncio.sleep(0.2)) + + assert logging_obj.success_calls == 1 + assert logging_obj.async_success_calls == 1 + assert hook_calls["post_call"] == 1 + assert hook_calls["metadata"] == 1 + + +def test_log_completed_response_falls_back_when_model_validate_fails(monkeypatch): + class _BadSerializableResponse: + @classmethod + def model_validate(cls, value): + raise RuntimeError("nope") + + def model_dump(self): + return {"id": "bad"} + + logging_obj = _FakeLoggingObj() + iterator = SyncResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=SimpleNamespace(), + logging_obj=logging_obj, + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + iterator._persist_completed_response_before_logging = False + iterator.completed_response = _BadSerializableResponse() + monkeypatch.setattr(iterator, "_run_post_success_hooks", MagicMock()) + + iterator._log_completed_response(is_async=False) + asyncio.run(asyncio.sleep(0.2)) + + assert logging_obj.success_calls == 1 + assert logging_obj.async_success_calls == 1 + + +@pytest.mark.parametrize( + "scenario", + [ + "already_cached", + "not_completed", + "missing_caching_handler", + "not_streaming", + "store_disabled", + "missing_cache_backend", + ], +) +def test_persist_completed_response_to_cache_guard_branches(monkeypatch, scenario): + logging_obj = _FakeLoggingObj() + iterator = SyncResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=SimpleNamespace(), + logging_obj=logging_obj, + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + openai_types = streaming_module._get_openai_response_types() + completed_event = _make_completed_response("resp_guard") + iterator.completed_response = completed_event + + if scenario == "already_cached": + iterator._completed_response_cached = True + elif scenario == "not_completed": + iterator.completed_response = openai_types.ResponseIncompleteEvent( + type=openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, + response=completed_event.response, + ) + elif scenario == "missing_caching_handler": + logging_obj._llm_caching_handler = None + else: + logging_obj._llm_caching_handler = SimpleNamespace( + request_kwargs={ + "model": "test-model", + "input": "hello", + "stream": scenario != "not_streaming", + "cache_key": "request-cache-key", + "metadata": None, + "custom_llm_provider": "openai", + }, + preset_cache_key=None, + original_function=litellm.responses, + dual_cache=None, + _should_store_result_in_cache=lambda original_function, kwargs: ( + scenario != "store_disabled" + ), + ) + if scenario == "missing_cache_backend": + monkeypatch.setattr(streaming_module.litellm, "cache", None) + else: + monkeypatch.setattr( + streaming_module.litellm, + "cache", + SimpleNamespace(add_cache=MagicMock(), async_add_cache=AsyncMock()), + ) + + iterator._persist_completed_response_to_cache(is_async=False) + + expected_cached_flag = scenario == "already_cached" + assert iterator._completed_response_cached is expected_cached_flag + + +def test_build_synthetic_response_events_covers_annotations_function_calls_and_refusals(): + original_include_cost = litellm.include_cost_in_streaming_usage + litellm.include_cost_in_streaming_usage = True + logging_obj = _FakeLoggingObj() + logging_obj._response_cost_calculator = MagicMock(side_effect=RuntimeError("boom")) + transformed = ResponsesAPIResponse( + id="resp_events", + created_at=int(datetime.now().timestamp()), + status="completed", + model="gpt-4.1-mini", + object="response", + output=[ + { + "type": "message", + "id": "msg_events", + "status": "completed", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "hello world", + "annotations": [{"type": "file_citation", "file_id": "file_1"}], + }, + { + "type": "refusal", + "refusal": "no thanks", + }, + ], + }, + { + "type": "function_call", + "id": "fc_events", + "call_id": "call_123", + "name": "lookup", + "arguments": '{"id":1}', + }, + ], + ) + + try: + events = streaming_module._build_synthetic_response_events( + transformed=transformed, + logging_obj=logging_obj, + chunk_size=5, + ) + finally: + litellm.include_cost_in_streaming_usage = original_include_cost + + event_types = [ + event.type.value if hasattr(event.type, "value") else str(event.type) + for event in events + ] + + assert "response.output_text.annotation.added" in event_types + assert "response.refusal.delta" in event_types + assert "response.refusal.done" in event_types + assert "response.function_call_arguments.delta" in event_types + assert "response.function_call_arguments.done" in event_types + assert event_types[-1] == "response.completed" + + +@pytest.mark.asyncio +async def test_mock_responses_streaming_iterator_async_iteration_logs_completion( + monkeypatch, +): + hook_calls = {"post_call": 0, "metadata": 0} + + async def fake_post_call(request_data, response, call_type): + hook_calls["post_call"] += 1 + + def fake_update_metadata(**kwargs): + hook_calls["metadata"] += 1 + + monkeypatch.setattr( + streaming_module, + "async_post_call_success_deployment_hook", + fake_post_call, + ) + monkeypatch.setattr( + streaming_module, + "update_response_metadata", + fake_update_metadata, + ) + + class _MockTransformConfig: + def transform_response_api_response(self, **kwargs): + return _make_completed_response("resp_mock").response + + logging_obj = _FakeLoggingObj() + + iterator = MockResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=_MockTransformConfig(), + logging_obj=logging_obj, + request_data={"model": "test-model", "stream": True}, + call_type=CallTypes.responses.value, + ) + + streamed_events = [event async for event in iterator] + await asyncio.sleep(0.2) + + assert streamed_events[0].type == ResponsesAPIStreamEvents.RESPONSE_CREATED + assert streamed_events[-1].type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + assert logging_obj.success_calls == 1 + assert logging_obj.async_success_calls == 1 + assert hook_calls["post_call"] == 1 + assert hook_calls["metadata"] == 1 + + +def test_mock_responses_streaming_iterator_sync_iteration_logs_completion(monkeypatch): + hook_calls = {"post_call": 0, "metadata": 0} + + async def fake_post_call(request_data, response, call_type): + hook_calls["post_call"] += 1 + + def fake_update_metadata(**kwargs): + hook_calls["metadata"] += 1 + + monkeypatch.setattr( + streaming_module, + "async_post_call_success_deployment_hook", + fake_post_call, + ) + monkeypatch.setattr( + streaming_module, + "update_response_metadata", + fake_update_metadata, + ) + + class _MockTransformConfig: + def transform_response_api_response(self, **kwargs): + return _make_completed_response("resp_mock_sync").response + + logging_obj = _FakeLoggingObj() + iterator = MockResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=_MockTransformConfig(), + logging_obj=logging_obj, + request_data={"model": "test-model", "stream": True}, + call_type=CallTypes.responses.value, + ) + + streamed_events = list(iterator) + asyncio.run(asyncio.sleep(0.2)) + + assert streamed_events[0].type == ResponsesAPIStreamEvents.RESPONSE_CREATED + assert streamed_events[-1].type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + assert logging_obj.success_calls == 1 + assert logging_obj.async_success_calls == 1 + assert hook_calls["post_call"] == 1 + assert hook_calls["metadata"] == 1 + + +@pytest.mark.asyncio +async def test_cached_responses_stream_async_hit_triggers_success_callbacks( + monkeypatch, +): + hook_calls = {"post_call": 0, "metadata": 0} + + async def fake_post_call(request_data, response, call_type): + hook_calls["post_call"] += 1 + + def fake_update_metadata(**kwargs): + hook_calls["metadata"] += 1 + + monkeypatch.setattr( + streaming_module, + "async_post_call_success_deployment_hook", + fake_post_call, + ) + monkeypatch.setattr( + streaming_module, + "update_response_metadata", + fake_update_metadata, + ) + + logging_obj = _FakeLoggingObj() + original_cache = litellm.cache + litellm.cache = SimpleNamespace( + async_add_cache=AsyncMock(), + add_cache=MagicMock(), + ) + logging_obj._llm_caching_handler = SimpleNamespace( + request_kwargs={"model": "test-model", "input": "hello", "stream": True}, + preset_cache_key="responses-stream-cache-key", + original_function=litellm.aresponses, + _should_store_result_in_cache=lambda original_function, kwargs: True, + ) + + iterator = CachedResponsesAPIStreamingIterator( + response=_make_completed_response("resp_cached_async").response, + logging_obj=logging_obj, + request_data={"model": "test-model", "input": "hello", "stream": True}, + call_type=CallTypes.aresponses.value, + ) + + streamed_events = [event async for event in iterator] + await asyncio.sleep(0.2) + + assert streamed_events[-1].type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + assert logging_obj.success_calls == 1 + assert logging_obj.async_success_calls == 1 + assert logging_obj.last_success_kwargs["cache_hit"] is True + assert logging_obj.last_async_success_kwargs["cache_hit"] is True + assert hook_calls["post_call"] == 1 + assert hook_calls["metadata"] == 1 + litellm.cache.async_add_cache.assert_not_called() + litellm.cache.add_cache.assert_not_called() + litellm.cache = original_cache + + +def test_cached_responses_stream_sync_hit_triggers_success_callbacks(monkeypatch): + hook_calls = {"post_call": 0, "metadata": 0} + + async def fake_post_call(request_data, response, call_type): + hook_calls["post_call"] += 1 + + def fake_update_metadata(**kwargs): + hook_calls["metadata"] += 1 + + monkeypatch.setattr( + streaming_module, + "async_post_call_success_deployment_hook", + fake_post_call, + ) + monkeypatch.setattr( + streaming_module, + "update_response_metadata", + fake_update_metadata, + ) + + logging_obj = _FakeLoggingObj() + original_cache = litellm.cache + litellm.cache = SimpleNamespace( + async_add_cache=AsyncMock(), + add_cache=MagicMock(), + ) + logging_obj._llm_caching_handler = SimpleNamespace( + request_kwargs={"model": "test-model", "input": "hello", "stream": True}, + preset_cache_key="responses-stream-cache-key", + original_function=litellm.responses, + _should_store_result_in_cache=lambda original_function, kwargs: True, + ) + + iterator = CachedResponsesAPIStreamingIterator( + response=_make_completed_response("resp_cached_sync").response, + logging_obj=logging_obj, + request_data={"model": "test-model", "input": "hello", "stream": True}, + call_type=CallTypes.responses.value, + ) + + streamed_events = list(iterator) + asyncio.run(asyncio.sleep(0.2)) + + assert streamed_events[-1].type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + assert logging_obj.success_calls == 1 + assert logging_obj.async_success_calls == 1 + assert logging_obj.last_success_kwargs["cache_hit"] is True + assert logging_obj.last_async_success_kwargs["cache_hit"] is True + assert hook_calls["post_call"] == 1 + assert hook_calls["metadata"] == 1 + litellm.cache.async_add_cache.assert_not_called() + litellm.cache.add_cache.assert_not_called() + litellm.cache = original_cache diff --git a/tests/llm_translation/Readme.md b/tests/llm_translation/Readme.md index db84e7c33c9..958adbd9755 100644 --- a/tests/llm_translation/Readme.md +++ b/tests/llm_translation/Readme.md @@ -1,3 +1,41 @@ -Unit tests for individual LLM providers. +Unit tests for individual LLM providers. -Name of the test file is the name of the LLM provider - e.g. `test_openai.py` is for OpenAI. \ No newline at end of file +Name of the test file is the name of the LLM provider - e.g. `test_openai.py` is for OpenAI. + +## Redis-backed VCR cache + +Every test in this directory is auto-decorated with `@pytest.mark.vcr` (via +`conftest.py`). The first time a test runs we hit the live provider and +record the HTTP exchange into Redis under +`litellm:vcr:cassette:`. Every subsequent run within 24h replays +from Redis without touching the network. The 24h TTL means each new day's +first run records again, so upstream API drift surfaces within a day. + +The persister, header scrubbing, and 2xx-only filtering are defined in +`tests/_vcr_redis_persister.py`. Files that already use `respx` (which +patches the same httpx transport vcrpy does) are excluded from the +auto-marker — see `_RESPX_CONFLICTING_FILES` in `conftest.py`. + +### Required environment + +`REDIS_HOST`, `REDIS_PORT`, `REDIS_PASSWORD` — same vars CircleCI uses for +its other Redis-backed jobs. Provider credentials +(`ANTHROPIC_API_KEY`, `OPENAI_API_KEY`, `AWS_*`, etc.) are needed only on +cache-miss (the daily re-record), not on replay. + +### Flushing the cache + +When you want the next run to re-record immediately instead of waiting +for the 24h TTL: + +```bash +make test-llm-translation-flush-vcr-cache +``` + +### Disabling VCR + +Skip the cache entirely (every call goes live, no recording): + +```bash +LITELLM_VCR_DISABLE=1 uv run pytest tests/llm_translation/test_.py +``` diff --git a/tests/llm_translation/base_llm_unit_tests.py b/tests/llm_translation/base_llm_unit_tests.py index f3b1895323c..aedc4f810cd 100644 --- a/tests/llm_translation/base_llm_unit_tests.py +++ b/tests/llm_translation/base_llm_unit_tests.py @@ -401,7 +401,7 @@ class BaseLLMChatTest(ABC): { "type": "file", "file": { - "file_id": "https://upload.wikimedia.org/wikipedia/commons/2/20/Re_example.pdf" + "file_id": "https://cdn.jsdelivr.net/gh/BerriAI/litellm@d769e81c90d453240c61fc572cdb27fae06a89d0/tests/llm_translation/fixtures/dummy.pdf" }, }, ] diff --git a/tests/llm_translation/conftest.py b/tests/llm_translation/conftest.py index d315dc63bcd..09da0520be0 100644 --- a/tests/llm_translation/conftest.py +++ b/tests/llm_translation/conftest.py @@ -5,6 +5,7 @@ # - Function-scoped fixture resets litellm globals to true defaults # - Module-scoped reload only in single-process mode +import asyncio import importlib import os import sys @@ -14,9 +15,195 @@ import pytest sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path -import litellm -import asyncio +import litellm # noqa: E402 + +from tests._vcr_redis_persister import ( # noqa: E402 + filter_non_2xx_response, + format_vcr_verdict, + make_redis_persister, + mark_test_outcome_for_cassette, + patch_vcrpy_aiohttp_record_path, + vcr_verbose_enabled, +) + + +_controller_pluginmanager = None +_controller_terminal_reporter = None + + +# vcrpy and respx both patch the httpx transport — applying both makes one +# silently win, so respx-using files opt out of the auto-marker. +_RESPX_CONFLICTING_FILES = frozenset( + { + "test_azure_o_series.py", + "test_gpt4o_audio.py", + "test_nvidia_nim.py", + "test_openai.py", + "test_openai_o1.py", + "test_prompt_caching.py", + "test_text_completion_unit_tests.py", + "test_xai.py", + } +) +_VCR_AUTO_MARKER_SKIP_FILES = _RESPX_CONFLICTING_FILES | frozenset( + {"test_vcr_redis_persister.py"} +) + +# Tests that observe live cross-call provider state (e.g. prompt-cache +# warm-up between two consecutive calls); replay can't reproduce that state. +_VCR_INCOMPATIBLE_NODEID_SUFFIXES = frozenset( + { + "::test_prompt_caching", + "TestBedrockInvokeNovaJson::test_json_response_pydantic_obj", + "::test_bedrock_converse__streaming_passthrough", + } +) + + +def _is_vcr_incompatible(nodeid: str) -> bool: + return any(nodeid.endswith(suffix) for suffix in _VCR_INCOMPATIBLE_NODEID_SUFFIXES) + + +_FILTERED_REQUEST_HEADERS = ( + "authorization", + "x-api-key", + "anthropic-api-key", + "anthropic-version", + "openai-api-key", + "azure-api-key", + "api-key", + "cookie", + "x-amz-security-token", + "x-amz-date", + "x-amz-content-sha256", + "amz-sdk-invocation-id", + "amz-sdk-request", + "x-goog-api-key", + "x-goog-user-project", +) + +_FILTERED_RESPONSE_HEADERS = ( + "set-cookie", + "x-request-id", + "request-id", + "cf-ray", + "anthropic-organization-id", + "openai-organization", + "x-amzn-requestid", + "x-amzn-trace-id", + "date", +) + + +def _scrub_response(response): + if not isinstance(response, dict): + return response + headers = response.get("headers") or {} + if isinstance(headers, dict): + for header in list(headers): + if header.lower() in _FILTERED_RESPONSE_HEADERS: + headers.pop(header, None) + return response + + +def _before_record_response(response): + return filter_non_2xx_response(_scrub_response(response)) + + +@pytest.fixture(scope="module") +def vcr_config(): + return { + "filter_headers": list(_FILTERED_REQUEST_HEADERS), + "decode_compressed_response": True, + "record_mode": "new_episodes", + "allow_playback_repeats": True, + "match_on": ( + "method", + "scheme", + "host", + "port", + "path", + "query", + "body", + ), + "before_record_response": _before_record_response, + } + + +def _vcr_disabled() -> bool: + if os.environ.get("LITELLM_VCR_DISABLE") == "1": + return True + return not os.environ.get("CASSETTE_REDIS_URL") + + +def pytest_recording_configure(config, vcr): + if _vcr_disabled(): + return + vcr.register_persister(make_redis_persister()) + patch_vcrpy_aiohttp_record_path() + + +@pytest.hookimpl(hookwrapper=True) +def pytest_runtest_makereport(item, call): + outcome = yield + rep = outcome.get_result() + setattr(item, f"rep_{rep.when}", rep) + + +@pytest.fixture(autouse=True) +def _vcr_outcome_gate(request, vcr): + yield + cassette = vcr + rep_call = getattr(request.node, "rep_call", None) + test_passed = bool(rep_call and rep_call.passed) + cassette_path = getattr(cassette, "_path", None) if cassette is not None else None + if cassette_path: + mark_test_outcome_for_cassette(cassette_path, test_passed) + + if not vcr_verbose_enabled(): + return + verdict = format_vcr_verdict(cassette) + request.node.user_properties.append(("vcr_verdict", verdict)) + + +def pytest_configure(config): + global _controller_pluginmanager + if os.environ.get("PYTEST_XDIST_WORKER"): + return + _controller_pluginmanager = config.pluginmanager + + +def _resolve_terminal_reporter(): + global _controller_terminal_reporter + if _controller_terminal_reporter is not None: + return _controller_terminal_reporter + if _controller_pluginmanager is None: + return None + _controller_terminal_reporter = _controller_pluginmanager.getplugin( + "terminalreporter" + ) + return _controller_terminal_reporter + + +def pytest_runtest_logreport(report): + if report.when != "teardown": + return + if os.environ.get("PYTEST_XDIST_WORKER"): + return + if not vcr_verbose_enabled(): + return + reporter = _resolve_terminal_reporter() + if reporter is None: + return + verdict = next( + (v for k, v in (report.user_properties or []) if k == "vcr_verdict"), + None, + ) + if not verdict: + return + reporter.write_line(f"{verdict} :: {report.nodeid}") + # --------------------------------------------------------------------------- # Capture TRUE defaults at conftest import time (before test modules pollute). @@ -48,7 +235,6 @@ def event_loop(): @pytest.fixture(scope="function", autouse=True) def setup_and_teardown(event_loop): # Add event_loop as a dependency - curr_dir = os.getcwd() sys.path.insert(0, os.path.abspath("../..")) import litellm @@ -97,15 +283,23 @@ def setup_and_teardown(event_loop): # Add event_loop as a dependency def pytest_collection_modifyitems(config, items): - # Separate tests in 'test_amazing_proxy_custom_logger.py' and other tests + if not _vcr_disabled(): + for item in items: + filename = os.path.basename(str(item.fspath)) + if filename in _VCR_AUTO_MARKER_SKIP_FILES: + continue + if _is_vcr_incompatible(item.nodeid): + continue + if item.get_closest_marker("vcr") is not None: + continue + item.add_marker(pytest.mark.vcr) + custom_logger_tests = [ item for item in items if "custom_logger" in item.parent.name ] other_tests = [item for item in items if "custom_logger" not in item.parent.name] - # Sort tests based on their names custom_logger_tests.sort(key=lambda x: x.name) other_tests.sort(key=lambda x: x.name) - # Reorder the items list items[:] = custom_logger_tests + other_tests diff --git a/tests/llm_translation/test_anthropic_completion.py b/tests/llm_translation/test_anthropic_completion.py index 7b2b6bed6a4..371b27c5b21 100644 --- a/tests/llm_translation/test_anthropic_completion.py +++ b/tests/llm_translation/test_anthropic_completion.py @@ -1885,3 +1885,42 @@ def test_metadata_filter_applies_to_azure_anthropic(): headers={}, ) assert data.get("metadata") == {"user_id": "u2"} + + +def test_anthropic_basic_completion_replay(): + response = litellm.completion( + model="anthropic/claude-sonnet-4-5-20250929", + messages=[{"role": "user", "content": "Hello!"}], + ) + + assert response is not None + content = response.choices[0].message.content + assert isinstance(content, str) and content.strip(), content + assert response.usage.prompt_tokens > 0 + assert response.usage.completion_tokens > 0 + assert response.choices[0].finish_reason in {"stop", "length"} + + +def test_anthropic_streaming_completion_replay(): + stream = litellm.completion( + model="anthropic/claude-sonnet-4-5-20250929", + messages=[{"role": "user", "content": "Hello!"}], + stream=True, + ) + + collected_text = "" + finish_reason = None + chunk_count = 0 + for chunk in stream: + chunk_count += 1 + if not chunk.choices: + continue + delta = chunk.choices[0].delta + if delta and delta.content: + collected_text += delta.content + if chunk.choices[0].finish_reason: + finish_reason = chunk.choices[0].finish_reason + + assert chunk_count > 1, "expected multiple SSE chunks from streaming response" + assert collected_text.strip(), collected_text + assert finish_reason in {"stop", "length"} diff --git a/tests/llm_translation/test_vcr_redis_persister.py b/tests/llm_translation/test_vcr_redis_persister.py new file mode 100644 index 00000000000..853558150c1 --- /dev/null +++ b/tests/llm_translation/test_vcr_redis_persister.py @@ -0,0 +1,243 @@ +from __future__ import annotations + +import os +import sys + +import fakeredis +import pytest +from redis.exceptions import ConnectionError as RedisConnectionError +from vcr.persisters.filesystem import CassetteNotFoundError +from vcr.request import Request +from vcr.serializers import yamlserializer + +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", ".."))) + +from tests._vcr_redis_persister import ( # noqa: E402 + CASSETTE_TTL_SECONDS, + MAX_EPISODES_PER_CASSETTE, + filter_non_2xx_response, + make_redis_persister, + mark_test_outcome_for_cassette, + redis_key_for, +) + + +def _sample_cassette_dict(): + request = Request( + method="POST", + uri="https://api.anthropic.com/v1/messages", + body=b'{"model":"claude","messages":[{"role":"user","content":"hi"}]}', + headers={"content-type": "application/json"}, + ) + response = { + "status": {"code": 200, "message": "OK"}, + "headers": {"content-type": ["application/json"]}, + "body": {"string": b'{"id":"msg_1","type":"message"}'}, + } + return {"requests": [request], "responses": [response]} + + +def _persister_with_fake_redis(): + fake = fakeredis.FakeStrictRedis() + return fake, make_redis_persister(client=fake) + + +def test_save_then_load_roundtrips_cassette_content(): + _, persister = _persister_with_fake_redis() + cassette_id = "tests/llm_translation/test_x/test_y" + + persister.save_cassette(cassette_id, _sample_cassette_dict(), yamlserializer) + requests, responses = persister.load_cassette(cassette_id, yamlserializer) + + assert len(requests) == 1 + assert len(responses) == 1 + assert requests[0].method == "POST" + assert requests[0].uri == "https://api.anthropic.com/v1/messages" + assert responses[0]["status"]["code"] == 200 + assert responses[0]["body"]["string"] == b'{"id":"msg_1","type":"message"}' + + +def test_saved_key_has_24h_ttl(): + fake, persister = _persister_with_fake_redis() + cassette_id = "tests/llm_translation/test_x/test_ttl" + + persister.save_cassette(cassette_id, _sample_cassette_dict(), yamlserializer) + + ttl = fake.ttl(redis_key_for(cassette_id)) + assert CASSETTE_TTL_SECONDS - 5 <= ttl <= CASSETTE_TTL_SECONDS + + +def test_load_missing_key_raises_cassette_not_found(): + _, persister = _persister_with_fake_redis() + with pytest.raises(CassetteNotFoundError): + persister.load_cassette("never/recorded", yamlserializer) + + +def test_redis_key_normalizes_path_passed_by_pytest_recording(): + raw = "tests/llm_translation/cassettes/test_anthropic/test_streaming.yaml" + assert ( + redis_key_for(raw) + == "litellm:vcr:cassette:tests/llm_translation/test_anthropic/test_streaming" + ) + + +class _FlakyRedis: + def __init__(self, inner, fail_on: str): + self._inner = inner + self._fail_on = fail_on + + def get(self, *args, **kwargs): + if self._fail_on == "get": + raise RedisConnectionError("simulated outage") + return self._inner.get(*args, **kwargs) + + def set(self, *args, **kwargs): + if self._fail_on == "set": + raise RedisConnectionError("simulated outage") + return self._inner.set(*args, **kwargs) + + +def test_save_swallows_connection_errors_so_teardown_does_not_fail(): + flaky = _FlakyRedis(fakeredis.FakeStrictRedis(), fail_on="set") + persister = make_redis_persister(client=flaky) + + persister.save_cassette( + "tests/llm_translation/test_x/test_save_outage", + _sample_cassette_dict(), + yamlserializer, + ) + + +def test_save_skipped_when_test_marked_failed_and_prior_cassette_preserved(): + fake, persister = _persister_with_fake_redis() + cassette_id = "tests/llm_translation/test_x/test_flaky" + key = redis_key_for(cassette_id) + + good = _sample_cassette_dict() + persister.save_cassette(cassette_id, good, yamlserializer) + good_payload = fake.get(key) + assert good_payload is not None + + mark_test_outcome_for_cassette(cassette_id, passed=False) + bad_response = { + "status": {"code": 200, "message": "OK"}, + "headers": {}, + "body": {"string": b'{"id":"BAD","type":"message"}'}, + } + bad = {"requests": good["requests"], "responses": [bad_response]} + persister.save_cassette(cassette_id, bad, yamlserializer) + + assert fake.get(key) == good_payload + + +def test_save_proceeds_when_test_marked_passed(): + fake, persister = _persister_with_fake_redis() + cassette_id = "tests/llm_translation/test_x/test_passed" + key = redis_key_for(cassette_id) + + mark_test_outcome_for_cassette(cassette_id, passed=True) + persister.save_cassette(cassette_id, _sample_cassette_dict(), yamlserializer) + + assert fake.get(key) is not None + + +def test_save_refused_when_cassette_exceeds_max_episodes(): + fake, persister = _persister_with_fake_redis() + cassette_id = "tests/llm_translation/test_x/test_runaway" + key = redis_key_for(cassette_id) + + persister.save_cassette(cassette_id, _sample_cassette_dict(), yamlserializer) + seed_payload = fake.get(key) + + request = Request( + method="POST", + uri="https://api.anthropic.com/v1/messages", + body=b"x", + headers={"content-type": "application/json"}, + ) + response = { + "status": {"code": 200, "message": "OK"}, + "headers": {}, + "body": {"string": b"{}"}, + } + bloated = { + "requests": [request] * (MAX_EPISODES_PER_CASSETTE + 1), + "responses": [response] * (MAX_EPISODES_PER_CASSETTE + 1), + } + persister.save_cassette(cassette_id, bloated, yamlserializer) + + assert fake.get(key) == seed_payload + + +def test_save_proceeds_at_max_episodes_threshold(): + fake, persister = _persister_with_fake_redis() + cassette_id = "tests/llm_translation/test_x/test_at_threshold" + key = redis_key_for(cassette_id) + + request = Request( + method="POST", + uri="https://api.anthropic.com/v1/messages", + body=b"x", + headers={"content-type": "application/json"}, + ) + response = { + "status": {"code": 200, "message": "OK"}, + "headers": {}, + "body": {"string": b"{}"}, + } + at_threshold = { + "requests": [request] * MAX_EPISODES_PER_CASSETTE, + "responses": [response] * MAX_EPISODES_PER_CASSETTE, + } + persister.save_cassette(cassette_id, at_threshold, yamlserializer) + + assert fake.get(key) is not None + + +def test_save_proceeds_when_outcome_unknown(): + fake, persister = _persister_with_fake_redis() + cassette_id = "tests/llm_translation/test_x/test_no_marker" + key = redis_key_for(cassette_id) + + persister.save_cassette(cassette_id, _sample_cassette_dict(), yamlserializer) + + assert fake.get(key) is not None + + +def test_load_treats_connection_errors_as_cassette_miss(): + flaky = _FlakyRedis(fakeredis.FakeStrictRedis(), fail_on="get") + persister = make_redis_persister(client=flaky) + + with pytest.raises(CassetteNotFoundError): + persister.load_cassette( + "tests/llm_translation/test_x/test_load_outage", yamlserializer + ) + + +@pytest.mark.parametrize( + ("status_code", "expect_dropped"), + [ + (200, False), + (201, False), + (204, False), + (299, False), + (300, True), + (400, True), + (401, True), + (404, True), + (429, True), + (500, True), + (502, True), + (503, True), + ], +) +def test_only_2xx_responses_are_cached(status_code, expect_dropped): + response = { + "status": {"code": status_code, "message": "X"}, + "headers": {}, + "body": {"string": ""}, + } + result = filter_non_2xx_response(response) + assert (result is None) == expect_dropped + if not expect_dropped: + assert result is response diff --git a/tests/local_testing/test_caching_handler.py b/tests/local_testing/test_caching_handler.py index 806f72bfde8..2b6712cbaa3 100644 --- a/tests/local_testing/test_caching_handler.py +++ b/tests/local_testing/test_caching_handler.py @@ -19,9 +19,14 @@ import pytest import litellm from litellm import aembedding, completion, embedding, aresponses, responses from litellm.caching.caching import Cache +from litellm.responses.streaming_iterator import CachedResponsesAPIStreamingIterator from unittest.mock import AsyncMock, patch, MagicMock -from litellm.caching.caching_handler import LLMCachingHandler, CachingHandlerResponse +from litellm.caching.caching_handler import ( + LLMCachingHandler, + CachingHandlerResponse, + _should_defer_streaming_cache_hit_callbacks, +) from litellm.caching.caching import LiteLLMCacheType from litellm.types.utils import CallTypes from litellm.types.rerank import RerankResponse @@ -627,6 +632,55 @@ async def test_async_responses_api_caching(): assert cached_response.cached_result._hidden_params["cache_hit"] == True +@pytest.mark.asyncio +async def test_async_get_cache_updates_request_kwargs_for_streaming_responses(): + """ + Ensure streamed responses retain the normalized lookup kwargs so a later + cache write can reuse the exact cache key from the read path. + """ + setup_cache() + + caching_handler = LLMCachingHandler( + original_function=aresponses, + request_kwargs={"stale": True}, + start_time=datetime.now(), + ) + + logging_obj = LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=CallTypes.aresponses.value, + model="gpt-4o", + messages=[], + function_id=str(uuid.uuid4()), + stream=True, + start_time=datetime.now(), + ) + + kwargs = { + "model": "gpt-4o", + "input": "hello", + "stream": True, + "caching": True, + } + + await caching_handler._async_get_cache( + model="gpt-4o", + original_function=aresponses, + logging_obj=logging_obj, + start_time=datetime.now(), + call_type=CallTypes.aresponses.value, + kwargs=kwargs, + ) + + assert "stale" not in caching_handler.request_kwargs + assert caching_handler.request_kwargs["model"] == "gpt-4o" + assert caching_handler.request_kwargs["input"] == "hello" + assert caching_handler.request_kwargs["stream"] is True + assert caching_handler.request_kwargs["cache_key"] == litellm.cache.get_cache_key( + **caching_handler.request_kwargs + ) + + def test_sync_responses_api_caching(): """ Test that synchronous responses API calls are properly cached and retrieved. @@ -769,6 +823,339 @@ def test_convert_cached_responses_api_result_to_model_response(): assert len(result.output) == 1 +def test_sync_get_cache_does_not_eagerly_log_streaming_responses_hits(): + litellm.set_verbose = True + setup_cache() + caching_handler = LLMCachingHandler( + original_function=responses, request_kwargs={}, start_time=datetime.now() + ) + + original_model = "gpt-4o" + responses_api_response = ResponsesAPIResponse( + id="resp_stream_sync_hit", + created_at=int(time.time()), + status="completed", + model=original_model, + object="response", + output=[ + { + "type": "message", + "id": "msg_stream_sync_hit", + "status": "completed", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "Sync streamed cache hit response.", + "annotations": [], + } + ], + } + ], + ) + + logging_obj = LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=CallTypes.responses.value, + model=original_model, + messages=[], + function_id=str(uuid.uuid4()), + stream=True, + start_time=datetime.now(), + ) + logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock() + + kwargs = { + "model": original_model, + "input": "Tell me a cached story", + "stream": True, + "caching": True, + } + + caching_handler.sync_set_cache(result=responses_api_response, kwargs=kwargs) + time.sleep(0.2) + + cached_response = caching_handler._sync_get_cache( + model=original_model, + original_function=responses, + logging_obj=logging_obj, + start_time=datetime.now(), + call_type=CallTypes.responses.value, + kwargs=kwargs, + ) + + assert cached_response.cached_result is not None + assert isinstance( + cached_response.cached_result, CachedResponsesAPIStreamingIterator + ) + logging_obj.handle_sync_success_callbacks_for_async_calls.assert_not_called() + + +def test_sync_get_cache_defers_streaming_completion_hit_callbacks(): + litellm.set_verbose = True + setup_cache() + caching_handler = LLMCachingHandler( + original_function=completion, request_kwargs={}, start_time=datetime.now() + ) + + original_model = "gpt-4o" + logging_obj = LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=CallTypes.completion.value, + model=original_model, + messages=[], + function_id=str(uuid.uuid4()), + stream=True, + start_time=datetime.now(), + ) + logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock() + + kwargs = { + "model": original_model, + "messages": [{"role": "user", "content": "Tell me a cached joke"}], + "stream": True, + "caching": True, + } + + caching_handler.sync_set_cache(result=chat_completion_response, kwargs=kwargs) + time.sleep(0.2) + + cached_response = caching_handler._sync_get_cache( + model=original_model, + original_function=completion, + logging_obj=logging_obj, + start_time=datetime.now(), + call_type=CallTypes.completion.value, + kwargs=kwargs, + ) + + assert cached_response.cached_result is not None + logging_obj.handle_sync_success_callbacks_for_async_calls.assert_not_called() + + +def test_should_defer_streaming_cache_hit_callbacks_for_any_streaming_request(): + assert ( + _should_defer_streaming_cache_hit_callbacks( + kwargs={"stream": True}, + ) + is True + ) + assert ( + _should_defer_streaming_cache_hit_callbacks( + kwargs={"stream": False}, + ) + is False + ) + assert ( + _should_defer_streaming_cache_hit_callbacks( + kwargs={}, + ) + is False + ) + + +@pytest.mark.asyncio +async def test_async_get_cache_defers_streaming_completion_hit_callbacks(): + litellm.set_verbose = True + setup_cache() + caching_handler = LLMCachingHandler( + original_function=completion, request_kwargs={}, start_time=datetime.now() + ) + + original_model = "gpt-4o" + kwargs = { + "model": original_model, + "messages": [{"role": "user", "content": "Tell me a cached joke"}], + "stream": True, + "caching": True, + } + + await caching_handler.async_set_cache( + result=chat_completion_response, + original_function=litellm.acompletion, + kwargs=kwargs, + ) + await asyncio.sleep(0.2) + + logging_obj = LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=CallTypes.acompletion.value, + model=original_model, + messages=[], + function_id=str(uuid.uuid4()), + stream=True, + start_time=datetime.now(), + ) + caching_handler._async_log_cache_hit_on_callbacks = MagicMock() + + cached_response = await caching_handler._async_get_cache( + model=original_model, + original_function=litellm.acompletion, + logging_obj=logging_obj, + start_time=datetime.now(), + call_type=CallTypes.acompletion.value, + kwargs=kwargs, + ) + + assert cached_response is not None + assert cached_response.cached_result is not None + caching_handler._async_log_cache_hit_on_callbacks.assert_not_called() + + +def test_convert_cached_streaming_responses_result_to_iterator(): + """ + Test that cached streaming Responses results are replayed through a synthetic + streaming iterator instead of being returned as a full response object. + """ + caching_handler = LLMCachingHandler( + original_function=responses, request_kwargs={}, start_time=datetime.now() + ) + + logging_obj = LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=CallTypes.responses.value, + model="gpt-4o", + messages=[], + function_id=str(uuid.uuid4()), + stream=True, + start_time=datetime.now(), + ) + + cached_result = { + "id": "resp_stream_cache_test", + "created_at": int(time.time()), + "status": "completed", + "model": "gpt-4o", + "object": "response", + "output": [ + { + "type": "message", + "id": "msg_stream_cache_test", + "status": "completed", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "Streaming cache replay test.", + "annotations": [], + } + ], + } + ], + } + + result = caching_handler._convert_cached_result_to_model_response( + cached_result=cached_result, + call_type=CallTypes.responses.value, + kwargs={"model": "gpt-4o", "input": "test", "stream": True}, + logging_obj=logging_obj, + model="gpt-4o", + args=(), + ) + + assert isinstance(result, CachedResponsesAPIStreamingIterator) + assert result.completed_response is not None + assert result.completed_response.response.id == cached_result["id"] + + streamed_events = list(result) + assert streamed_events[0].type == "response.created" + assert streamed_events[1].type == "response.in_progress" + assert streamed_events[2].type == "response.output_item.added" + assert streamed_events[3].type == "response.content_part.added" + assert streamed_events[-4].type == "response.output_text.done" + assert streamed_events[-3].type == "response.content_part.done" + assert streamed_events[-2].type == "response.output_item.done" + assert streamed_events[-1].type == "response.completed" + assert streamed_events[-1].response.id == cached_result["id"] + assert streamed_events[-1].response.output[0].content[0].text == ( + "Streaming cache replay test." + ) + + +def test_convert_cached_streaming_reasoning_result_to_iterator(): + caching_handler = LLMCachingHandler( + original_function=responses, request_kwargs={}, start_time=datetime.now() + ) + + logging_obj = LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=CallTypes.responses.value, + model="gpt-4o", + messages=[], + function_id=str(uuid.uuid4()), + stream=True, + start_time=datetime.now(), + ) + + cached_result = { + "id": "resp_stream_reasoning_cache_test", + "created_at": int(time.time()), + "status": "completed", + "model": "gpt-4o", + "object": "response", + "output": [ + { + "type": "reasoning", + "id": "rs_stream_cache_test", + "summary": [ + { + "type": "summary_text", + "text": "Cached reasoning summary.", + } + ], + } + ], + } + + result = caching_handler._convert_cached_result_to_model_response( + cached_result=cached_result, + call_type=CallTypes.responses.value, + kwargs={"model": "gpt-4o", "input": "test", "stream": True}, + logging_obj=logging_obj, + model="gpt-4o", + args=(), + ) + + assert isinstance(result, CachedResponsesAPIStreamingIterator) + + streamed_events = list(result) + streamed_event_types = [ + event.type.value if hasattr(event.type, "value") else str(event.type) + for event in streamed_events + ] + + assert streamed_event_types[:3] == [ + "response.created", + "response.in_progress", + "response.output_item.added", + ] + assert streamed_event_types[-4:] == [ + "response.reasoning_summary_text.done", + "response.reasoning_summary_part.done", + "response.output_item.done", + "response.completed", + ] + assert streamed_event_types.count("response.reasoning_summary_text.delta") >= 1 + + delta_events = [ + event + for event in streamed_events + if (event.type.value if hasattr(event.type, "value") else str(event.type)) + == "response.reasoning_summary_text.delta" + ] + text_done_event = streamed_events[-4] + part_done_event = streamed_events[-3] + output_item_done_event = streamed_events[-2] + + assert all(delta_event.summary_index == 0 for delta_event in delta_events) + assert text_done_event.text == "Cached reasoning summary." + assert text_done_event.summary_index == 0 + assert part_done_event.part.type == "summary_text" + assert part_done_event.part.text == "Cached reasoning summary." + assert output_item_done_event.item.type == "reasoning" + assert output_item_done_event.item.summary[0]["text"] == "Cached reasoning summary." + + @pytest.mark.asyncio async def test_responses_api_cache_with_different_inputs(): """ diff --git a/tests/local_testing/test_embedding.py b/tests/local_testing/test_embedding.py index 3f1a397ebcb..82510b6f4fd 100644 --- a/tests/local_testing/test_embedding.py +++ b/tests/local_testing/test_embedding.py @@ -1257,22 +1257,13 @@ def test_jina_ai_img_embeddings(input_data, expected_payload_input): assert sent_data["input"] == expected_payload_input -def test_encoding_format_none_not_omitted_from_openai_sdk(): +def test_encoding_format_defaults_to_float_for_openai_sdk(monkeypatch): """ - Test that encoding_format=None is explicitly sent to OpenAI SDK. + When encoding_format is not provided, LiteLLM sends `float` for OpenAI-path embeddings. - This test verifies that when encoding_format is not provided by the user, - liteLLM explicitly sets it to None rather than omitting it. This prevents - the OpenAI SDK from adding its default value of 'base64'. - - Without this fix: - - OpenAI SDK adds encoding_format='base64' as default when parameter is missing - - This causes issues with providers that don't support encoding_format (like Gemini) - - With this fix: - - encoding_format=None is explicitly passed - - OpenAI SDK respects the explicit None and doesn't add defaults + Optional global override: `LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT`. """ + monkeypatch.delenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", raising=False) with patch( "litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client" ) as mock_get_client: @@ -1310,17 +1301,12 @@ def test_encoding_format_none_not_omitted_from_openai_sdk(): call_kwargs = call_args[1] # Get kwargs - # The key assertion: encoding_format should be in the request with value None - # This prevents OpenAI SDK from adding its default 'base64' value - assert "encoding_format" in call_kwargs, ( - "encoding_format should be explicitly passed to OpenAI SDK " - "(even if None) to prevent SDK from adding default value" - ) + assert "encoding_format" in call_kwargs assert ( - call_kwargs["encoding_format"] is None - ), "encoding_format should be None when not provided by user" + call_kwargs["encoding_format"] == "float" + ), "encoding_format should default to float when not provided by user" - print("✅ PASS: encoding_format=None is correctly passed to OpenAI SDK") + print("✅ PASS: encoding_format='float' is correctly passed to OpenAI SDK") def test_encoding_format_explicit_value_preserved(): diff --git a/tests/local_testing/test_get_llm_provider.py b/tests/local_testing/test_get_llm_provider.py index 010a071f73e..14b9e8cd136 100644 --- a/tests/local_testing/test_get_llm_provider.py +++ b/tests/local_testing/test_get_llm_provider.py @@ -477,3 +477,4 @@ def test_get_llm_provider_use_proxy_arg_true_with_direct_args(): assert provider == "litellm_proxy" assert key == arg_api_key # Should use the argument key assert base == arg_api_base # Should use the argument base + diff --git a/tests/local_testing/test_responses_stream_cache_keys.py b/tests/local_testing/test_responses_stream_cache_keys.py new file mode 100644 index 00000000000..5637028f550 --- /dev/null +++ b/tests/local_testing/test_responses_stream_cache_keys.py @@ -0,0 +1,141 @@ +from datetime import datetime +from unittest.mock import AsyncMock, MagicMock + +import pytest + +import litellm +from litellm import aresponses +from litellm._uuid import uuid +from litellm.caching.caching_handler import LLMCachingHandler +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging +from litellm.types.llms import openai as openai_types +from litellm.types.utils import CallTypes + + +@pytest.mark.asyncio +async def test_async_get_cache_reuses_preset_cache_key_for_responses(): + caching_handler = LLMCachingHandler( + original_function=aresponses, + request_kwargs={}, + start_time=datetime.now(), + ) + logging_obj = LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=CallTypes.aresponses.value, + model="gpt-4.1-mini", + messages=[], + function_id=str(uuid.uuid4()), + stream=True, + start_time=datetime.now(), + ) + + original_cache = litellm.cache + mock_cache = MagicMock() + mock_cache.supported_call_types = [CallTypes.aresponses.value] + mock_cache._supports_async.return_value = True + mock_cache.get_cache_key.return_value = "responses-stream-cache-key" + mock_cache.async_get_cache = AsyncMock(return_value=None) + litellm.cache = mock_cache + + kwargs = { + "model": "gpt-4.1-mini", + "input": "hello", + "stream": True, + "litellm_params": {}, + } + await caching_handler._async_get_cache( + model="gpt-4.1-mini", + original_function=aresponses, + logging_obj=logging_obj, + start_time=datetime.now(), + call_type=CallTypes.aresponses.value, + kwargs=kwargs, + ) + + assert caching_handler.preset_cache_key == "responses-stream-cache-key" + mock_cache.async_get_cache.assert_awaited_once() + assert ( + mock_cache.async_get_cache.call_args.kwargs["cache_key"] + == "responses-stream-cache-key" + ) + + litellm.cache = original_cache + + +@pytest.mark.asyncio +async def test_async_get_cache_falls_back_to_sync_cache_for_responses(): + caching_handler = LLMCachingHandler( + original_function=aresponses, + request_kwargs={}, + start_time=datetime.now(), + ) + logging_obj = LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=CallTypes.aresponses.value, + model="gpt-4.1-mini", + messages=[], + function_id=str(uuid.uuid4()), + stream=True, + start_time=datetime.now(), + ) + + original_cache = litellm.cache + mock_cache = MagicMock() + mock_cache.supported_call_types = [CallTypes.aresponses.value] + mock_cache._supports_async.return_value = False + mock_cache.get_cache_key.return_value = "responses-stream-cache-key" + mock_cache.get_cache.return_value = None + litellm.cache = mock_cache + + kwargs = { + "model": "gpt-4.1-mini", + "input": "hello", + "stream": True, + "litellm_params": {}, + } + await caching_handler._async_get_cache( + model="gpt-4.1-mini", + original_function=aresponses, + logging_obj=logging_obj, + start_time=datetime.now(), + call_type=CallTypes.aresponses.value, + kwargs=kwargs, + ) + + assert caching_handler.preset_cache_key == "responses-stream-cache-key" + mock_cache.get_cache.assert_called_once() + assert mock_cache.get_cache.call_args.kwargs["cache_key"] == ( + "responses-stream-cache-key" + ) + + litellm.cache = original_cache + + +def test_reasoning_summary_events_default_summary_index(): + delta_event = openai_types.ReasoningSummaryTextDeltaEvent( + type=openai_types.ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DELTA, + item_id="rs_1", + output_index=0, + delta="abc", + ) + text_done_event = openai_types.ReasoningSummaryTextDoneEvent( + type=openai_types.ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DONE, + item_id="rs_1", + output_index=0, + sequence_number=1, + text="abc", + ) + part_done_event = openai_types.ReasoningSummaryPartDoneEvent( + type=openai_types.ResponsesAPIStreamEvents.REASONING_SUMMARY_PART_DONE, + item_id="rs_1", + output_index=0, + sequence_number=2, + part=openai_types.BaseLiteLLMOpenAIResponseObject( + type="summary_text", + text="abc", + ), + ) + + assert delta_event.summary_index == 0 + assert text_done_event.summary_index == 0 + assert part_done_event.summary_index == 0 diff --git a/tests/pass_through_unit_tests/test_assemblyai_unit_tests_passthrough.py b/tests/pass_through_unit_tests/test_assemblyai_unit_tests_passthrough.py index 963f1ad6ef9..67bc4423d8c 100644 --- a/tests/pass_through_unit_tests/test_assemblyai_unit_tests_passthrough.py +++ b/tests/pass_through_unit_tests/test_assemblyai_unit_tests_passthrough.py @@ -134,3 +134,62 @@ def test_is_assemblyai_route(): == False ) assert handler.is_assemblyai_route("") == False + + +# --- Security: SSRF via transcript_id path traversal --- + + +def test_get_assembly_transcript_rejects_slash_in_id(assembly_handler): + with patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials", + return_value="test-key", + ): + with pytest.raises(ValueError, match="disallowed characters"): + assembly_handler._get_assembly_transcript("../../admin/credentials") + + +def test_get_assembly_transcript_rejects_dotdot_in_id(assembly_handler): + with patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials", + return_value="test-key", + ): + with pytest.raises(ValueError, match="disallowed characters"): + assembly_handler._get_assembly_transcript("..evil") + + +def test_get_assembly_transcript_rejects_fragment_in_id(assembly_handler): + with patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials", + return_value="test-key", + ): + with pytest.raises(ValueError, match="disallowed characters"): + assembly_handler._get_assembly_transcript("abc#suffix") + + +def test_get_assembly_transcript_rejects_query_in_id(assembly_handler): + with patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials", + return_value="test-key", + ): + with pytest.raises(ValueError, match="disallowed characters"): + assembly_handler._get_assembly_transcript("abc?x=1") + + +def test_get_assembly_transcript_allows_valid_id( + assembly_handler, mock_transcript_response +): + with patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials", + return_value="test-key", + ): + with patch("httpx.get") as mock_get: + mock_get.return_value.json.return_value = mock_transcript_response + mock_get.return_value.raise_for_status.return_value = None + + transcript = assembly_handler._get_assembly_transcript( + "abc123-valid-id_xyz" + ) + assert transcript == mock_transcript_response + called_url = mock_get.call_args[0][0] + assert "abc123-valid-id_xyz" in called_url + assert ".." not in called_url diff --git a/tests/proxy_unit_tests/test_auth_checks.py b/tests/proxy_unit_tests/test_auth_checks.py index 86cd5c0c413..5636a55c95a 100644 --- a/tests/proxy_unit_tests/test_auth_checks.py +++ b/tests/proxy_unit_tests/test_auth_checks.py @@ -16,6 +16,7 @@ import httpx from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_checks import get_end_user_object from litellm.caching.caching import DualCache +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy._types import ( LiteLLM_EndUserTable, LiteLLM_BudgetTable, @@ -48,9 +49,15 @@ async def test_get_end_user_object(customer_spend, customer_budget): litellm_budget_table=_budget, blocked=False, ) - _cache = DualCache() + # UserApiKeyCache applies model_type on get/set; plain DualCache returns raw dicts + # and breaks get_end_user_object's typed async_get_cache path. + _cache = UserApiKeyCache() _key = "end_user_id:{}".format(end_user_id) - _cache.set_cache(key=_key, value=end_user_obj.model_dump()) + await _cache.async_set_cache( + key=_key, + value=end_user_obj, + model_type=LiteLLM_EndUserTable, + ) try: await get_end_user_object( end_user_id=end_user_id, diff --git a/tests/proxy_unit_tests/test_proxy_server.py b/tests/proxy_unit_tests/test_proxy_server.py index 86bbc5170ec..cdcdc89e7f0 100644 --- a/tests/proxy_unit_tests/test_proxy_server.py +++ b/tests/proxy_unit_tests/test_proxy_server.py @@ -2768,40 +2768,40 @@ async def test_update_config_success_callback_normalization(): import litellm.proxy.proxy_server as proxy_server from litellm.proxy._types import ConfigYAML - # Ensure feature is enabled and prisma_client is set - setattr(proxy_server, "store_model_in_db", True) setattr(proxy_server, "proxy_logging_obj", MagicMock()) + existing_litellm_settings = {"success_callback": ["langfuse"]} + + class FakeRow: + def __init__(self, name, value): + self.param_name = name + self.param_value = value + + upserted = {} + + async def fake_find_first(where=None): + if where and where.get("param_name") == "litellm_settings": + return FakeRow("litellm_settings", existing_litellm_settings) + return None + + async def fake_upsert(where=None, data=None): + upserted[where["param_name"]] = json.loads(data["update"]["param_value"]) + class MockPrisma: def __init__(self): self.db = MagicMock() self.db.litellm_config = MagicMock() - self.db.litellm_config.upsert = AsyncMock() - - # proxy_server.update_config expects this to be sync returning a dict - def jsonify_object(self, obj): - return obj + self.db.litellm_config.find_first = AsyncMock(side_effect=fake_find_first) + self.db.litellm_config.upsert = AsyncMock(side_effect=fake_upsert) setattr(proxy_server, "prisma_client", MockPrisma()) class MockProxyConfig: - def __init__(self): - self.saved_config = None - - async def get_config(self): - # Existing config has one lowercase callback already - return {"litellm_settings": {"success_callback": ["langfuse"]}} - - async def save_config(self, new_config: dict): - self.saved_config = new_config - async def add_deployment(self, prisma_client=None, proxy_logging_obj=None): return None - mock_proxy_config = MockProxyConfig() - setattr(proxy_server, "proxy_config", mock_proxy_config) + setattr(proxy_server, "proxy_config", MockProxyConfig()) - # Update config with mixed-case callbacks - expect normalization to lowercase config_update = ConfigYAML(litellm_settings={"success_callback": ["SQS", "sQs"]}) from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth @@ -2810,9 +2810,10 @@ async def test_update_config_success_callback_normalization(): ) await proxy_server.update_config(config_update, user_api_key_dict=admin_user) - saved = mock_proxy_config.saved_config - assert saved is not None, "save_config was not called" - callbacks = saved["litellm_settings"]["success_callback"] + assert ( + "litellm_settings" in upserted + ), "litellm_config.upsert was not called for litellm_settings" + callbacks = upserted["litellm_settings"]["success_callback"] # Deduped and normalized assert "sqs" in callbacks diff --git a/tests/proxy_unit_tests/test_user_api_key_auth.py b/tests/proxy_unit_tests/test_user_api_key_auth.py index e51f81561aa..543cabb6b4c 100644 --- a/tests/proxy_unit_tests/test_user_api_key_auth.py +++ b/tests/proxy_unit_tests/test_user_api_key_auth.py @@ -268,7 +268,12 @@ async def test_aaauser_personal_budgets(key_ownership): test_user_cache = getattr(litellm.proxy.proxy_server, "user_api_key_cache") - assert test_user_cache.get_cache(key=hash_token(user_key)) == valid_token + assert ( + test_user_cache.get_cache( + key=hash_token(user_key), model_type=UserAPIKeyAuth + ) + == valid_token + ) try: await user_api_key_auth(request=request, api_key="Bearer " + user_key) diff --git a/tests/router_unit_tests/test_router_endpoints.py b/tests/router_unit_tests/test_router_endpoints.py index b93502e8152..0ce2dec9b56 100644 --- a/tests/router_unit_tests/test_router_endpoints.py +++ b/tests/router_unit_tests/test_router_endpoints.py @@ -1110,7 +1110,7 @@ def test_initialize_skills_endpoints(): async def test_init_containers_api_endpoints(): """ Test that _init_containers_api_endpoints calls the original function - directly without model-based routing. + directly when there is no managed container ID (no embedded model_id). """ router = Router(model_list=[]) @@ -1127,3 +1127,112 @@ async def test_init_containers_api_endpoints(): custom_llm_provider="openai", name="Test Container" ) assert result == mock_response + + +@pytest.mark.asyncio +async def test_init_containers_api_endpoints_managed_id_routes_via_generic_fallbacks(): + """ + Managed ``cntr_`` IDs embed ``model_id``; router should decode and use + ``_ageneric_api_call_with_fallbacks`` so deployment credentials apply. + """ + from litellm.responses.utils import ResponsesAPIRequestUtils + + router = Router( + model_list=[ + { + "model_name": "azure-router-model", + "litellm_params": { + "model": "azure/gpt-4", + "api_key": "fake-key", + "api_base": "https://westus.api.cognitive.microsoft.com", + }, + } + ] + ) + router._ageneric_api_call_with_fallbacks = AsyncMock() + + managed_id = ResponsesAPIRequestUtils._build_container_id( + custom_llm_provider="azure", + model_id="azure-router-model", + container_id="cfile_upstream_abc", + ) + + await router._init_containers_api_endpoints( + original_function=AsyncMock(), + custom_llm_provider="openai", + container_id=managed_id, + file_id="cfile_xyz", + ) + + router._ageneric_api_call_with_fallbacks.assert_called_once() + call_kw = router._ageneric_api_call_with_fallbacks.call_args.kwargs + assert call_kw["model"] == "azure-router-model" + assert call_kw["container_id"] == "cfile_upstream_abc" + assert call_kw["file_id"] == "cfile_xyz" + assert call_kw["custom_llm_provider"] == "azure" + + +@pytest.mark.asyncio +async def test_init_containers_api_endpoints_managed_id_without_model_id_unwraps(): + """ + Managed ``cntr_`` IDs may be encoded with an empty ``model_id`` (e.g. when a + streaming response had no router metadata). The router must still unwrap the + managed ID before calling the upstream provider — otherwise the raw + ``cntr_...`` token leaks downstream and the provider rejects it. + """ + from litellm.responses.utils import ResponsesAPIRequestUtils + + router = Router(model_list=[]) + mock_original_function = AsyncMock(return_value={"ok": True}) + + managed_id = ResponsesAPIRequestUtils._build_container_id( + custom_llm_provider="openai", + model_id=None, + container_id="cfile_upstream_abc", + ) + + await router._init_containers_api_endpoints( + original_function=mock_original_function, + custom_llm_provider="openai", + container_id=managed_id, + file_id="cfile_xyz", + ) + + mock_original_function.assert_called_once() + call_kw = mock_original_function.call_args.kwargs + assert call_kw["container_id"] == "cfile_upstream_abc" + assert call_kw["file_id"] == "cfile_xyz" + assert call_kw["custom_llm_provider"] == "openai" + + +@pytest.mark.asyncio +async def test_init_containers_api_endpoints_managed_id_without_model_id_applies_decoded_provider(): + """ + A managed ``cntr_`` ID can encode a non-OpenAI provider (e.g. ``azure``) with + an empty ``model_id`` (streaming events without router ``model_info.id``). + The router must still apply the decoded provider so the request routes to + the correct upstream — not stay on the default ``openai``. + """ + from litellm.responses.utils import ResponsesAPIRequestUtils + + router = Router(model_list=[]) + mock_original_function = AsyncMock(return_value={"ok": True}) + + managed_id = ResponsesAPIRequestUtils._build_container_id( + custom_llm_provider="azure", + model_id=None, + container_id="cfile_upstream_abc", + ) + + await router._init_containers_api_endpoints( + original_function=mock_original_function, + custom_llm_provider="openai", + container_id=managed_id, + file_id="cfile_xyz", + ) + + mock_original_function.assert_called_once() + call_kw = mock_original_function.call_args.kwargs + assert call_kw["container_id"] == "cfile_upstream_abc" + assert call_kw["file_id"] == "cfile_xyz" + assert call_kw["custom_llm_provider"] == "azure" diff --git a/tests/test_litellm/caching/test_dual_cache.py b/tests/test_litellm/caching/test_dual_cache.py index 8e502175761..64774726201 100644 --- a/tests/test_litellm/caching/test_dual_cache.py +++ b/tests/test_litellm/caching/test_dual_cache.py @@ -1,5 +1,6 @@ import asyncio import time +import uuid from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -260,3 +261,72 @@ async def test_async_increment_cache_returns_none_when_no_in_memory_cache_and_re f"Expected None when in_memory_cache is absent and Redis fails, got {result!r}. " "Returning the delta (1.0) would silently miscalculate rate-limit counters." ) + + +def test_dual_cache_late_attach_redis_wires_writes_and_ttl_sync(): + """ + Typical lazy startup (sync): DualCache runs with in-memory only, then Redis + becomes available and is attached. New writes must reach Redis; keys written + before attach are not backfilled. Optional default_redis_ttl is applied on attach. + """ + in_memory = InMemoryCache() + dual_cache = DualCache(in_memory_cache=in_memory, redis_cache=None) + + mock_redis = MagicMock() + mock_redis.set_cache = MagicMock() + mock_redis.async_set_cache = AsyncMock() + + key_before = f"before_attach_{uuid.uuid4()}" + val_before = {"phase": "memory_only"} + dual_cache.set_cache(key_before, val_before) + + assert in_memory.get_cache(key_before) == val_before + + dual_cache.attach_redis_cache(mock_redis, default_redis_ttl=99.0) + assert dual_cache.redis_cache is mock_redis + assert dual_cache.default_redis_ttl == 99.0 + + mock_redis.set_cache.assert_not_called() + + key_after = f"after_attach_{uuid.uuid4()}" + val_after = {"phase": "memory_and_redis"} + dual_cache.set_cache(key_after, val_after) + mock_redis.set_cache.assert_called_once() + assert mock_redis.set_cache.call_args[0][:2] == (key_after, val_after) + + assert in_memory.get_cache(key_after) == val_after + + +@pytest.mark.asyncio +async def test_dual_cache_late_attach_redis_wires_writes_and_ttl_async(): + """ + Typical lazy startup (async): DualCache runs with in-memory only, then Redis + becomes available and is attached. New writes must reach Redis; keys written + before attach are not backfilled. Optional default_redis_ttl is applied on attach. + """ + in_memory = InMemoryCache() + dual_cache = DualCache(in_memory_cache=in_memory, redis_cache=None) + + mock_redis = MagicMock() + mock_redis.set_cache = MagicMock() + mock_redis.async_set_cache = AsyncMock() + + key_before = f"before_attach_{uuid.uuid4()}" + val_before = {"phase": "memory_only"} + await dual_cache.async_set_cache(key_before, val_before) + + assert in_memory.get_cache(key_before) == val_before + + dual_cache.attach_redis_cache(mock_redis, default_redis_ttl=99.0) + assert dual_cache.redis_cache is mock_redis + assert dual_cache.default_redis_ttl == 99.0 + + mock_redis.async_set_cache.assert_not_called() + + key_after = f"after_attach_{uuid.uuid4()}" + val_after = {"phase": "memory_and_redis"} + await dual_cache.async_set_cache(key_after, val_after) + mock_redis.async_set_cache.assert_called_once() + assert mock_redis.async_set_cache.call_args[0][:2] == (key_after, val_after) + + assert in_memory.get_cache(key_after) == val_after diff --git a/tests/test_litellm/containers/test_azure_container_transformation.py b/tests/test_litellm/containers/test_azure_container_transformation.py index a46046b318b..32caaaffef5 100644 --- a/tests/test_litellm/containers/test_azure_container_transformation.py +++ b/tests/test_litellm/containers/test_azure_container_transformation.py @@ -11,6 +11,7 @@ sys.path.insert(0, os.path.abspath("../../../")) import litellm from litellm.llms.azure.containers.transformation import AzureContainerConfig from litellm.llms.base_llm.containers.transformation import BaseContainerConfig +from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.types.containers.main import ( ContainerFileListResponse, ContainerListResponse, @@ -323,6 +324,29 @@ class TestAzureContainerConfig: assert url_fc == expected_fc assert url_fc.index("/content") < url_fc.index("?") + def test_transform_requests_encode_path_ids_before_query_string(self): + from litellm.types.router import GenericLiteLLMParams + + api_base = ( + "https://my-resource.openai.azure.com/openai/v1/containers" + "?api-version=v1" + ) + + url, _ = self.config.transform_container_file_content_request( + container_id="../../other", + file_id="file?download=1#frag", + api_base=api_base, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + expected_url = ( + "https://my-resource.openai.azure.com/openai/v1/containers/" + "..%2F..%2Fother/files/file%3Fdownload%3D1%23frag/content" + "?api-version=v1" + ) + assert url == expected_url + def test_provider_config_manager_returns_azure_config(self): from litellm.types.utils import LlmProviders from litellm.utils import ProviderConfigManager @@ -518,3 +542,206 @@ class TestAzureContainerKnownFailureRegressions: c2 = _get_container_provider_config("azure_text") assert type(c1) is type(c2) assert isinstance(c1, AzureContainerConfig) + + @pytest.mark.asyncio + async def test_proxy_process_request_preserves_managed_container_id( + self, monkeypatch + ): + from starlette.requests import Request + + from litellm.proxy.container_endpoints import handler_factory + + encoded_id = ResponsesAPIRequestUtils._build_container_id( + custom_llm_provider="azure", + model_id="model_abc123", + container_id="cntr_123", + ) + captured = {} + + async def _mock_base_process_llm_request( + self, + request, + fastapi_response, + user_api_key_dict, + route_type, + **kwargs, + ): + captured["data"] = self.data + captured["route_type"] = route_type + return {"id": "cfile_abc"} + + from litellm.proxy.common_request_processing import ( + ProxyBaseLLMRequestProcessing, + ) + + monkeypatch.setattr( + ProxyBaseLLMRequestProcessing, + "base_process_llm_request", + _mock_base_process_llm_request, + ) + + request = Request( + { + "type": "http", + "method": "GET", + "path": "/v1/containers/id/files/id/content", + "headers": [], + "query_string": b"", + } + ) + fastapi_response = MagicMock() + + await handler_factory._process_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=MagicMock(), + route_type="alist_container_files", + path_params={"container_id": encoded_id}, + ) + + assert captured["route_type"] == "alist_container_files" + assert captured["data"]["container_id"] == encoded_id + assert captured["data"]["custom_llm_provider"] == "openai" + assert "model_id" not in captured["data"] + assert "api_base" not in captured["data"] + + @pytest.mark.asyncio + async def test_regression_binary_file_request_routes_through_proxy_processor( + self, monkeypatch + ): + from fastapi import Response + from starlette.requests import Request + + from litellm.proxy.container_endpoints import handler_factory + + encoded_id = ResponsesAPIRequestUtils._build_container_id( + custom_llm_provider="azure", + model_id="model_abc123", + container_id="cntr_123", + ) + captured = {} + + async def _mock_base_process_llm_request( + self, + request, + fastapi_response, + user_api_key_dict, + route_type, + **kwargs, + ): + captured["data"] = self.data + captured["route_type"] = route_type + fastapi_response.headers["x-litellm-call-id"] = "call-123" + return b"csv-bytes" + + from litellm.proxy.common_request_processing import ( + ProxyBaseLLMRequestProcessing, + ) + + monkeypatch.setattr( + ProxyBaseLLMRequestProcessing, + "base_process_llm_request", + _mock_base_process_llm_request, + ) + + request = Request( + { + "type": "http", + "method": "GET", + "path": "/v1/containers/id/files/id/content", + "headers": [], + "query_string": b"", + } + ) + fastapi_response = Response() + + response = await handler_factory._process_binary_request( + request=request, + fastapi_response=fastapi_response, + container_id=encoded_id, + file_id="cfile_abc", + user_api_key_dict=MagicMock(), + ) + + assert captured["route_type"] == "aretrieve_container_file_content" + assert captured["data"]["container_id"] == encoded_id + assert captured["data"]["file_id"] == "cfile_abc" + assert captured["data"]["custom_llm_provider"] == "openai" + assert response.status_code == 200 + assert response.body == b"csv-bytes" + assert response.headers["x-litellm-call-id"] == "call-123" + + @pytest.mark.asyncio + async def test_regression_multipart_upload_request_uses_provider_from_managed_id( + self, monkeypatch + ): + from starlette.requests import Request + + from litellm.proxy.common_request_processing import ( + ProxyBaseLLMRequestProcessing, + ) + from litellm.proxy.common_utils import http_parsing_utils + from litellm.proxy.container_endpoints import handler_factory + + encoded_id = ResponsesAPIRequestUtils._build_container_id( + custom_llm_provider="azure", + model_id="model_abc123", + container_id="cntr_123", + ) + captured = {} + + async def _mock_get_form_data(request): + return {"file": "ignored"} + + async def _mock_convert_upload_files_to_file_data(form_data): + return {"file": [("data.csv", b"csv-bytes", "text/csv")]} + + async def _mock_base_process_llm_request( + self, + request, + fastapi_response, + user_api_key_dict, + route_type, + **kwargs, + ): + captured["data"] = self.data + captured["route_type"] = route_type + return {"id": "cfile_abc"} + + monkeypatch.setattr( + http_parsing_utils, + "get_form_data", + _mock_get_form_data, + ) + monkeypatch.setattr( + http_parsing_utils, + "convert_upload_files_to_file_data", + _mock_convert_upload_files_to_file_data, + ) + monkeypatch.setattr( + ProxyBaseLLMRequestProcessing, + "base_process_llm_request", + _mock_base_process_llm_request, + ) + + request = Request( + { + "type": "http", + "method": "POST", + "path": "/v1/containers/id/files", + "headers": [], + "query_string": b"", + } + ) + + await handler_factory._process_multipart_upload_request( + request=request, + fastapi_response=MagicMock(), + user_api_key_dict=MagicMock(), + route_type="aupload_container_file", + container_id=encoded_id, + ) + + assert captured["route_type"] == "aupload_container_file" + assert captured["data"]["container_id"] == encoded_id + assert captured["data"]["custom_llm_provider"] == "openai" diff --git a/tests/test_litellm/containers/test_container_handler_url.py b/tests/test_litellm/containers/test_container_handler_url.py new file mode 100644 index 00000000000..19275800680 --- /dev/null +++ b/tests/test_litellm/containers/test_container_handler_url.py @@ -0,0 +1,28 @@ +import pytest + +from litellm.llms.custom_httpx.container_handler import _build_url + + +def test_build_url_encodes_path_params_and_preserves_query(): + url = _build_url( + api_base="https://example.com/v1/containers?api-version=v1", + path_template="/containers/{container_id}/files/{file_id}/content", + path_params={ + "container_id": "../../containers/other", + "file_id": "file?download=1#frag", + }, + ) + + assert ( + url + == "https://example.com/v1/containers/..%2F..%2Fcontainers%2Fother/files/file%3Fdownload%3D1%23frag/content?api-version=v1" + ) + + +def test_build_url_rejects_dot_segment_path_param(): + with pytest.raises(ValueError, match="container_id cannot be a dot path segment"): + _build_url( + api_base="https://example.com/v1/containers", + path_template="/containers/{container_id}", + path_params={"container_id": ".."}, + ) diff --git a/tests/test_litellm/containers/test_container_transformation.py b/tests/test_litellm/containers/test_container_transformation.py index 47b5b8dc56f..555fe7773f0 100644 --- a/tests/test_litellm/containers/test_container_transformation.py +++ b/tests/test_litellm/containers/test_container_transformation.py @@ -230,6 +230,23 @@ class TestOpenAIContainerTransformation: assert url == f"{api_base}/{container_id}" assert params == {} # No query params for retrieve + def test_transform_container_retrieve_request_encodes_path_traversal(self): + """Test container IDs are treated as a single upstream path segment.""" + api_base = "https://api.openai.com/v1/containers" + + url, params = self.config.transform_container_retrieve_request( + container_id="../../vector_stores?x=1#frag", + api_base=api_base, + litellm_params={}, + headers={}, + ) + + assert ( + url + == "https://api.openai.com/v1/containers/..%2F..%2Fvector_stores%3Fx%3D1%23frag" + ) + assert params == {} + def test_transform_container_retrieve_response(self): """Test container retrieve response transformation.""" # Mock HTTP response diff --git a/tests/test_litellm/integrations/arize/test_arize_phoenix.py b/tests/test_litellm/integrations/arize/test_arize_phoenix.py index 01f85af2620..4a2eab29e8e 100644 --- a/tests/test_litellm/integrations/arize/test_arize_phoenix.py +++ b/tests/test_litellm/integrations/arize/test_arize_phoenix.py @@ -280,3 +280,55 @@ class TestDynamicProjectNameOnSpan: if __name__ == "__main__": unittest.main() + + +# --- Security: SSRF via prompt_version_id path traversal --- + + +def test_arize_phoenix_client_sanitize_id_rejects_traversal(): + from litellm.integrations.arize.arize_phoenix_client import _sanitize_id + + # dotdot without slashes + with pytest.raises(ValueError, match="path traversal"): + _sanitize_id("..something") + # full traversal (slash caught first) + with pytest.raises(ValueError, match="disallowed characters"): + _sanitize_id("../../projects") + + +def test_arize_phoenix_client_sanitize_id_rejects_slash(): + from litellm.integrations.arize.arize_phoenix_client import _sanitize_id + + with pytest.raises(ValueError, match="disallowed characters"): + _sanitize_id("valid/extra") + + +def test_arize_phoenix_client_sanitize_id_rejects_fragment(): + from litellm.integrations.arize.arize_phoenix_client import _sanitize_id + + with pytest.raises(ValueError, match="disallowed characters"): + _sanitize_id("abc#suffix") + + +def test_arize_phoenix_client_sanitize_id_rejects_query(): + from litellm.integrations.arize.arize_phoenix_client import _sanitize_id + + with pytest.raises(ValueError, match="disallowed characters"): + _sanitize_id("abc?x=1") + + +def test_arize_phoenix_client_sanitize_id_allows_uuid(): + from litellm.integrations.arize.arize_phoenix_client import _sanitize_id + + uid = "550e8400-e29b-41d4-a716-446655440000" + assert _sanitize_id(uid) == uid + + +def test_arize_phoenix_client_get_prompt_version_rejects_traversal(): + from litellm.integrations.arize.arize_phoenix_client import ArizePhoenixClient + + client = ArizePhoenixClient( + api_key="test-key", api_base="https://app.phoenix.arize.com" + ) + with pytest.raises(ValueError, match="disallowed characters"): + client.get_prompt_version("../../projects") diff --git a/tests/test_litellm/integrations/bitbucket/test_bitbucket_integration.py b/tests/test_litellm/integrations/bitbucket/test_bitbucket_integration.py index a7b2d362ed2..46cd1d6e765 100644 --- a/tests/test_litellm/integrations/bitbucket/test_bitbucket_integration.py +++ b/tests/test_litellm/integrations/bitbucket/test_bitbucket_integration.py @@ -11,6 +11,7 @@ sys.path.insert( import litellm from litellm.integrations.bitbucket import BitBucketPromptManager +from litellm.integrations.bitbucket.bitbucket_client import _sanitize_file_path @patch("litellm.integrations.bitbucket.bitbucket_prompt_manager.BitBucketClient") @@ -370,3 +371,45 @@ def test_bitbucket_prompt_manager_list_templates(mock_client_class): templates = manager.prompt_manager.list_templates() assert isinstance(templates, list) assert "test_prompt" in templates + + +# --- Security: path traversal / SSRF --- + + +def test_sanitize_file_path_rejects_traversal(): + with pytest.raises(ValueError, match="path traversal"): + _sanitize_file_path("../../etc/passwd") + + +def test_sanitize_file_path_rejects_fragment(): + with pytest.raises(ValueError, match="URL special characters"): + _sanitize_file_path("secret#.prompt") + + +def test_sanitize_file_path_rejects_query(): + with pytest.raises(ValueError, match="URL special characters"): + _sanitize_file_path("secret?.prompt") + + +def test_sanitize_file_path_encodes_special_chars(): + result = _sanitize_file_path("prompts/my prompt.prompt") + assert result == "prompts/my%20prompt.prompt" + + +def test_sanitize_file_path_allows_normal_paths(): + assert _sanitize_file_path("prompts/my-prompt") == "prompts/my-prompt" + assert _sanitize_file_path("simple") == "simple" + + +def test_bitbucket_client_rejects_traversal_in_get_file_content(): + from litellm.integrations.bitbucket.bitbucket_client import BitBucketClient + + client = BitBucketClient( + { + "workspace": "ws", + "repository": "repo", + "access_token": "tok", + } + ) + with pytest.raises(ValueError, match="path traversal"): + client.get_file_content("../../admin/credentials") diff --git a/tests/test_litellm/interactions/test_gemini_interactions_transformation.py b/tests/test_litellm/interactions/test_gemini_interactions_transformation.py index 37dc491c26a..758ff3ea38e 100644 --- a/tests/test_litellm/interactions/test_gemini_interactions_transformation.py +++ b/tests/test_litellm/interactions/test_gemini_interactions_transformation.py @@ -147,6 +147,21 @@ class TestInteractionOperationUrls: assert "secret-key" not in url assert expected_suffix in url + def test_interaction_id_is_encoded_as_one_path_segment(self, config): + with patch(_PATCH_GET_API_KEY, return_value="secret-key"): + url, params = config.transform_cancel_interaction_request( + interaction_id="../../interactions/other?x=1#frag", + api_base="https://generativelanguage.googleapis.com", + litellm_params=GenericLiteLLMParams(api_key="secret-key"), + headers={}, + ) + + assert ( + url + == "https://generativelanguage.googleapis.com/v1beta/interactions/..%2F..%2Finteractions%2Fother%3Fx%3D1%23frag:cancel" + ) + assert params == {} + def test_get_interaction_raises_without_key(self, config): with patch(_PATCH_GET_API_KEY, return_value=None): with pytest.raises(ValueError, match="Google API key is required"): diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index d424cd8599f..27a3ddb553d 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -9,8 +9,10 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( BAD_MESSAGE_ERROR_STR, BedrockConverseMessagesProcessor, BedrockImageProcessor, - anthropic_messages_pt, + _bedrock_converse_messages_pt, _convert_to_bedrock_tool_call_invoke, + _convert_to_bedrock_tool_call_result, + anthropic_messages_pt, convert_to_gemini_tool_call_result, ollama_pt, sanitize_messages_for_tool_calling, @@ -2485,10 +2487,6 @@ def test_convert_to_anthropic_tool_result_openai_file_pdf_becomes_document(): inside the tool_result content. Reuses anthropic_process_openai_file_message, which already handles this for user messages. """ - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_anthropic_tool_result, - ) - pdf_b64 = "JVBERi0xLjQKJeLjz9MK" message = { "tool_call_id": "toolu_pdf_1", @@ -2505,157 +2503,105 @@ def test_convert_to_anthropic_tool_result_openai_file_pdf_becomes_document(): ], } - result = convert_to_anthropic_tool_result(message) + result = _convert_to_bedrock_tool_call_result(message) - assert result["type"] == "tool_result" - assert result["tool_use_id"] == "toolu_pdf_1" - content = result["content"] - assert isinstance(content, list) and len(content) == 1 - block = content[0] - assert block["type"] == "document" - assert block["source"]["type"] == "base64" - assert block["source"]["media_type"] == "application/pdf" - assert block["source"]["data"] == pdf_b64 + tool_result = result["toolResult"] + assert len(tool_result["content"]) == 1 + assert "document" in tool_result["content"][0] + assert tool_result["content"][0]["document"]["format"] == "pdf" + assert tool_result["content"][0]["document"]["source"]["bytes"] == pdf_b64 -def test_convert_to_anthropic_tool_result_image_url_pdf_data_uri_becomes_document(): - """ - Regression: a PDF sent as an `image_url` data URI on the tool-result path - must translate to an Anthropic document block (not an image block — Anthropic - rejects image blocks whose media_type is a non-image like application/pdf). - """ - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_anthropic_tool_result, - ) +def test_bedrock_converse_messages_pt_document_various_formats(): + """Test that various document media types produce the correct format value.""" + test_cases = [ + ("application/pdf", "pdf"), + ("text/csv", "csv"), + ("text/html", "html"), + ("text/plain", "txt"), + ("text/markdown", "md"), + ( + "application/vnd.openxmlformats-officedocument.wordprocessingml.document", + "docx", + ), + ] - pdf_b64 = "JVBERi0xLjQKJeLjz9MK" - message = { - "tool_call_id": "toolu_pdf_img_1", - "role": "tool", - "name": "fetch_document", - "content": [ + for media_type, expected_format in test_cases: + messages = [ { - "type": "image_url", - "image_url": { - "url": f"data:application/pdf;base64,{pdf_b64}", + "role": "user", + "content": [ + { + "type": "document", + "source": { + "type": "base64", + "media_type": media_type, + "data": "dGVzdA==", + }, + }, + ], + } + ] + + result = _bedrock_converse_messages_pt( + messages, "anthropic.claude-sonnet-4-6", "bedrock" + ) + + doc_block = result[0]["content"][0] + assert doc_block["document"]["format"] == expected_format, ( + f"Expected format '{expected_format}' for media_type '{media_type}', " + f"got '{doc_block['document']['format']}'" + ) + + +def test_bedrock_converse_messages_pt_document_deterministic_name(): + """Test that the same document data always produces the same name.""" + messages = [ + { + "role": "user", + "content": [ + { + "type": "document", + "source": { + "type": "base64", + "media_type": "application/pdf", + "data": "dGVzdA==", + }, }, - }, - ], - } + ], + } + ] - result = convert_to_anthropic_tool_result(message) - - content = result["content"] - assert isinstance(content, list) and len(content) == 1 - block = content[0] - assert block["type"] == "document" - assert block["source"]["media_type"] == "application/pdf" - assert block["source"]["data"] == pdf_b64 - - -def test_convert_to_anthropic_tool_result_image_url_unsupported_mime_stays_image_path(): - """ - An `image_url` data URI whose mime is neither application/pdf nor text/plain - (e.g. application/json) must NOT be routed through the document path. Anthropic - only accepts application/pdf and text/plain as base64 document media_types — - anything else would produce a document block the API rejects. The old - (pre-fix) behavior was to wrap such data as an image block, which also - fails but stays on the image code path; preserve that failure mode rather - than switching to a document path that is equally broken. - """ - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_anthropic_tool_result, + result1 = _bedrock_converse_messages_pt( + messages, "anthropic.claude-sonnet-4-6", "bedrock" + ) + result2 = _bedrock_converse_messages_pt( + messages, "anthropic.claude-sonnet-4-6", "bedrock" ) - message = { - "tool_call_id": "toolu_json_1", - "role": "tool", - "name": "fetch_json", - "content": [ - { - "type": "image_url", - "image_url": { - "url": "data:application/json;base64,eyJrIjoidiJ9", + name1 = result1[0]["content"][0]["document"]["name"] + name2 = result2[0]["content"][0]["document"]["name"] + assert name1 == name2 + + +def test_bedrock_converse_messages_pt_document_rejects_url_source(): + """Test that a URL-type document source raises a clear error instead of KeyError.""" + messages = [ + { + "role": "user", + "content": [ + { + "type": "document", + "source": { + "type": "url", + "url": "https://example.com/doc.pdf", + }, }, - }, - ], - } + ], + } + ] - result = convert_to_anthropic_tool_result(message) - - content = result["content"] - assert isinstance(content, list) and len(content) == 1 - block = content[0] - assert block["type"] == "image", ( - f"unsupported mime {block.get('source', {}).get('media_type')!r} " - f"should not be routed to document path; got {block}" - ) - - -def test_convert_to_anthropic_tool_result_image_url_text_plain_data_uri_becomes_document(): - """ - text/plain is one of the two mimes Anthropic accepts as a base64 document - media_type. Confirm it routes through the document path so tightening the - gate to {application/pdf, text/plain} (not "application/*") covers both. - """ - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_anthropic_tool_result, - ) - - txt_b64 = "aGVsbG8=" # "hello" - message = { - "tool_call_id": "toolu_txt_1", - "role": "tool", - "name": "fetch_text", - "content": [ - { - "type": "image_url", - "image_url": { - "url": f"data:text/plain;base64,{txt_b64}", - }, - }, - ], - } - - result = convert_to_anthropic_tool_result(message) - - content = result["content"] - assert isinstance(content, list) and len(content) == 1 - block = content[0] - assert block["type"] == "document" - assert block["source"]["media_type"] == "text/plain" - assert block["source"]["data"] == txt_b64 - - -def test_convert_to_anthropic_tool_result_image_url_png_still_becomes_image(): - """ - Regression: image_url with a real image mime type must continue to translate - to an Anthropic image block. Locks in existing behavior after the - data-URI-mime-type branching for PDFs. - """ - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_anthropic_tool_result, - ) - - png_b64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR4nGNgYGBgAAAABQABXvMqOgAAAABJRU5ErkJggg==" - message = { - "tool_call_id": "toolu_png_1", - "role": "tool", - "name": "fetch_image", - "content": [ - { - "type": "image_url", - "image_url": { - "url": f"data:image/png;base64,{png_b64}", - }, - }, - ], - } - - result = convert_to_anthropic_tool_result(message) - - content = result["content"] - assert isinstance(content, list) and len(content) == 1 - block = content[0] - assert block["type"] == "image" - assert block["source"]["media_type"] == "image/png" + with pytest.raises(ValueError, match="only supports base64-encoded"): + _bedrock_converse_messages_pt( + messages, "anthropic.claude-sonnet-4-6", "bedrock" + ) diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py index d6281703a0a..49d3c51e340 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py @@ -878,6 +878,39 @@ def test_sync_streaming_bad_request_not_midstream(logging_obj: Logging): assert "invalid maxOutputTokens" in str(excinfo.value) +@pytest.mark.asyncio +async def test_async_streaming_read_timeout_triggers_midstream_fallback( + logging_obj: Logging, +): + """A mid-stream httpx.ReadTimeout must wrap into MidStreamFallbackError so + the Router's FallbackStreamWrapper can switch to a fallback model. + + Previously __anext__ caught httpx.TimeoutException and re-raised it raw, + which bypassed _handle_stream_fallback_error and prevented stream_timeout + from triggering fallbacks the way connection-phase timeout does. + """ + import httpx + + from litellm.exceptions import MidStreamFallbackError + + async def _raise_read_timeout(**kwargs): + raise httpx.ReadTimeout("Timeout on reading data from socket") + + response = CustomStreamWrapper( + completion_stream=None, + model="gpt-4", + logging_obj=logging_obj, + custom_llm_provider="openai", + make_call=_raise_read_timeout, + ) + + with pytest.raises(MidStreamFallbackError) as excinfo: + await response.__anext__() + + assert excinfo.value.is_pre_first_chunk is True + assert isinstance(excinfo.value.original_exception, Exception) + + def test_streaming_handler_with_created_time_propagation( initialized_custom_stream_wrapper: CustomStreamWrapper, logging_obj: Logging ): diff --git a/tests/test_litellm/litellm_core_utils/test_url_utils.py b/tests/test_litellm/litellm_core_utils/test_url_utils.py index 4579c203218..e363418e774 100644 --- a/tests/test_litellm/litellm_core_utils/test_url_utils.py +++ b/tests/test_litellm/litellm_core_utils/test_url_utils.py @@ -4,7 +4,13 @@ import pytest import litellm from litellm.litellm_core_utils import url_utils -from litellm.litellm_core_utils.url_utils import SSRFError, _is_blocked_ip, validate_url +from litellm.litellm_core_utils.url_utils import ( + SSRFError, + _is_blocked_ip, + encode_url_path_segment, + encode_url_path_segments, + validate_url, +) @pytest.fixture @@ -80,6 +86,28 @@ class TestIsBlockedIp: assert _is_blocked_ip("::ffff:168.63.129.16") is True +class TestEncodeUrlPathSegment: + def test_encodes_path_delimiters_and_query_markers(self): + encoded = encode_url_path_segment("../../v1/files?limit=1#frag") + + assert encoded == "..%2F..%2Fv1%2Ffiles%3Flimit%3D1%23frag" + + def test_encodes_path_segments_without_collapsing_valid_model_paths(self): + encoded = encode_url_path_segments("@cf/meta/model?debug=1") + + assert encoded == "%40cf/meta/model%3Fdebug%3D1" + + @pytest.mark.parametrize("value", ["", ".", "..", None]) + def test_rejects_empty_and_dot_segments(self, value): + with pytest.raises(ValueError): + encode_url_path_segment(value, field_name="resource_id") + + @pytest.mark.parametrize("value", ["../model", "model/../other", "/model"]) + def test_rejects_dot_segments_in_multi_segment_paths(self, value): + with pytest.raises(ValueError): + encode_url_path_segments(value, field_name="model") + + class TestValidateUrl: def test_blocks_loopback(self): with pytest.raises(SSRFError): @@ -394,3 +422,76 @@ class TestHostAllowlist: monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake) validate_url("http://internal.corp/") + + +# ── assert_same_origin ──────────────────────────────────────────────────────── + + +from litellm.litellm_core_utils.url_utils import assert_same_origin + + +def test_assert_same_origin_matches_scheme_host_port(): + """A polling URL on the same scheme + host + port as the api_base + passes — the upstream is trusted; the URL it returned points back at + the same upstream.""" + assert_same_origin( + "https://api.example.com/v1/operations/abc", + "https://api.example.com/v1/generate", + ) + + +def test_assert_same_origin_treats_default_ports_as_explicit(): + """``https://x/`` and ``https://x:443/`` are the same origin.""" + assert_same_origin("https://api.example.com/poll", "https://api.example.com:443/") + assert_same_origin("https://api.example.com:443/poll", "https://api.example.com/") + assert_same_origin("http://api.example.com/poll", "http://api.example.com:80/") + + +def test_assert_same_origin_rejects_different_host(): + with pytest.raises(SSRFError, match="host"): + assert_same_origin( + "https://attacker.example.com/poll", + "https://api.example.com/generate", + ) + + +def test_assert_same_origin_rejects_different_scheme(): + with pytest.raises(SSRFError, match="scheme"): + assert_same_origin( + "http://api.example.com/poll", "https://api.example.com/generate" + ) + + +def test_assert_same_origin_rejects_different_port(): + with pytest.raises(SSRFError, match="port"): + assert_same_origin( + "https://api.example.com:8443/poll", "https://api.example.com/generate" + ) + + +def test_assert_same_origin_rejects_non_http_scheme(): + """``file://`` polling URLs are rejected outright — the upstream + should never return a non-HTTP scheme.""" + with pytest.raises(SSRFError, match="scheme"): + assert_same_origin("file:///etc/passwd", "https://api.example.com/") + + +def test_assert_same_origin_case_insensitive_host(): + assert_same_origin( + "https://API.example.com/poll", "https://api.example.com/generate" + ) + + +def test_assert_same_origin_error_message_does_not_leak_hostnames(): + """Greptile P2: in the SSRF threat model the caller is the attacker. + The error message must not echo the operator's expected host or the + attacker-supplied candidate host back to the caller — only identify + *which* component mismatched.""" + with pytest.raises(SSRFError) as exc: + assert_same_origin( + "https://attacker.example.com:1234/poll", + "https://api.internal-corp.example/generate", + ) + detail = str(exc.value) + assert "attacker.example.com" not in detail + assert "api.internal-corp.example" not in detail diff --git a/tests/test_litellm/llms/anthropic/files/test_anthropic_files_transformation.py b/tests/test_litellm/llms/anthropic/files/test_anthropic_files_transformation.py index e9509be9e11..9fc4981510f 100644 --- a/tests/test_litellm/llms/anthropic/files/test_anthropic_files_transformation.py +++ b/tests/test_litellm/llms/anthropic/files/test_anthropic_files_transformation.py @@ -180,6 +180,19 @@ class TestAnthropicFilesConfig: assert url == "https://custom.api.com/v1/files/file-abc123" assert params == {} + def test_transform_retrieve_file_request_encodes_path_traversal(self): + url, params = self.config.transform_retrieve_file_request( + file_id="../../v1/messages/batches?limit=1#frag", + optional_params={}, + litellm_params={}, + ) + + assert ( + url + == f"{ANTHROPIC_FILES_API_BASE}/v1/files/..%2F..%2Fv1%2Fmessages%2Fbatches%3Flimit%3D1%23frag" + ) + assert params == {} + def test_transform_retrieve_file_response(self): mock_response = Mock(spec=httpx.Response) mock_response.json.return_value = { @@ -296,6 +309,14 @@ class TestAnthropicFilesConfig: assert url == f"{ANTHROPIC_FILES_API_BASE}/v1/files/file-abc123/content" assert params == {} + def test_transform_file_content_request_rejects_dot_segment(self): + with pytest.raises(ValueError, match="file_id cannot be a dot path segment"): + self.config.transform_file_content_request( + file_content_request={"file_id": ".."}, + optional_params={}, + litellm_params={}, + ) + def test_transform_file_content_response(self): mock_response = Mock(spec=httpx.Response) result = self.config.transform_file_content_response( diff --git a/tests/test_litellm/llms/azure/response/test_azure_transformation.py b/tests/test_litellm/llms/azure/response/test_azure_transformation.py index 9519b7c8a5b..a4bd14d69ff 100644 --- a/tests/test_litellm/llms/azure/response/test_azure_transformation.py +++ b/tests/test_litellm/llms/azure/response/test_azure_transformation.py @@ -96,6 +96,28 @@ def test_get_complete_url(): assert result == expected +@pytest.mark.serial +def test_response_id_path_requests_encode_response_id(): + config = AzureOpenAIResponsesAPIConfig() + api_base = ( + "https://litellm8397336933.openai.azure.com/openai/responses" + "?api-version=2024-05-01-preview" + ) + + url, params = config.transform_cancel_response_api_request( + response_id="../../responses/other?x=1#frag", + api_base=api_base, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert ( + url + == "https://litellm8397336933.openai.azure.com/openai/responses/..%2F..%2Fresponses%2Fother%3Fx%3D1%23frag/cancel?api-version=2024-05-01-preview" + ) + assert params == {} + + @pytest.mark.serial def test_azure_o_series_responses_api_supported_params(): """Test that Azure OpenAI O-series responses API excludes temperature from supported parameters.""" diff --git a/tests/test_litellm/llms/azure_ai/test_azure_ai_agents_handler.py b/tests/test_litellm/llms/azure_ai/test_azure_ai_agents_handler.py new file mode 100644 index 00000000000..f65573b7ae1 --- /dev/null +++ b/tests/test_litellm/llms/azure_ai/test_azure_ai_agents_handler.py @@ -0,0 +1,57 @@ +import pytest + +from litellm.llms.azure_ai.agents.handler import AzureAIAgentsHandler + + +def test_should_encode_thread_id_in_azure_ai_agent_urls(): + handler = AzureAIAgentsHandler() + + assert ( + handler._build_messages_url( + "https://example.services.ai.azure.com/api/projects/proj", + "../../threads/other?x=1#frag", + "2024-05-01-preview", + ) + == "https://example.services.ai.azure.com/api/projects/proj/threads/..%2F..%2Fthreads%2Fother%3Fx%3D1%23frag/messages?api-version=2024-05-01-preview" + ) + assert ( + handler._build_runs_url( + "https://example.services.ai.azure.com/api/projects/proj", + "thread/abc", + "2024-05-01-preview", + ) + == "https://example.services.ai.azure.com/api/projects/proj/threads/thread%2Fabc/runs?api-version=2024-05-01-preview" + ) + + +def test_should_encode_thread_and_run_ids_in_azure_ai_agent_status_url(): + handler = AzureAIAgentsHandler() + + assert ( + handler._build_run_status_url( + "https://example.services.ai.azure.com/api/projects/proj", + "thread/abc", + "../runs/other#frag", + "2024-05-01-preview", + ) + == "https://example.services.ai.azure.com/api/projects/proj/threads/thread%2Fabc/runs/..%2Fruns%2Fother%23frag?api-version=2024-05-01-preview" + ) + + +def test_should_reject_dot_segments_in_azure_ai_agent_urls(): + handler = AzureAIAgentsHandler() + + with pytest.raises(ValueError, match="thread_id cannot be a dot path segment"): + handler._build_messages_url( + "https://example.services.ai.azure.com/api/projects/proj", + "..", + "2024-05-01-preview", + ) + + with pytest.raises(ValueError, match="run_id cannot be a dot path segment"): + handler._build_run_status_url( + "https://example.services.ai.azure.com/api/projects/proj", + "thread_123", + "..", + "2024-05-01-preview", + ) diff --git a/tests/test_litellm/llms/azure_ai/test_azure_document_intelligence_ocr_transformation.py b/tests/test_litellm/llms/azure_ai/test_azure_document_intelligence_ocr_transformation.py new file mode 100644 index 00000000000..e638be68ec0 --- /dev/null +++ b/tests/test_litellm/llms/azure_ai/test_azure_document_intelligence_ocr_transformation.py @@ -0,0 +1,33 @@ +import pytest + +from litellm.llms.azure_ai.ocr.document_intelligence.transformation import ( + AzureDocumentIntelligenceOCRConfig, +) + + +def test_should_encode_azure_document_intelligence_model_id(): + config = AzureDocumentIntelligenceOCRConfig() + + url = config.get_complete_url( + api_base="https://example.cognitiveservices.azure.com", + model="prebuilt-layout?x=1#frag", + optional_params={}, + litellm_params={}, + ) + + assert ( + url + == "https://example.cognitiveservices.azure.com/documentintelligence/documentModels/prebuilt-layout%3Fx%3D1%23frag:analyze?api-version=2024-11-30" + ) + + +def test_should_reject_dot_segment_azure_document_intelligence_model_id(): + config = AzureDocumentIntelligenceOCRConfig() + + with pytest.raises(ValueError, match="model_id cannot be a dot path segment"): + config.get_complete_url( + api_base="https://example.cognitiveservices.azure.com", + model="azure_ai/doc-intelligence/..", + optional_params={}, + litellm_params={}, + ) diff --git a/tests/test_litellm/llms/bedrock/count_tokens/test_bedrock_count_tokens_transformation.py b/tests/test_litellm/llms/bedrock/count_tokens/test_bedrock_count_tokens_transformation.py index eac022ec237..0de3f833a37 100644 --- a/tests/test_litellm/llms/bedrock/count_tokens/test_bedrock_count_tokens_transformation.py +++ b/tests/test_litellm/llms/bedrock/count_tokens/test_bedrock_count_tokens_transformation.py @@ -167,3 +167,24 @@ def test_tool_name_sanitization(): ] # Should be sanitized: only [a-zA-Z0-9_] assert tool_name == "my_tool_" + + +def test_count_tokens_endpoint_encodes_model_id(monkeypatch): + """Test model IDs are treated as a single Bedrock path segment.""" + config = BedrockCountTokensConfig() + + monkeypatch.setattr( + config, + "get_runtime_endpoint", + lambda **kwargs: ("https://bedrock-runtime.us-east-1.amazonaws.com", None), + ) + + endpoint = config.get_bedrock_count_tokens_endpoint( + model="bedrock/../../model/other?x=1#frag", + aws_region_name="us-east-1", + ) + + assert ( + endpoint + == "https://bedrock-runtime.us-east-1.amazonaws.com/model/..%2F..%2Fmodel%2Fother%3Fx%3D1%23frag/count-tokens" + ) diff --git a/tests/test_litellm/llms/bedrock/invoke_agent/test_bedrock_agent_transformation.py b/tests/test_litellm/llms/bedrock/invoke_agent/test_bedrock_agent_transformation.py index ffd3526698a..9e526e47784 100644 --- a/tests/test_litellm/llms/bedrock/invoke_agent/test_bedrock_agent_transformation.py +++ b/tests/test_litellm/llms/bedrock/invoke_agent/test_bedrock_agent_transformation.py @@ -1,12 +1,9 @@ import base64 -import json import os import sys -from litellm._uuid import uuid -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import patch import pytest -from fastapi.testclient import TestClient sys.path.insert( 0, os.path.abspath("../../../..") @@ -15,12 +12,7 @@ sys.path.insert( from litellm.llms.bedrock.chat.invoke_agent.transformation import ( AmazonInvokeAgentConfig, ) -from litellm.types.llms.bedrock_invoke_agents import ( - InvokeAgentEvent, - InvokeAgentEventHeaders, - InvokeAgentUsage, -) -from litellm.types.utils import Message, ModelResponse, Usage +from litellm.types.utils import ModelResponse class TestAmazonInvokeAgentConfig: @@ -270,3 +262,28 @@ class TestAmazonInvokeAgentConfig: "https://bedrock-runtime.us-east-1.amazonaws.com/agents/L1RT58GYRW/agentAliases/MFPSBCXYTW/sessions" in result ) + + @patch( + "litellm.llms.bedrock.chat.invoke_agent.transformation.convert_content_list_to_str" + ) + @patch.object(AmazonInvokeAgentConfig, "get_runtime_endpoint") + @patch.object(AmazonInvokeAgentConfig, "_get_aws_region_name") + def test_get_complete_url_encodes_session_id( + self, mock_region, mock_endpoint, mock_convert, config + ): + """Test get_complete_url encodes session ID path segment.""" + mock_endpoint.return_value = ( + "https://bedrock-runtime.us-east-1.amazonaws.com", + None, + ) + mock_region.return_value = "us-east-1" + + result = config.get_complete_url( + api_base=None, + api_key=None, + model="agent/L1RT58GYRW/MFPSBCXYTW", + optional_params={"sessionID": "../../sessions/other?x=1#frag"}, + litellm_params={}, + ) + + assert "sessions/..%2F..%2Fsessions%2Fother%3Fx%3D1%23frag/text" in result diff --git a/tests/test_litellm/llms/bedrock/vector_stores/test_bedrock_vector_store_transformation.py b/tests/test_litellm/llms/bedrock/vector_stores/test_bedrock_vector_store_transformation.py index d60d0487d06..7b04efa17dc 100644 --- a/tests/test_litellm/llms/bedrock/vector_stores/test_bedrock_vector_store_transformation.py +++ b/tests/test_litellm/llms/bedrock/vector_stores/test_bedrock_vector_store_transformation.py @@ -28,6 +28,28 @@ def test_transform_search_request(): assert body["retrievalQuery"].get("text") == "hello" +def test_transform_search_request_encodes_vector_store_id(): + config = BedrockVectorStoreConfig() + mock_log = MagicMock() + mock_log.model_call_details = {} + + url, body = config.transform_search_vector_store_request( + vector_store_id="../../knowledgebases/other?x=1#frag", + query="hello", + vector_store_search_optional_params={}, + api_base="https://bedrock-agent-runtime.us-west-2.amazonaws.com/knowledgebases", + litellm_logging_obj=mock_log, + litellm_params={}, + extra_body=None, + ) + + assert ( + url + == "https://bedrock-agent-runtime.us-west-2.amazonaws.com/knowledgebases/..%2F..%2Fknowledgebases%2Fother%3Fx%3D1%23frag/retrieve" + ) + assert body["retrievalQuery"].get("text") == "hello" + + def test_transform_search_request_uses_only_retrieval_config_from_extra_body(): config = BedrockVectorStoreConfig() mock_log = MagicMock() diff --git a/tests/test_litellm/llms/bytez/chat/test_bytez_chat_transformation.py b/tests/test_litellm/llms/bytez/chat/test_bytez_chat_transformation.py index dd388fedc61..2f8cc5484ba 100644 --- a/tests/test_litellm/llms/bytez/chat/test_bytez_chat_transformation.py +++ b/tests/test_litellm/llms/bytez/chat/test_bytez_chat_transformation.py @@ -87,6 +87,29 @@ class TestBytezChatConfig: assert response.choices[0].message.content == output_content # type: ignore + def test_get_complete_url_encodes_model_path_segment(self): + config = BytezChatConfig() + + assert ( + config.get_complete_url( + api_base=API_BASE, + api_key=TEST_API_KEY, + model="google/gemma?x=1#frag", + optional_params={}, + litellm_params={}, + ) + == f"{API_BASE}/google/gemma%3Fx%3D1%23frag" + ) + + with pytest.raises(ValueError, match="dot path segment"): + config.get_complete_url( + api_base=API_BASE, + api_key=TEST_API_KEY, + model="../../models/other", + optional_params={}, + litellm_params={}, + ) + def test_bytez_messages_adaptation(self): cases = [ dict( diff --git a/tests/test_litellm/llms/cloudflare/test_cloudflare_transformation.py b/tests/test_litellm/llms/cloudflare/test_cloudflare_transformation.py new file mode 100644 index 00000000000..cecb6024de1 --- /dev/null +++ b/tests/test_litellm/llms/cloudflare/test_cloudflare_transformation.py @@ -0,0 +1,27 @@ +import pytest + +from litellm.llms.cloudflare.chat.transformation import CloudflareChatConfig + + +def test_get_complete_url_encodes_model_path_segment(): + config = CloudflareChatConfig() + + assert ( + config.get_complete_url( + api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/run/", + api_key="cf-key", + model="@cf/meta/llama?x=1#frag", + optional_params={}, + litellm_params={}, + ) + == "https://api.cloudflare.com/client/v4/accounts/acct/ai/run/%40cf/meta/llama%3Fx%3D1%23frag" + ) + + with pytest.raises(ValueError, match="dot path segment"): + config.get_complete_url( + api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/run/", + api_key="cf-key", + model="../../accounts/other", + optional_params={}, + litellm_params={}, + ) diff --git a/tests/test_litellm/llms/elevenlabs/test_elevenlabs_text_to_speech_transformation.py b/tests/test_litellm/llms/elevenlabs/test_elevenlabs_text_to_speech_transformation.py new file mode 100644 index 00000000000..54e689dea6b --- /dev/null +++ b/tests/test_litellm/llms/elevenlabs/test_elevenlabs_text_to_speech_transformation.py @@ -0,0 +1,33 @@ +import pytest + +from litellm.llms.elevenlabs.text_to_speech.transformation import ( + ElevenLabsTextToSpeechConfig, +) + + +def test_should_encode_elevenlabs_voice_id_path_segment(): + config = ElevenLabsTextToSpeechConfig() + + url = config.get_complete_url( + model="elevenlabs/tts", + api_base="https://api.elevenlabs.io", + litellm_params={ + config.ELEVENLABS_VOICE_ID_KEY: "voice/../../models?x=1#frag", + }, + ) + + assert ( + url + == "https://api.elevenlabs.io/v1/text-to-speech/voice%2F..%2F..%2Fmodels%3Fx%3D1%23frag" + ) + + +def test_should_reject_dot_segment_elevenlabs_voice_id(): + config = ElevenLabsTextToSpeechConfig() + + with pytest.raises(ValueError, match="voice_id cannot be a dot path segment"): + config.get_complete_url( + model="elevenlabs/tts", + api_base="https://api.elevenlabs.io", + litellm_params={config.ELEVENLABS_VOICE_ID_KEY: ".."}, + ) diff --git a/tests/test_litellm/llms/gemini/files/test_gemini_files_transformation.py b/tests/test_litellm/llms/gemini/files/test_gemini_files_transformation.py index 2431c9a9c4f..a2f95724688 100644 --- a/tests/test_litellm/llms/gemini/files/test_gemini_files_transformation.py +++ b/tests/test_litellm/llms/gemini/files/test_gemini_files_transformation.py @@ -2,7 +2,6 @@ Test Google AI Studio (Gemini) files transformation functionality """ -import os from unittest.mock import Mock, patch import httpx @@ -93,6 +92,30 @@ class TestGoogleAIStudioFilesTransformation: assert "key=" not in url assert params == {} + def test_transform_retrieve_file_request_encodes_file_id_path_segment(self): + file_id = "files/../../models/gemini-pro?x=1#frag" + litellm_params = {"api_key": "test-api-key"} + + url, params = self.handler.transform_retrieve_file_request( + file_id=file_id, + optional_params={}, + litellm_params=litellm_params, + ) + + assert ( + url + == "https://generativelanguage.googleapis.com/v1beta/files/..%2F..%2Fmodels%2Fgemini-pro%3Fx%3D1%23frag" + ) + assert params == {} + + def test_transform_retrieve_file_request_rejects_dot_path_segment(self): + with pytest.raises(ValueError, match="file_id cannot be a dot path segment"): + self.handler.transform_retrieve_file_request( + file_id="files/..", + optional_params={}, + litellm_params={"api_key": "test-api-key"}, + ) + @patch.dict("os.environ", {}, clear=True) @patch("litellm.llms.gemini.common_utils.get_secret_str", return_value=None) def test_transform_retrieve_file_request_missing_api_key(self, mock_get_secret): @@ -297,9 +320,7 @@ class TestGoogleAIStudioFilesTransformation: litellm_params=litellm_params, ) - # Verify URL extraction - assert "files/test123" in url - assert "generativelanguage.googleapis.com" in url + assert url == "https://generativelanguage.googleapis.com/v1beta/files/test123" # Params should be empty (API key goes in header via validate_environment) assert params == {} @@ -322,3 +343,22 @@ class TestGoogleAIStudioFilesTransformation: assert file_id in url assert "generativelanguage.googleapis.com" in url assert params == {} + + def test_transform_delete_file_request_encodes_file_id_path_segment(self): + file_id = "files/../../models/gemini-pro?x=1#frag" + litellm_params = { + "api_key": "test-api-key", + "api_base": "https://generativelanguage.googleapis.com", + } + + url, params = self.handler.transform_delete_file_request( + file_id=file_id, + optional_params={}, + litellm_params=litellm_params, + ) + + assert ( + url + == "https://generativelanguage.googleapis.com/v1beta/files/..%2F..%2Fmodels%2Fgemini-pro%3Fx%3D1%23frag" + ) + assert params == {} diff --git a/tests/test_litellm/llms/manus/responses/test_manus_responses_transformation.py b/tests/test_litellm/llms/manus/responses/test_manus_responses_transformation.py index 10d66174c59..43ce030323b 100644 --- a/tests/test_litellm/llms/manus/responses/test_manus_responses_transformation.py +++ b/tests/test_litellm/llms/manus/responses/test_manus_responses_transformation.py @@ -58,3 +58,18 @@ def test_transform_responses_api_request_adds_manus_params(): assert result["agent_profile"] == "manus-1.6" assert "input" in result assert "model" in result + + +def test_get_response_request_encodes_response_id(): + """Test response IDs are encoded before being appended to Manus URLs.""" + config = ManusResponsesAPIConfig() + + url, params = config.transform_get_response_api_request( + response_id="../../files?x=1#frag", + api_base="https://api.manus.im/v1/responses", + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert url == "https://api.manus.im/v1/responses/..%2F..%2Ffiles%3Fx%3D1%23frag" + assert params == {} diff --git a/tests/test_litellm/llms/openai/evals/test_openai_evals_transformation.py b/tests/test_litellm/llms/openai/evals/test_openai_evals_transformation.py index 751e48ff7a5..f39be511b97 100644 --- a/tests/test_litellm/llms/openai/evals/test_openai_evals_transformation.py +++ b/tests/test_litellm/llms/openai/evals/test_openai_evals_transformation.py @@ -49,6 +49,17 @@ def test_get_complete_url_with_eval_id(config: OpenAIEvalsConfig): assert url == "https://api.openai.com/v1/evals/eval_123" +def test_get_complete_url_encodes_eval_id(config: OpenAIEvalsConfig): + """Test eval_id is treated as a single path segment.""" + url = config.get_complete_url( + api_base="https://api.openai.com", + endpoint="evals", + eval_id="../../files?x=1#frag", + ) + + assert url == "https://api.openai.com/v1/evals/..%2F..%2Ffiles%3Fx%3D1%23frag" + + def test_get_complete_url_without_eval_id(config: OpenAIEvalsConfig): """Test URL construction without eval_id""" url = config.get_complete_url( @@ -253,3 +264,20 @@ def test_transform_cancel_eval_response(config: OpenAIEvalsConfig): assert result.id == "eval_123" assert result.object == "eval" + + +def test_transform_run_requests_encode_eval_and_run_ids(config: OpenAIEvalsConfig): + """Test run path IDs are treated as single path segments.""" + url, _, request_body = config.transform_cancel_run_request( + eval_id="../../evals?x=1#frag", + run_id="../runs#other", + api_base="https://api.openai.com", + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert ( + url + == "https://api.openai.com/v1/evals/..%2F..%2Fevals%3Fx%3D1%23frag/runs/..%2Fruns%23other/cancel" + ) + assert request_body == {} diff --git a/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py b/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py index dae87842832..acb9fa9b64c 100644 --- a/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py +++ b/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py @@ -265,6 +265,24 @@ class TestOpenAIResponsesAPIConfig: assert result == "https://custom-openai.example.com/v1/responses" + def test_response_id_path_requests_encode_response_id(self): + """Test response_id is treated as one upstream URL path segment.""" + api_base = "https://custom-openai.example.com/v1/responses" + response_id = "../../files?x=1#frag" + + url, data = self.config.transform_list_input_items_request( + response_id=response_id, + api_base=api_base, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert ( + url + == "https://custom-openai.example.com/v1/responses/..%2F..%2Ffiles%3Fx%3D1%23frag/input_items" + ) + assert data["limit"] == 20 + def test_get_event_model_class_generic_event(self): """Test that get_event_model_class returns the correct event model class""" from litellm.types.llms.openai import GenericEvent @@ -547,7 +565,12 @@ class TestOpenAIResponsesAPIConfig: """Base helper strips ``namespace`` from custom_tool_call for every provider path.""" inp = [ {"type": "function_call", "call_id": "a", "name": "f", "namespace": "keep"}, - {"type": "custom_tool_call", "call_id": "b", "name": "c", "namespace": "drop"}, + { + "type": "custom_tool_call", + "call_id": "b", + "name": "c", + "namespace": "drop", + }, ] out = BaseResponsesAPIConfig.strip_custom_tool_call_namespace_from_responses_input( inp diff --git a/tests/test_litellm/llms/openai/vector_store_files/test_openai_vector_store_files_transformation.py b/tests/test_litellm/llms/openai/vector_store_files/test_openai_vector_store_files_transformation.py index a8e07cd3645..aebd93d062e 100644 --- a/tests/test_litellm/llms/openai/vector_store_files/test_openai_vector_store_files_transformation.py +++ b/tests/test_litellm/llms/openai/vector_store_files/test_openai_vector_store_files_transformation.py @@ -32,6 +32,21 @@ def test_get_complete_url(config: OpenAIVectorStoreFilesConfig): assert url == "https://api.example.com/v1/vector_stores/vs_123/files" +def test_get_complete_url_encodes_vector_store_id( + config: OpenAIVectorStoreFilesConfig, +): + url = config.get_complete_url( + api_base="https://api.example.com/v1", + vector_store_id="../vs_123?x=1#frag", + litellm_params={}, + ) + + assert ( + url + == "https://api.example.com/v1/vector_stores/..%2Fvs_123%3Fx%3D1%23frag/files" + ) + + def test_transform_create_request(config: OpenAIVectorStoreFilesConfig): api_base = "https://api.example.com/v1/vector_stores/vs_123/files" url, payload = config.transform_create_vector_store_file_request( @@ -60,6 +75,22 @@ def test_transform_list_request(config: OpenAIVectorStoreFilesConfig): assert params == {"limit": 2, "order": "asc"} +def test_transform_file_request_encodes_file_id(config: OpenAIVectorStoreFilesConfig): + api_base = "https://api.example.com/v1/vector_stores/vs_123/files" + + url, params = config.transform_retrieve_vector_store_file_content_request( + vector_store_id="vs_123", + file_id="../../files?x=1#frag", + api_base=api_base, + ) + + assert ( + url + == "https://api.example.com/v1/vector_stores/vs_123/files/..%2F..%2Ffiles%3Fx%3D1%23frag/content" + ) + assert params == {} + + def test_transform_create_response(config: OpenAIVectorStoreFilesConfig): response = httpx.Response( status_code=200, diff --git a/tests/test_litellm/llms/openai/vector_stores/test_openai_vector_stores_transformation.py b/tests/test_litellm/llms/openai/vector_stores/test_openai_vector_stores_transformation.py index 053b107afbd..e7b1aab45b4 100644 --- a/tests/test_litellm/llms/openai/vector_stores/test_openai_vector_stores_transformation.py +++ b/tests/test_litellm/llms/openai/vector_stores/test_openai_vector_stores_transformation.py @@ -64,3 +64,21 @@ class TestOpenAIVectorStoreAPIConfig: for i in range(16): assert f"key_{i}" in request_body["metadata"] assert request_body["metadata"][f"key_{i}"] == f"value_{i}" + + def test_transform_search_vector_store_request_encodes_vector_store_id(self): + config = OpenAIVectorStoreConfig() + + url, request_body = config.transform_search_vector_store_request( + vector_store_id="../../files?x=1#frag", + query="hello", + vector_store_search_optional_params={}, + api_base="https://api.openai.com/v1/vector_stores", + litellm_logging_obj=None, # type: ignore[arg-type] + litellm_params={}, + ) + + assert ( + url + == "https://api.openai.com/v1/vector_stores/..%2F..%2Ffiles%3Fx%3D1%23frag/search" + ) + assert request_body["query"] == "hello" diff --git a/tests/test_litellm/llms/openai/videos/test_openai_video_transformation.py b/tests/test_litellm/llms/openai/videos/test_openai_video_transformation.py new file mode 100644 index 00000000000..c15554a46a5 --- /dev/null +++ b/tests/test_litellm/llms/openai/videos/test_openai_video_transformation.py @@ -0,0 +1,70 @@ +from litellm.llms.openai.videos.transformation import OpenAIVideoConfig +from litellm.types.router import GenericLiteLLMParams +from litellm.types.videos.utils import encode_character_id_with_provider + + +def test_video_content_request_encodes_video_id_path_segment(): + config = OpenAIVideoConfig() + + url, params = config.transform_video_content_request( + video_id="../../responses?x=1#frag", + api_base="https://api.openai.com/v1/videos", + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert ( + url + == "https://api.openai.com/v1/videos/..%2F..%2Fresponses%3Fx%3D1%23frag/content" + ) + assert params == {} + + +def test_video_content_request_encodes_variant_query_param(): + """``variant`` is user-controlled and was previously interpolated raw + into the query string. A value like ``thumbnail&extra=1`` would + inject additional query parameters into the upstream request.""" + config = OpenAIVideoConfig() + + url, _ = config.transform_video_content_request( + video_id="vid_123", + api_base="https://api.openai.com/v1/videos", + litellm_params=GenericLiteLLMParams(), + headers={}, + variant="thumbnail&extra=1", + ) + + # ``&`` and ``=`` must be percent-encoded so they cannot terminate + # the ``variant`` value or open a new query parameter. + assert "?variant=thumbnail%26extra%3D1" in url + # Sanity: the legitimate "thumbnail" value still round-trips cleanly. + url2, _ = config.transform_video_content_request( + video_id="vid_123", + api_base="https://api.openai.com/v1/videos", + litellm_params=GenericLiteLLMParams(), + headers={}, + variant="thumbnail", + ) + assert url2.endswith("?variant=thumbnail") + + +def test_wrapped_character_id_is_decoded_then_encoded_as_path_segment(): + config = OpenAIVideoConfig() + character_id = encode_character_id_with_provider( + "../../characters?x=1#frag", + provider="openai", + model_id="sora", + ) + + url, params = config.transform_video_get_character_request( + character_id=character_id, + api_base="https://api.openai.com/v1/videos", + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert ( + url + == "https://api.openai.com/v1/videos/characters/..%2F..%2Fcharacters%3Fx%3D1%23frag" + ) + assert params == {} diff --git a/tests/test_litellm/llms/pg_vector/vector_stores/test_pg_vector_transformation.py b/tests/test_litellm/llms/pg_vector/vector_stores/test_pg_vector_transformation.py index e3343c3037a..56953a574d6 100644 --- a/tests/test_litellm/llms/pg_vector/vector_stores/test_pg_vector_transformation.py +++ b/tests/test_litellm/llms/pg_vector/vector_stores/test_pg_vector_transformation.py @@ -141,6 +141,24 @@ class TestPGVectorStoreConfig: assert headers["Authorization"] == "Bearer test_key" assert url == "https://example.com/v1/vector_stores" + def test_search_request_encodes_vector_store_id(self): + config = PGVectorStoreConfig() + + url, request_body = config.transform_search_vector_store_request( + vector_store_id="../../files?x=1#frag", + query="hello", + vector_store_search_optional_params={}, + api_base="https://example.com/v1/vector_stores", + litellm_logging_obj=Mock(), + litellm_params={}, + ) + + assert ( + url + == "https://example.com/v1/vector_stores/..%2F..%2Ffiles%3Fx%3D1%23frag/search" + ) + assert request_body["query"] == "hello" + def test_environment_variable_support(self): """ Test that environment variables are supported for configuration. diff --git a/tests/test_litellm/llms/ragflow/chat/test_ragflow_chat_transformation.py b/tests/test_litellm/llms/ragflow/chat/test_ragflow_chat_transformation.py index ae43eac7ffc..baf2ab33910 100644 --- a/tests/test_litellm/llms/ragflow/chat/test_ragflow_chat_transformation.py +++ b/tests/test_litellm/llms/ragflow/chat/test_ragflow_chat_transformation.py @@ -117,6 +117,24 @@ class TestRAGFlowChatTransformation: == "http://localhost:9380/api/v1/agents_openai/my-agent-id/chat/completions" ) + def test_get_complete_url_encodes_entity_id(self): + """Test RAGFlow chat IDs are encoded as one upstream path segment.""" + config = RAGFlowConfig() + + url = config.get_complete_url( + api_base="http://localhost:9380", + api_key=None, + model="ragflow/chat/..%2F..%2Fagents_openai%2Fother/gpt-4o-mini", + optional_params={}, + litellm_params={}, + stream=False, + ) + + assert ( + url + == "http://localhost:9380/api/v1/chats_openai/..%252F..%252Fagents_openai%252Fother/chat/completions" + ) + def test_get_complete_url_strips_v1(self): """Test URL construction when api_base ends with /v1.""" config = RAGFlowConfig() diff --git a/tests/test_litellm/llms/runwayml/videos/test_runway_video_transformation.py b/tests/test_litellm/llms/runwayml/videos/test_runway_video_transformation.py index 755716c9da5..24879ce83f9 100644 --- a/tests/test_litellm/llms/runwayml/videos/test_runway_video_transformation.py +++ b/tests/test_litellm/llms/runwayml/videos/test_runway_video_transformation.py @@ -134,6 +134,21 @@ class TestRunwayMLVideoTransformation: with pytest.raises(ValueError, match="still processing"): self.config._extract_video_url_from_response(processing_response) + def test_transform_video_status_encodes_video_id_path_segment(self): + """Test task IDs are encoded before being appended to Runway URLs.""" + url, params = self.config.transform_video_status_retrieve_request( + video_id="../../tasks/other?x=1#frag", + api_base="https://api.dev.runwayml.com/v1", + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert ( + url + == "https://api.dev.runwayml.com/v1/tasks/..%2F..%2Ftasks%2Fother%3Fx%3D1%23frag" + ) + assert params == {} + def test_full_video_workflow(self): """Test complete video generation workflow from creation to status check.""" config = RunwayMLVideoConfig() diff --git a/tests/test_litellm/llms/test_polling_url_origin_match.py b/tests/test_litellm/llms/test_polling_url_origin_match.py new file mode 100644 index 00000000000..f1f910bc731 --- /dev/null +++ b/tests/test_litellm/llms/test_polling_url_origin_match.py @@ -0,0 +1,177 @@ +""" +VERIA-51: polling URLs returned by upstream APIs (Azure DALL-E, +Azure Document Intelligence, Black Forest Labs) used to be followed +without origin validation. The handlers attached the operator's API +key to the polling request, so an attacker who could influence the +upstream response (or a compromised upstream) could redirect the proxy +to send credentials anywhere. + +These tests assert each handler now rejects polling URLs that don't +share an origin with the original request URL. +""" + +from unittest.mock import MagicMock, patch + +import httpx +import pytest + + +# Azure DALL-E sync + async paths route through ``assert_same_origin`` +# the same way as the cases below. The helper itself is unit-tested in +# ``tests/test_litellm/litellm_core_utils/test_url_utils.py``; the +# tests here exercise the wiring at sites with simpler signatures. + + +# ── Azure Document Intelligence polling ─────────────────────────────────────── + + +def test_azure_di_sync_rejects_cross_origin_polling(): + from litellm.llms.azure_ai.ocr.document_intelligence.transformation import ( + AzureDocumentIntelligenceOCRConfig, + ) + + config = AzureDocumentIntelligenceOCRConfig() + + raw_response = MagicMock() + raw_response.status_code = 202 + raw_response.headers = { + "Operation-Location": "https://attacker.example.com/results/xyz", + } + raw_response.request = MagicMock() + raw_response.request.url = ( + "https://eastus.cognitiveservices.azure.com/documentintelligence/.../analyze" + ) + raw_response.request.headers = {"Ocp-Apim-Subscription-Key": "leak-me"} + + with pytest.raises(ValueError, match="rejected polling URL"): + config.transform_ocr_response( + model="azure-doc-intel", + raw_response=raw_response, + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + response={}, + ) + + +# ── Black Forest Labs polling ───────────────────────────────────────────────── + + +def test_bfl_image_generation_sync_rejects_cross_origin_polling(): + from litellm.llms.black_forest_labs.image_generation.handler import ( + BlackForestLabsImageGeneration, + ) + + handler = BlackForestLabsImageGeneration() + + initial_response = MagicMock() + initial_response.status_code = 200 + initial_response.json = MagicMock( + return_value={"polling_url": "https://attacker.example.com/get_result"} + ) + initial_response.request = MagicMock() + initial_response.request.url = "https://api.bfl.ai/v1/flux-pro" + + sync_client = MagicMock() + sync_client.get = MagicMock() + + with pytest.raises(Exception, match="Rejected polling URL"): + handler._poll_for_result_sync( + initial_response=initial_response, + headers={"x-key": "secret"}, + sync_client=sync_client, + ) + + sync_client.get.assert_not_called() + + +@pytest.mark.asyncio +async def test_bfl_image_generation_async_rejects_cross_origin_polling(): + from litellm.llms.black_forest_labs.image_generation.handler import ( + BlackForestLabsImageGeneration, + ) + + handler = BlackForestLabsImageGeneration() + + initial_response = MagicMock() + initial_response.status_code = 200 + initial_response.json = MagicMock( + return_value={"polling_url": "https://attacker.example.com/get_result"} + ) + initial_response.request = MagicMock() + initial_response.request.url = "https://api.bfl.ai/v1/flux-pro" + + async_client = MagicMock() + async_client.get = MagicMock() + + with pytest.raises(Exception, match="Rejected polling URL"): + await handler._poll_for_result_async( + initial_response=initial_response, + headers={"x-key": "secret"}, + async_client=async_client, + ) + + async_client.get.assert_not_called() + + +def test_bfl_image_edit_sync_rejects_cross_origin_polling(): + from litellm.llms.black_forest_labs.image_edit.handler import ( + BlackForestLabsImageEdit, + ) + + handler = BlackForestLabsImageEdit() + + initial_response = MagicMock() + initial_response.status_code = 200 + initial_response.json = MagicMock( + return_value={"polling_url": "https://attacker.example.com/get_result"} + ) + initial_response.request = MagicMock() + initial_response.request.url = "https://api.bfl.ai/v1/flux-pro/edit" + + sync_client = MagicMock() + sync_client.get = MagicMock() + + with pytest.raises(Exception, match="Rejected polling URL"): + handler._poll_for_result_sync( + initial_response=initial_response, + headers={"x-key": "secret"}, + sync_client=sync_client, + ) + + sync_client.get.assert_not_called() + + +def test_bfl_image_generation_same_origin_polling_passes(): + """Sanity check: when the polling URL shares origin with the original + request, the origin check passes and polling proceeds.""" + from litellm.llms.black_forest_labs.image_generation.handler import ( + BlackForestLabsImageGeneration, + ) + + handler = BlackForestLabsImageGeneration() + + initial_response = MagicMock() + initial_response.status_code = 200 + initial_response.json = MagicMock( + return_value={"polling_url": "https://api.bfl.ai/v1/get_result?id=abc"} + ) + initial_response.request = MagicMock() + initial_response.request.url = "https://api.bfl.ai/v1/flux-pro" + + sync_client = MagicMock() + poll_response = MagicMock() + poll_response.status_code = 200 + poll_response.json = MagicMock(return_value={"status": "Ready"}) + sync_client.get = MagicMock(return_value=poll_response) + + result = handler._poll_for_result_sync( + initial_response=initial_response, + headers={"x-key": "secret"}, + sync_client=sync_client, + ) + + sync_client.get.assert_called_once() + assert result is poll_response diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py index 780e7d2ba96..977c53280a9 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py @@ -86,6 +86,76 @@ def test_check_if_part_exists_in_parts_camel_case_snake_case(): assert check_if_part_exists_in_parts(parts_mixed, part_mixed_casing) +def test_cached_content_respects_modify_params_for_cache_incompatible_fields(): + """Regression: cachedContent drops system/tools/toolConfig only when modify_params=True.""" + import litellm + + cache_name = "projects/p/locations/us-central1/cachedContents/abc123" + messages = [ + {"role": "system", "content": "You are helpful"}, + {"role": "user", "content": "hi"}, + ] + optional_params = { + "tools": [ + { + "functionDeclarations": [ + {"name": "get_weather", "description": "Get weather"}, + ] + } + ], + "tool_choice": {"functionCallingConfig": {"mode": "AUTO"}}, + } + + original_modify_params = litellm.modify_params + try: + # With modify_params=False (default), keep fields even with cachedContent. + litellm.modify_params = False + result = _transform_request_body( + messages=list(messages), + model="gemini-2.5-pro", + optional_params=dict(optional_params), + custom_llm_provider="vertex_ai", + litellm_params={}, + cached_content=cache_name, + ) + assert result.get("cachedContent") == cache_name + assert "system_instruction" in result + assert "tools" in result + assert "toolConfig" in result + assert "contents" in result + + # With modify_params=True, drop cache-incompatible fields. + litellm.modify_params = True + result_modify_true = _transform_request_body( + messages=list(messages), + model="gemini-2.5-pro", + optional_params=dict(optional_params), + custom_llm_provider="vertex_ai", + litellm_params={}, + cached_content=cache_name, + ) + assert result_modify_true.get("cachedContent") == cache_name + assert "system_instruction" not in result_modify_true + assert "tools" not in result_modify_true + assert "toolConfig" not in result_modify_true + assert "contents" in result_modify_true + + # Without cache, fields are always included. + result_no_cache = _transform_request_body( + messages=list(messages), + model="gemini-2.5-pro", + optional_params=dict(optional_params), + custom_llm_provider="vertex_ai", + litellm_params={}, + cached_content=None, + ) + assert "system_instruction" in result_no_cache + assert "tools" in result_no_cache + assert "toolConfig" in result_no_cache + finally: + litellm.modify_params = original_modify_params + + # Tests for issue #14556: Labels field provider-aware filtering def test_google_genai_excludes_labels(): """Test that Google GenAI/AI Studio endpoints exclude labels when custom_llm_provider='gemini'""" diff --git a/tests/test_litellm/llms/vertex_ai/gemini_embeddings/__init__.py b/tests/test_litellm/llms/vertex_ai/gemini_embeddings/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py b/tests/test_litellm/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py new file mode 100644 index 00000000000..bb4e6c67e9e --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py @@ -0,0 +1,290 @@ +""" +Tests for Gemini batchEmbedContents transformation logic. + +Covers: +- Text-only inputs (single and batch) +- Multimodal inputs (data URIs, GCS URLs, file references) +- Mixed text + multimodal inputs +- Response processing with correct indices +""" + +import pytest + +from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import ( + _build_part_for_input, + _is_multimodal_input, + process_response, + transform_openai_input_gemini_content, + transform_openai_input_gemini_embed_content, +) +from litellm.types.llms.vertex_ai import VertexAIBatchEmbeddingsResponseObject +from litellm.types.utils import EmbeddingResponse + + +IMAGE_DATA_URI = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAgAAAAIAQMAAAD+wSzIAAAABlBMVEX///+/v7+jQ3Y5AAAADklEQVQI12P4AIX8EAgALgAD/aNpbtEAAAAASUVORK5CYII" +GCS_URL = "gs://my-bucket/image.png" + + +class TestIsMultimodalInput: + def test_text_only_string(self): + assert _is_multimodal_input("hello world") is False + + def test_text_only_list(self): + assert _is_multimodal_input(["hello", "world"]) is False + + def test_data_uri(self): + assert _is_multimodal_input([IMAGE_DATA_URI]) is True + + def test_gcs_url(self): + assert _is_multimodal_input([GCS_URL]) is True + + def test_file_reference(self): + assert _is_multimodal_input(["files/abc123"]) is True + + def test_mixed_text_and_image(self): + assert _is_multimodal_input(["hello", IMAGE_DATA_URI]) is True + + def test_nested_text_is_not_multimodal(self): + """Nested list with text is not multimodal.""" + assert _is_multimodal_input([["text_a", "text_b"]]) is False + + def test_nested_list_with_image_is_multimodal(self): + assert _is_multimodal_input([["a red shoe", IMAGE_DATA_URI]]) is True + + +class TestBuildPartForInput: + def test_text_input(self): + part = _build_part_for_input("hello") + assert part["text"] == "hello" + assert part.get("inline_data") is None + + def test_data_uri_input(self): + part = _build_part_for_input(IMAGE_DATA_URI) + assert part.get("text") is None + assert part["inline_data"] is not None + assert part["inline_data"]["mime_type"] == "image/png" + + def test_gcs_url_input(self): + part = _build_part_for_input(GCS_URL) + assert part.get("text") is None + assert part["file_data"] is not None + assert part["file_data"]["mime_type"] == "image/png" + assert part["file_data"]["file_uri"] == GCS_URL + + def test_file_reference_resolved(self): + resolved = {"files/abc": {"mime_type": "image/jpeg", "uri": "https://example.com/abc"}} + part = _build_part_for_input("files/abc", resolved_files=resolved) + assert part["file_data"] is not None + assert part["file_data"]["mime_type"] == "image/jpeg" + + def test_file_reference_unresolved_raises(self): + with pytest.raises(ValueError, match="not resolved"): + _build_part_for_input("files/abc") + + +class TestTransformOpenaiInputGeminiContent: + """Test that transform_openai_input_gemini_content creates separate requests per input.""" + + def test_single_text(self): + result = transform_openai_input_gemini_content( + input="hello", model="gemini-embedding-2-preview", optional_params={} + ) + assert len(result["requests"]) == 1 + assert result["requests"][0]["content"]["parts"][0]["text"] == "hello" + + def test_multiple_texts(self): + result = transform_openai_input_gemini_content( + input=["hello", "world"], model="gemini-embedding-2-preview", optional_params={} + ) + assert len(result["requests"]) == 2 + assert result["requests"][0]["content"]["parts"][0]["text"] == "hello" + assert result["requests"][1]["content"]["parts"][0]["text"] == "world" + + def test_multimodal_inputs_are_separate_requests(self): + """Key regression test for #24209: each input becomes its own request.""" + result = transform_openai_input_gemini_content( + input=["The food was delicious", IMAGE_DATA_URI], + model="gemini-embedding-2-preview", + optional_params={}, + ) + assert len(result["requests"]) == 2 + # First request is text + assert result["requests"][0]["content"]["parts"][0]["text"] == "The food was delicious" + # Second request is image + assert result["requests"][1]["content"]["parts"][0]["inline_data"] is not None + + def test_dimensions_mapped_to_output_dimensionality(self): + result = transform_openai_input_gemini_content( + input="hello", + model="gemini-embedding-2-preview", + optional_params={"dimensions": 256}, + ) + assert result["requests"][0]["outputDimensionality"] == 256 + + def test_model_name_prefixed(self): + result = transform_openai_input_gemini_content( + input="hello", model="gemini-embedding-2-preview", optional_params={} + ) + assert result["requests"][0]["model"] == "models/gemini-embedding-2-preview" + + def test_gcs_url_input(self): + result = transform_openai_input_gemini_content( + input=[GCS_URL], model="gemini-embedding-2-preview", optional_params={} + ) + assert len(result["requests"]) == 1 + assert result["requests"][0]["content"]["parts"][0]["file_data"] is not None + + def test_mixed_text_image_gcs(self): + result = transform_openai_input_gemini_content( + input=["hello", IMAGE_DATA_URI, GCS_URL], + model="gemini-embedding-2-preview", + optional_params={}, + ) + assert len(result["requests"]) == 3 + + def test_nested_input_combined_embedding(self): + """Nested list produces one request with multiple parts (combined embedding).""" + result = transform_openai_input_gemini_content( + input=[["a red shoe", IMAGE_DATA_URI]], + model="gemini-embedding-2-preview", + optional_params={}, + ) + assert len(result["requests"]) == 1 + parts = result["requests"][0]["content"]["parts"] + assert len(parts) == 2 + assert parts[0]["text"] == "a red shoe" + assert parts[1]["inline_data"] is not None + + def test_mixed_nested_and_flat(self): + """Mixed nested + flat produces correct number of requests.""" + result = transform_openai_input_gemini_content( + input=[["text", IMAGE_DATA_URI], "standalone"], + model="gemini-embedding-2-preview", + optional_params={}, + ) + assert len(result["requests"]) == 2 + # First: combined (2 parts) + assert len(result["requests"][0]["content"]["parts"]) == 2 + # Second: standalone (1 part) + assert len(result["requests"][1]["content"]["parts"]) == 1 + assert result["requests"][1]["content"]["parts"][0]["text"] == "standalone" + + +class TestTransformOpenaiInputGeminiEmbedContent: + """Test transform_openai_input_gemini_embed_content (vertex_ai / embedContent path).""" + + def test_text_and_image_combined(self): + result = transform_openai_input_gemini_embed_content( + input=["hello", IMAGE_DATA_URI], + model="gemini-embedding-2-preview", + optional_params={}, + ) + assert "content" in result + parts = result["content"]["parts"] + assert len(parts) == 2 + assert parts[0]["text"] == "hello" + assert parts[1]["inline_data"] is not None + + def test_gcs_url(self): + result = transform_openai_input_gemini_embed_content( + input=[GCS_URL], + model="gemini-embedding-2-preview", + optional_params={}, + ) + parts = result["content"]["parts"] + assert len(parts) == 1 + assert parts[0]["file_data"]["file_uri"] == GCS_URL + + def test_dimensions_mapped(self): + result = transform_openai_input_gemini_embed_content( + input="hello", + model="gemini-embedding-2-preview", + optional_params={"dimensions": 256}, + ) + assert result["outputDimensionality"] == 256 + + +class TestProcessResponse: + """Test that process_response sets correct indices.""" + + def test_single_embedding_index(self): + predictions: VertexAIBatchEmbeddingsResponseObject = { + "embeddings": [{"values": [0.1, 0.2]}] + } + model_response = EmbeddingResponse() + result = process_response( + input="hello", + model_response=model_response, + model="gemini-embedding-2-preview", + _predictions=predictions, + ) + assert len(result.data) == 1 + assert result.data[0]["index"] == 0 + + def test_multiple_embeddings_have_correct_indices(self): + """Regression test: indices should be 0, 1, 2... not all 0.""" + predictions: VertexAIBatchEmbeddingsResponseObject = { + "embeddings": [ + {"values": [0.1, 0.2]}, + {"values": [0.3, 0.4]}, + {"values": [0.5, 0.6]}, + ] + } + model_response = EmbeddingResponse() + result = process_response( + input=["a", "b", "c"], + model_response=model_response, + model="gemini-embedding-2-preview", + _predictions=predictions, + ) + assert len(result.data) == 3 + assert result.data[0]["index"] == 0 + assert result.data[1]["index"] == 1 + assert result.data[2]["index"] == 2 + + def test_multimodal_mixed_input(self): + """process_response works with mixed text + multimodal inputs.""" + predictions: VertexAIBatchEmbeddingsResponseObject = { + "embeddings": [{"values": [0.1, 0.2]}, {"values": [0.3, 0.4]}] + } + result = process_response( + input=["hello", IMAGE_DATA_URI], + model_response=EmbeddingResponse(), + model="gemini-embedding-2-preview", + _predictions=predictions, + ) + assert len(result.data) == 2 + assert result.data[0]["index"] == 0 + assert result.data[1]["index"] == 1 + # Should count tokens only for the text element, not the image + assert result.usage.prompt_tokens > 0 + + def test_nested_input_token_counting(self): + """Nested list: only plain-text sub-elements should be counted.""" + predictions: VertexAIBatchEmbeddingsResponseObject = { + "embeddings": [{"values": [0.1, 0.2]}] + } + result = process_response( + input=[["a red shoe", IMAGE_DATA_URI]], + model_response=EmbeddingResponse(), + model="gemini-embedding-2-preview", + _predictions=predictions, + ) + assert len(result.data) == 1 + assert result.usage.prompt_tokens > 0 + + def test_nested_empty_list_raises(self): + with pytest.raises(ValueError, match="must not be empty"): + transform_openai_input_gemini_content( + input=[[]], + model="gemini-embedding-2-preview", + optional_params={}, + ) + + def test_nested_non_string_element_raises(self): + with pytest.raises(ValueError, match="must be strings"): + transform_openai_input_gemini_content( + input=[[["doubly", "nested"]]], + model="gemini-embedding-2-preview", + optional_params={}, + ) diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_batch_transformation.py b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_batch_transformation.py index 33b3bfce44a..8a135dac3bb 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_batch_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_batch_transformation.py @@ -93,6 +93,51 @@ def test_vertex_ai_cancel_batch(): assert ":cancel" in call_args.kwargs["url"] +def test_vertex_ai_cancel_batch_encodes_batch_id(): + """Test that vertex_ai cancel_batch encodes user-controlled batch IDs.""" + handler = VertexAIBatchPrediction(gcs_bucket_name="test-bucket") + + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = { + "name": "projects/test-project/locations/us-central1/batchPredictionJobs/123456", + "state": "JOB_STATE_CANCELLING", + "createTime": "2024-03-17T10:00:00.000000Z", + "inputConfig": {"gcsSource": {"uris": ["gs://test-bucket/input.jsonl"]}}, + "outputConfig": { + "gcsDestination": {"outputUriPrefix": "gs://test-bucket/output"} + }, + } + + with patch( + "litellm.llms.vertex_ai.batches.handler._get_httpx_client" + ) as mock_client: + mock_client.return_value.post.return_value = mock_response + mock_client.return_value.get.return_value = mock_response + + with patch.object(handler, "_ensure_access_token") as mock_auth: + mock_auth.return_value = ("fake-token", "test-project") + + handler.cancel_batch( + _is_async=False, + batch_id="../../batchPredictionJobs/other?x=1#frag", + api_base=None, + vertex_credentials=None, + vertex_project="test-project", + vertex_location="us-central1", + timeout=600.0, + max_retries=None, + ) + + post_url = mock_client.return_value.post.call_args.kwargs["url"] + get_url = mock_client.return_value.get.call_args.kwargs["url"] + assert ( + "/..%2F..%2FbatchPredictionJobs%2Fother%3Fx%3D1%23frag:cancel" + in post_url + ) + assert "/..%2F..%2FbatchPredictionJobs%2Fother%3Fx%3D1%23frag" in get_url + + def test_vertex_ai_cancel_batch_forwards_timeout(): """Test that timeout is forwarded to the POST (cancel) HTTP call. diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py new file mode 100644 index 00000000000..b6329f33ae4 --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py @@ -0,0 +1,40 @@ +import pytest + +from litellm.llms.vertex_ai.vector_stores.search_api.transformation import ( + VertexSearchAPIVectorStoreConfig, +) + + +def test_should_encode_vertex_search_vector_store_id_in_complete_url(): + config = VertexSearchAPIVectorStoreConfig() + + url = config.get_complete_url( + api_base=None, + litellm_params={ + "vertex_project": "test-project", + "vertex_location": "global", + "vertex_collection_id": "default/collection", + "vector_store_id": "../../dataStores/other?x=1#frag", + }, + ) + + assert ( + url + == "https://discoveryengine.googleapis.com/v1/projects/test-project/locations/global/collections/default%2Fcollection/dataStores/..%2F..%2FdataStores%2Fother%3Fx%3D1%23frag/servingConfigs/default_config" + ) + + +def test_should_reject_dot_segment_vertex_search_vector_store_id(): + config = VertexSearchAPIVectorStoreConfig() + + with pytest.raises( + ValueError, match="vector_store_id cannot be a dot path segment" + ): + config.get_complete_url( + api_base=None, + litellm_params={ + "vertex_project": "test-project", + "vertex_location": "global", + "vector_store_id": "..", + }, + ) diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py b/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py new file mode 100644 index 00000000000..91261b63252 --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py @@ -0,0 +1,41 @@ +"""Vertex Model Garden: OpenAPI base URL for publisher/model ids vs per-endpoint path.""" + +import pytest + +from litellm.llms.vertex_ai.vertex_model_garden.main import ( + _vertex_model_garden_model_id_in_json_body, + create_vertex_url, +) + + +@pytest.mark.parametrize( + "model,expect_openapi_base", + [ + ("xai/grok-4.1-fast-reasoning", True), + ("openai/foo/bar", True), + ("5464397967697903616", False), + ("gpt-oss-20b-maas", False), + ], +) +def test_create_vertex_url_openapi_vs_deployed_endpoint( + model: str, expect_openapi_base: bool +) -> None: + url = create_vertex_url( + vertex_location="us-central1", + vertex_project="my-project", + stream=False, + model=model, + ) + if expect_openapi_base: + assert "/v1/projects/my-project/locations/us-central1/endpoints/openapi" in url + else: + assert ( + "/v1beta1/projects/my-project/locations/us-central1/endpoints/" + f"{model}" in url + ) + assert "openapi" not in url + + +def test_model_id_in_json_body_heuristic() -> None: + assert _vertex_model_garden_model_id_in_json_body("xai/grok-4.1-fast-reasoning") is True + assert _vertex_model_garden_model_id_in_json_body("5464397967697903616") is False diff --git a/tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py b/tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py index 623d162ccfe..13571e63c7d 100644 --- a/tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py +++ b/tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py @@ -101,6 +101,23 @@ class TestVolcengineResponsesAPITransformation: ) assert api_base_full == "https://custom.volc.com/api/v3/responses" + def test_response_id_path_requests_encode_response_id(self): + """response_id should be encoded before building Volcengine URLs.""" + config = VolcEngineResponsesAPIConfig() + + url, params = config.transform_cancel_response_api_request( + response_id="../../responses/other?x=1#frag", + api_base="https://custom.volc.com/api/v3/responses", + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert ( + url + == "https://custom.volc.com/api/v3/responses/..%2F..%2Fresponses%2Fother%3Fx%3D1%23frag/cancel" + ) + assert params == {} + @pytest.mark.parametrize( "litellm_params, expected_key", [ diff --git a/tests/test_litellm/llms/xai/test_xai_chat_transformation.py b/tests/test_litellm/llms/xai/test_xai_chat_transformation.py new file mode 100644 index 00000000000..5a236de900e --- /dev/null +++ b/tests/test_litellm/llms/xai/test_xai_chat_transformation.py @@ -0,0 +1,39 @@ +import os +import sys + +sys.path.insert( + 0, os.path.abspath("../../../..") +) # Adds the parent directory to the system path + +from litellm.llms.xai.chat.transformation import XAIChatConfig + + +class TestXAIParallelToolCalls: + """Test suite for XAI parallel tool calls functionality.""" + + def test_get_supported_openai_params_includes_parallel_tool_calls(self): + """Test that parallel_tool_calls is in supported parameters.""" + config = XAIChatConfig() + supported_params = config.get_supported_openai_params( + "xai/grok-4.20" + ) + assert "parallel_tool_calls" in supported_params + + def test_transform_request_preserves_parallel_tool_calls(self): + """Test that transform_request preserves parallel_tool_calls parameter.""" + config = XAIChatConfig() + + messages = [{"role": "user", "content": "What's the weather like?"}] + optional_params = {"parallel_tool_calls": True} + + result = config.transform_request( + model="xai/grok-4.20", + messages=messages, + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + + assert result.get("parallel_tool_calls") is True + assert len(result["messages"]) == 1 + assert result["messages"][0]["role"] == "user" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 558c677d2dc..85d5d6ba466 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -23,6 +23,21 @@ def mock_mcp_client_ip(): yield +@pytest.fixture +def trust_xff(): + """Force ``IPAddressUtils.is_request_from_trusted_proxy`` to True. + + Tests that exercise X-Forwarded-* parsing logic opt into this fixture. + The trust gate's own behaviour is covered by + ``test_get_request_base_url_xff_trust_gate``. + """ + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.IPAddressUtils.is_request_from_trusted_proxy", + return_value=True, + ): + yield + + @pytest.mark.asyncio async def test_authorize_endpoint_includes_response_type(): """Test that authorize endpoint includes response_type=code parameter (fixes #15684)""" @@ -505,6 +520,7 @@ async def test_register_client_remote_registration_success(): @pytest.mark.asyncio +@pytest.mark.usefixtures("trust_xff") async def test_authorize_endpoint_respects_x_forwarded_proto(): """Test that authorize endpoint uses X-Forwarded-Proto header to construct correct redirect_uri""" try: @@ -572,6 +588,7 @@ async def test_authorize_endpoint_respects_x_forwarded_proto(): @pytest.mark.asyncio +@pytest.mark.usefixtures("trust_xff") async def test_token_endpoint_respects_x_forwarded_proto(): """Test that token endpoint uses X-Forwarded-Proto header for redirect_uri""" try: @@ -650,6 +667,7 @@ async def test_token_endpoint_respects_x_forwarded_proto(): @pytest.mark.asyncio +@pytest.mark.usefixtures("trust_xff") async def test_oauth_protected_resource_respects_x_forwarded_proto(): """Test that oauth_protected_resource_mcp uses X-Forwarded-Proto for URLs""" try: @@ -704,6 +722,7 @@ async def test_oauth_protected_resource_respects_x_forwarded_proto(): @pytest.mark.asyncio +@pytest.mark.usefixtures("trust_xff") async def test_oauth_authorization_server_respects_x_forwarded_proto(): """Test that oauth_authorization_server_mcp uses X-Forwarded-Proto for URLs""" try: @@ -759,6 +778,7 @@ async def test_oauth_authorization_server_respects_x_forwarded_proto(): @pytest.mark.asyncio +@pytest.mark.usefixtures("trust_xff") async def test_register_client_respects_x_forwarded_proto(): """Test that register_client uses X-Forwarded-Proto for redirect_uris""" try: @@ -796,6 +816,7 @@ async def test_register_client_respects_x_forwarded_proto(): @pytest.mark.asyncio +@pytest.mark.usefixtures("trust_xff") async def test_authorize_endpoint_respects_x_forwarded_host(): """Test that authorize endpoint uses X-Forwarded-Host and X-Forwarded-Proto to construct correct redirect_uri""" try: @@ -869,6 +890,7 @@ async def test_authorize_endpoint_respects_x_forwarded_host(): @pytest.mark.asyncio +@pytest.mark.usefixtures("trust_xff") async def test_token_endpoint_respects_x_forwarded_host(): """Test that token endpoint uses X-Forwarded-Host and X-Forwarded-Proto for redirect_uri""" try: @@ -1071,7 +1093,12 @@ async def test_token_endpoint_respects_x_forwarded_host(): def test_get_request_base_url_comprehensive( base_url, x_forwarded_proto, x_forwarded_host, x_forwarded_port, expected_url ): - """Comprehensive test for get_request_base_url with various header combinations""" + """Comprehensive test for get_request_base_url with various header combinations. + + These cases exercise the X-Forwarded-* parsing logic, so the trust gate + is patched True; the gate's own behaviour is covered by the + ``test_get_request_base_url_xff_trust_gate`` matrix below. + """ try: from fastapi import Request @@ -1081,11 +1108,9 @@ def test_get_request_base_url_comprehensive( except ImportError: pytest.skip("MCP discoverable endpoints not available") - # Create mock request mock_request = MagicMock(spec=Request) mock_request.base_url = base_url - # Build headers dict headers = {} if x_forwarded_proto: headers["X-Forwarded-Proto"] = x_forwarded_proto @@ -1094,16 +1119,17 @@ def test_get_request_base_url_comprehensive( if x_forwarded_port: headers["X-Forwarded-Port"] = x_forwarded_port - # Mock headers.get() to return our test values def mock_get(header_name, default=None): return headers.get(header_name, default) mock_request.headers.get = mock_get - # Test the function - result = get_request_base_url(mock_request) + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.IPAddressUtils.is_request_from_trusted_proxy", + return_value=True, + ): + result = get_request_base_url(mock_request) - # Verify result assert result == expected_url, ( f"Expected '{expected_url}' but got '{result}'\n" f"Input: base_url={base_url}, " @@ -1113,6 +1139,131 @@ def test_get_request_base_url_comprehensive( ) +@pytest.mark.parametrize( + "general_settings,direct_ip,expect_xff_honoured", + [ + # Default: use_x_forwarded_for not set -> ignore X-Forwarded-* entirely. + ({}, "127.0.0.1", False), + # XFF enabled, no trusted ranges -> still ignored (no way to tell a trusted + # reverse proxy from a direct attacker). + ({"use_x_forwarded_for": True}, "127.0.0.1", False), + # XFF enabled, ranges set, but caller IP outside any range -> ignored. + ( + { + "use_x_forwarded_for": True, + "mcp_trusted_proxy_ranges": ["10.0.0.0/8"], + }, + "203.0.113.5", + False, + ), + # XFF enabled, caller in trusted range -> headers honoured. + ( + { + "use_x_forwarded_for": True, + "mcp_trusted_proxy_ranges": ["10.0.0.0/8"], + }, + "10.0.0.7", + True, + ), + # Loopback example (common dev / single-host deploy). + ( + { + "use_x_forwarded_for": True, + "mcp_trusted_proxy_ranges": ["127.0.0.0/8"], + }, + "127.0.0.1", + True, + ), + ], +) +def test_get_request_base_url_xff_trust_gate( + general_settings, direct_ip, expect_xff_honoured +): + """Verify the X-Forwarded-* trust gate. + + With XFF poisoning attempted, the helper must return either the literal + base_url (gate denies) or the forwarded URL (gate allows), never the + forwarded URL when the gate denies. + """ + try: + from fastapi import Request + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + get_request_base_url, + ) + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://localhost:4000/" + mock_request.client = MagicMock() + mock_request.client.host = direct_ip + + headers = { + "X-Forwarded-Proto": "https", + "X-Forwarded-Host": "attacker.example.com", + } + mock_request.headers.get = lambda name, default=None: headers.get(name, default) + mock_request.headers.__contains__ = lambda self_, name: name in headers + + with patch( + "litellm.proxy.proxy_server.general_settings", + general_settings, + create=True, + ): + result = get_request_base_url(mock_request) + + if expect_xff_honoured: + assert result == "https://attacker.example.com" + else: + assert result == "http://localhost:4000" + + +def test_xff_misconfig_warning_emitted_once(caplog): + """Operators upgrading from the old "always trust X-Forwarded-*" behaviour + get a one-shot warning when they have ``use_x_forwarded_for`` enabled + but no ``mcp_trusted_proxy_ranges`` configured. The warning must NOT + spam every request.""" + try: + from fastapi import Request + + from litellm.proxy import auth as proxy_auth_pkg # noqa: F401 + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + get_request_base_url, + ) + from litellm.proxy.auth import ip_address_utils + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Reset the module-level one-shot flag so the test is deterministic. + ip_address_utils._warned_xff_without_trusted_ranges = False + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://localhost:4000/" + mock_request.client = MagicMock() + mock_request.client.host = "203.0.113.5" + headers = {"X-Forwarded-Host": "attacker.example.com"} + mock_request.headers.get = lambda name, default=None: headers.get(name, default) + + misconfig = {"use_x_forwarded_for": True} + + import logging + + with ( + caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"), + patch("litellm.proxy.proxy_server.general_settings", misconfig, create=True), + ): + for _ in range(3): + get_request_base_url(mock_request) + + matching = [ + rec for rec in caplog.records if "mcp_trusted_proxy_ranges" in rec.getMessage() + ] + assert ( + len(matching) == 1 + ), f"expected exactly one warning, got {len(matching)}: {[r.getMessage() for r in matching]}" + + # ------------------------------------------------------------------- # Tests for scopes_supported when mcp_server.scopes is None # ------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py index 84c556b8ddc..649a08e8744 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py @@ -673,6 +673,106 @@ class TestHookHeaderMergePriority: assert headers["X-OAuth"] == "yes" assert headers["X-Trace-Id"] == "trace-123" + @pytest.mark.asyncio + async def test_m2m_oauth2_does_not_forward_litellm_caller_authorization(self): + """M2M must not put caller Bearer (LiteLLM API key) into extra_headers (#23652).""" + manager = MCPServerManager() + server = MCPServer( + server_id="test-id", + name="Test Server", + server_name="test_server", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + token_url="https://auth.example.com/token", + ) + + captured_extra_headers: Dict[str, Any] = {} + + async def fake_create_mcp_client( + server, mcp_auth_header=None, extra_headers=None, stdio_env=None + ): + captured_extra_headers["value"] = extra_headers + mock_client = MagicMock() + mock_client.call_tool = AsyncMock(return_value=MagicMock()) + return mock_client + + with patch.object( + manager, "_create_mcp_client", side_effect=fake_create_mcp_client + ): + with patch.object(manager, "_build_stdio_env", return_value=None): + try: + await manager._call_regular_mcp_tool( + mcp_server=server, + original_tool_name="test_tool", + arguments={"key": "val"}, + tasks=[], + mcp_auth_header=None, + mcp_server_auth_headers=None, + oauth2_headers={"Authorization": "Bearer sk-1234"}, + raw_headers={"authorization": "Bearer sk-1234"}, + proxy_logging_obj=None, + hook_extra_headers=None, + ) + except Exception: + pass + + assert captured_extra_headers.get("value") is None + + @pytest.mark.asyncio + async def test_m2m_oauth2_skips_authorization_in_configured_extra_headers(self): + """M2M must not take Authorization from raw_headers even if extra_headers lists it.""" + manager = MCPServerManager() + server = MCPServer( + server_id="test-id", + name="Test Server", + server_name="test_server", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + token_url="https://auth.example.com/token", + extra_headers=["Authorization", "X-Custom"], + ) + + captured_extra_headers: Dict[str, Any] = {} + + async def fake_create_mcp_client( + server, mcp_auth_header=None, extra_headers=None, stdio_env=None + ): + captured_extra_headers["value"] = extra_headers + mock_client = MagicMock() + mock_client.call_tool = AsyncMock(return_value=MagicMock()) + return mock_client + + with patch.object( + manager, "_create_mcp_client", side_effect=fake_create_mcp_client + ): + with patch.object(manager, "_build_stdio_env", return_value=None): + try: + await manager._call_regular_mcp_tool( + mcp_server=server, + original_tool_name="test_tool", + arguments={"key": "val"}, + tasks=[], + mcp_auth_header=None, + mcp_server_auth_headers=None, + oauth2_headers={"Authorization": "Bearer sk-1234"}, + raw_headers={ + "authorization": "Bearer sk-1234", + "x-custom": "from-client", + }, + proxy_logging_obj=None, + hook_extra_headers=None, + ) + except Exception: + pass + + headers = captured_extra_headers.get("value") or {} + assert "Authorization" not in headers + assert headers.get("X-Custom") == "from-client" + class TestUserAPIKeyAuthJwtClaims: """Tests that UserAPIKeyAuth correctly carries jwt_claims.""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 9df6408b0d7..06f95159c08 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -17,6 +17,7 @@ from litellm.proxy._types import ( MCPTransport, UserAPIKeyAuth, ) +from litellm.types.mcp import MCPAuth from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -135,6 +136,152 @@ def test_prepare_mcp_server_headers_case_insensitive_extra_headers(): assert extra_headers == {"Authorization": "Bearer token"} +def test_prepare_mcp_server_headers_oauth2_m2m_omits_litellm_caller_authorization(): + """M2M OAuth must not put caller Bearer (LiteLLM API key) into extra_headers (#23652).""" + try: + from litellm.proxy._experimental.mcp_server.server import ( + _prepare_mcp_server_headers, + ) + except ImportError: + pytest.skip("MCP server not available") + + server = MCPServer( + server_id="m2m-server", + name="m2m", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + token_url="https://auth.example.com/token", + ) + caller_key = {"Authorization": "Bearer sk-litellm-caller"} + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=None, + mcp_auth_header=None, + oauth2_headers=caller_key, + raw_headers=None, + ) + + assert server_auth_header is None + assert extra_headers is None + + +def test_prepare_mcp_server_headers_oauth2_interactive_copies_oauth2_headers(): + """Interactive OAuth still forwards the user's OAuth token in extra_headers.""" + try: + from litellm.proxy._experimental.mcp_server.server import ( + _prepare_mcp_server_headers, + ) + except ImportError: + pytest.skip("MCP server not available") + + user_oauth = {"Authorization": "Bearer upstream-user-token"} + + server = MCPServer( + server_id="3lo-server", + name="3lo", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow=None, + ) + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=None, + mcp_auth_header=None, + oauth2_headers=user_oauth, + raw_headers=None, + ) + + assert server_auth_header is None + assert extra_headers == user_oauth + + +def test_prepare_mcp_server_headers_m2m_skips_authorization_from_raw_extra_headers(): + """M2M must not merge caller Authorization from raw_headers when extra_headers lists it.""" + try: + from litellm.proxy._experimental.mcp_server.server import ( + _prepare_mcp_server_headers, + ) + except ImportError: + pytest.skip("MCP server not available") + + server = MCPServer( + server_id="m2m-raw", + name="m2m", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + token_url="https://auth.example.com/token", + extra_headers=["Authorization", "X-Custom"], + ) + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=None, + mcp_auth_header=None, + oauth2_headers={"Authorization": "Bearer sk-1234"}, + raw_headers={ + "authorization": "Bearer sk-1234", + "x-custom": "trace", + }, + ) + + assert server_auth_header is None + assert extra_headers is not None + assert "Authorization" not in extra_headers + assert extra_headers.get("X-Custom") == "trace" + + +@pytest.mark.asyncio +async def test_call_tool_m2m_skips_authorization_headers(): + """M2M call_tool must not forward caller Authorization in oauth2/raw headers.""" + try: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + ) + except ImportError: + pytest.skip("MCP server not available") + + manager = MCPServerManager() + server = MCPServer( + server_id="m2m-call-tool", + name="m2m-call-tool", + server_name="m2m-call-tool", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + token_url="https://auth.example.com/token", + client_id="cid", + client_secret="csecret", + extra_headers=["Authorization", "X-Custom"], + ) + + mock_client = MagicMock() + mock_client.call_tool = AsyncMock(return_value=MagicMock()) + + with patch.object( + manager, "_create_mcp_client", new=AsyncMock(return_value=mock_client) + ) as create_client_mock: + await manager._call_regular_mcp_tool( + mcp_server=server, + original_tool_name="echo", + arguments={"message": "hello"}, + tasks=[], + mcp_auth_header=None, + mcp_server_auth_headers=None, + oauth2_headers={"Authorization": "Bearer sk-1234"}, + raw_headers={"authorization": "Bearer sk-1234", "x-custom": "trace"}, + proxy_logging_obj=None, + ) + + create_kwargs = create_client_mock.await_args.kwargs + extra_headers = create_kwargs["extra_headers"] or {} + assert "Authorization" not in extra_headers + assert extra_headers.get("X-Custom") == "trace" + + @pytest.mark.asyncio async def test_get_prompts_from_mcp_servers_success(): try: @@ -2288,6 +2435,79 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab assert spend_meta["per_server_tool_counts"]["server_a"] == 1 +@pytest.mark.asyncio +async def test_get_tools_from_mcp_servers_returns_tools_when_success_logging_fails(): + """ + Regression test: list_tools should still return fetched tools even if + async_success_handler raises (e.g. serialization errors in logging path). + """ + try: + from litellm.proxy._experimental.mcp_server.server import ( + _get_tools_from_mcp_servers, + ) + from litellm.proxy._types import UserAPIKeyAuth + except ImportError: + pytest.skip("MCP server not available") + + user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user") + + server_a = MagicMock(name="server_a_obj") + server_a.name = "server_a" + server_a.alias = "server_a" + server_a.server_name = "server_a" + server_a.server_id = "a" + server_a.auth_type = None + server_a.extra_headers = None + + tool_1 = MagicMock() + tool_1.name = "server_a-tool_1" + + dummy_logging_obj = MagicMock() + dummy_logging_obj.model_call_details = {"metadata": {"spend_logs_metadata": {}}} + dummy_logging_obj.async_success_handler = AsyncMock( + side_effect=TypeError("Object of type Tool is not JSON serializable") + ) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server_a]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", + return_value=(None, None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + ) as mock_manager, + patch( + "litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools", + side_effect=lambda tools, _server: tools, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions", + new=AsyncMock(side_effect=lambda tools, **_: tools), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.function_setup", + return_value=(dummy_logging_obj, None), + ), + ): + mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1]) + + tools = await _get_tools_from_mcp_servers( + user_api_key_auth=user_auth, + mcp_auth_header=None, + mcp_servers=["server_a"], + mcp_server_auth_headers=None, + log_list_tools_to_spendlogs=True, + list_tools_log_source="mcp_protocol", + ) + + assert tools == [tool_1] + dummy_logging_obj.async_success_handler.assert_awaited_once() + + def test_tool_name_matches_case_insensitive(): """Test that _tool_name_matches performs case-insensitive comparison. @@ -2719,3 +2939,177 @@ class TestGatewayCreateInitializationOptions: _mcp_gateway_initialize_instructions.reset(tok) opts = server.create_initialization_options() assert getattr(opts, "instructions", None) is None + + +@pytest.mark.asyncio +async def test_list_tools_with_legacy_db_m2m_server_resolves_oauth2_flow(): + """ + P1 Regression: list_tools path must apply _resolve_oauth2_flow to legacy DB + rows where oauth2_flow is NULL but M2M credentials are present. + + Without this fix, has_client_credentials returns False and the caller's + Authorization header is forwarded upstream instead of being blocked. + """ + try: + from litellm.proxy._experimental.mcp_server.server import ( + _get_tools_from_mcp_servers, + ) + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.mcp import MCPAuth + except ImportError: + pytest.skip("MCP server not available") + + user_auth = UserAPIKeyAuth(api_key="sk-1234", user_id="test-user") + + # Simulate a legacy DB row: OAuth2 with M2M credentials but oauth2_flow=None + legacy_server = MagicMock(name="legacy_m2m_server") + legacy_server.name = "legacy_m2m" + legacy_server.alias = "legacy_m2m" + legacy_server.server_name = "legacy_m2m" + legacy_server.server_id = "legacy-m2m-id" + legacy_server.auth_type = MCPAuth.oauth2 + legacy_server.oauth2_flow = None # Legacy: field not set in DB + legacy_server.token_url = "https://oauth.example.com/token" + legacy_server.authorization_url = None + legacy_server.client_id = "client-id" + legacy_server.client_secret = "client-secret" + legacy_server.extra_headers = None + legacy_server.has_client_credentials = False # This is the bug: should be True + legacy_server.model_copy = MagicMock( + side_effect=lambda update: MCPServer( + server_id=legacy_server.server_id, + name=legacy_server.name, + transport=MCPTransport.http, + auth_type=legacy_server.auth_type, + oauth2_flow=update.get("oauth2_flow", legacy_server.oauth2_flow), + token_url=legacy_server.token_url, + authorization_url=legacy_server.authorization_url, + client_id=legacy_server.client_id, + client_secret=legacy_server.client_secret, + ) + ) + + tool_1 = MagicMock() + tool_1.name = "legacy_m2m-tool" + + captured_extra_headers = None + + async def capture_extra_headers(*args, **kwargs): + nonlocal captured_extra_headers + captured_extra_headers = kwargs.get("extra_headers") + return [tool_1] + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + ) as mock_manager, + patch( + "litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools", + side_effect=lambda tools, _server: tools, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions", + new=AsyncMock(side_effect=lambda tools, **_: tools), + ), + ): + mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["legacy-m2m-id"]) + mock_manager.get_mcp_server_by_id = MagicMock(return_value=legacy_server) + mock_manager.filter_server_ids_by_ip_with_info = MagicMock( + return_value=(["legacy-m2m-id"], 0) + ) + mock_manager._get_tools_from_server = AsyncMock( + side_effect=capture_extra_headers + ) + + tools = await _get_tools_from_mcp_servers( + user_api_key_auth=user_auth, + mcp_auth_header=None, + mcp_servers=["legacy_m2m"], + mcp_server_auth_headers=None, + oauth2_headers={"Authorization": "Bearer sk-1234"}, # Caller's token + ) + + # With P1 fix: _get_allowed_mcp_servers applies _resolve_oauth2_flow, + # so has_client_credentials becomes True and extra_headers should be None + # (caller's Authorization blocked) + assert captured_extra_headers is None, ( + "P1 security issue: caller's Authorization header was forwarded to M2M server. " + "Expected None, got: " + str(captured_extra_headers) + ) + assert tools == [tool_1] + + +@pytest.mark.asyncio +async def test_call_tool_empty_extra_headers_returns_none(): + """ + P2 Regression: When all configured extra_headers are filtered out (e.g. + Authorization for M2M), the resulting extra_headers should be None, not {}. + + Downstream code that checks `if extra_headers is None` will behave + differently if an empty dict is passed instead. + """ + try: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + ) + from litellm.types.mcp import MCPAuth + except ImportError: + pytest.skip("MCP server not available") + + manager = MCPServerManager() + + # M2M server with only Authorization in extra_headers + m2m_server = MCPServer( + server_id="m2m-srv", + name="m2m_test", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + token_url="https://oauth.example.com/token", + client_id="client-id", + client_secret="client-secret", + extra_headers=["Authorization"], # Will be filtered out for M2M + ) + + raw_headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"} + + captured_extra_headers = None + + async def capture_create_mcp_client(*args, **kwargs): + nonlocal captured_extra_headers + captured_extra_headers = kwargs.get("extra_headers") + # Return a mock client + mock_client = AsyncMock() + mock_client.call_tool = AsyncMock(return_value=MagicMock(content=[])) + return mock_client + + with ( + patch.object( + manager, + "_create_mcp_client", + side_effect=capture_create_mcp_client, + ), + patch.object( + manager, + "get_mcp_server_by_id", + return_value=m2m_server, + ), + ): + try: + await manager._call_regular_mcp_tool( + mcp_server=m2m_server, + original_tool_name="test_tool", + arguments={}, + mcp_auth_header=None, + oauth2_headers=None, + raw_headers=raw_headers, + ) + except Exception: + pass # We only care about the captured headers + + # With P2 fix: extra_headers should be None (not {}) when all headers filtered + assert captured_extra_headers is None, ( + "P2 API consistency issue: expected None for empty extra_headers, got: " + + str(captured_extra_headers) + ) + diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index a848db27fc1..4c21d0ec645 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -922,19 +922,21 @@ async def test_get_tag_objects_batch(): # Simulate 5 tags: 2 cached, 3 uncached tag_names = ["cached-1", "uncached-1", "cached-2", "uncached-2", "uncached-3"] - # Mock cached tags - cached_tag_1 = { - "tag_name": "cached-1", - "spend": 10.0, - "models": [], - "litellm_budget_table": None, - } - cached_tag_2 = { - "tag_name": "cached-2", - "spend": 20.0, - "models": [], - "litellm_budget_table": None, - } + # Mock cached tags — must be LiteLLM_TagTable instances: the mocked async_get_cache + # bypasses UserApiKeyCache deserialization, so returning plain dicts would flow through + # as dict (production returns models after Codec.deserialize inside the cache). + cached_tag_1 = LiteLLM_TagTable( + tag_name="cached-1", + spend=10.0, + models=[], + litellm_budget_table=None, + ) + cached_tag_2 = LiteLLM_TagTable( + tag_name="cached-2", + spend=20.0, + models=[], + litellm_budget_table=None, + ) # Mock DB response for uncached tags uncached_tag_1 = MagicMock() @@ -980,13 +982,13 @@ async def test_get_tag_objects_batch(): ) # Mock cache behavior - return cached tags, None for uncached - async def mock_get_cache(key): + async def mock_get_cache(*args, **kwargs): + key = kwargs.get("key") if key == "tag:cached-1": return cached_tag_1 - elif key == "tag:cached-2": + if key == "tag:cached-2": return cached_tag_2 - else: - return None + return None mock_cache.async_get_cache = AsyncMock(side_effect=mock_get_cache) mock_cache.async_set_cache = AsyncMock() diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 91f300b88ce..b82cb355192 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -964,3 +964,129 @@ class TestIsRequestBodySafeBlocksEndpointTargetingFields: ) is True ) + + +# ── is_request_body_safe nested-config recursion (VERIA-6) ──────────────────── + + +class TestIsRequestBodySafeNestedConfig: + """The Milvus vector store transformer unpacks + ``litellm_embedding_config`` as ``**kwargs`` into ``litellm.embedding(...)`` + — same SSRF / credential-exfil surface as a top-level ``api_base`` in + the request body. ``is_request_body_safe`` must recurse into this + nested dict so a banned param can't be smuggled in via nesting.""" + + def test_root_level_api_base_blocked_when_no_opt_in(self): + """Sanity check: pre-existing root-level enforcement still works.""" + with pytest.raises(ValueError, match="api_base"): + is_request_body_safe( + request_body={"api_base": "https://attacker.example.com"}, + general_settings={}, + llm_router=None, + model="gpt-4", + ) + + def test_nested_api_base_in_embedding_config_blocked(self): + """Smuggling ``api_base`` inside ``litellm_embedding_config`` is + the VERIA-6 bypass — must be blocked by the recursive check.""" + with pytest.raises(ValueError, match="api_base"): + is_request_body_safe( + request_body={ + "litellm_embedding_config": { + "api_base": "https://attacker.example.com", + "api_key": "leaked-key", + } + }, + general_settings={}, + llm_router=None, + model="milvus-store", + ) + + def test_nested_langfuse_host_in_embedding_config_blocked(self): + """The recursion uses the *full* banned-param list, not a special + subset — so any flag that's banned at the root is also banned + when nested.""" + with pytest.raises(ValueError, match="langfuse_host"): + is_request_body_safe( + request_body={ + "litellm_embedding_config": { + "langfuse_host": "https://attacker.example.com" + } + }, + general_settings={}, + llm_router=None, + model="milvus-store", + ) + + def test_nested_api_base_allowed_when_admin_opts_in(self): + """Admins who explicitly enable client-side credential passthrough + keep the existing escape hatch — same UX as for root-level.""" + assert ( + is_request_body_safe( + request_body={ + "litellm_embedding_config": { + "api_base": "https://my-azure.example.com" + } + }, + general_settings={"allow_client_side_credentials": True}, + llm_router=None, + model="milvus-store", + ) + is True + ) + + def test_safe_nested_config_accepted(self): + """A nested config without any banned params passes — there's no + false-positive on legitimate ``api_version`` / model params.""" + assert ( + is_request_body_safe( + request_body={ + "litellm_embedding_config": { + "api_version": "2024-02-15-preview", + } + }, + general_settings={}, + llm_router=None, + model="milvus-store", + ) + is True + ) + + def test_non_dict_nested_config_does_not_break_check(self): + """A bogus type for ``litellm_embedding_config`` (string, list, + None) must not crash the validator — it should just fall through.""" + assert ( + is_request_body_safe( + request_body={"litellm_embedding_config": "not-a-dict"}, + general_settings={}, + llm_router=None, + model="x", + ) + is True + ) + + def test_deeply_nested_config_does_not_recurse(self): + """Greptile P1: ``is_request_body_safe`` is iterative single-level — + a deeply-nested ``litellm_embedding_config`` cannot exhaust the + Python call stack to trigger a 500 ``RecursionError``. Build a + body 1000 levels deep; the validator must complete in O(1) + descent.""" + body = {"litellm_embedding_config": {}} + cur = body["litellm_embedding_config"] + for _ in range(1000): + cur["litellm_embedding_config"] = {} + cur = cur["litellm_embedding_config"] + # Banned param at the deepest level shouldn't be reached — single + # level only. + cur["api_base"] = "https://attacker.example.com" + + # No exception raised: deeper levels aren't checked. + assert ( + is_request_body_safe( + request_body=body, + general_settings={}, + llm_router=None, + model="x", + ) + is True + ) diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index 9085469268c..47e513dc593 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -405,9 +405,9 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_role_change(): mock_cache.async_set_cache.assert_called_once() call_kwargs = mock_cache.async_set_cache.call_args assert call_kwargs.kwargs["key"] == "u1" - assert ( - call_kwargs.kwargs["value"]["user_role"] == LitellmUserRoles.PROXY_ADMIN.value - ) + assert isinstance(call_kwargs.kwargs["value"], LiteLLM_UserTable) + assert call_kwargs.kwargs["value"].user_role == LitellmUserRoles.PROXY_ADMIN.value + assert call_kwargs.kwargs["model_type"] == LiteLLM_UserTable @pytest.mark.asyncio @@ -452,7 +452,9 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_team_change(): mock_cache.async_set_cache.assert_called_once() call_kwargs = mock_cache.async_set_cache.call_args assert call_kwargs.kwargs["key"] == "u1" - assert set(call_kwargs.kwargs["value"]["teams"]) == {"team1", "team2"} + assert isinstance(call_kwargs.kwargs["value"], LiteLLM_UserTable) + assert set(call_kwargs.kwargs["value"].teams) == {"team1", "team2"} + assert call_kwargs.kwargs["model_type"] == LiteLLM_UserTable @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/client/cli/test_auth_commands.py b/tests/test_litellm/proxy/client/cli/test_auth_commands.py index f7cb4d72d91..2e738ff900d 100644 --- a/tests/test_litellm/proxy/client/cli/test_auth_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_auth_commands.py @@ -231,6 +231,50 @@ class TestTokenUtilities: result = get_stored_api_key() assert result is None + def test_get_stored_api_key_base_url_match(self): + """Stored key is returned when expected_base_url matches stored origin""" + token_data = {"key": "sk-prod", "base_url": "https://real-proxy.com"} + with patch( + "litellm.litellm_core_utils.cli_token_utils.load_cli_token", + return_value=token_data, + ): + assert ( + get_stored_api_key(expected_base_url="https://real-proxy.com") + == "sk-prod" + ) + + def test_get_stored_api_key_base_url_match_trailing_slash(self): + """Trailing slash on expected_base_url is normalised before comparison""" + token_data = {"key": "sk-prod", "base_url": "https://real-proxy.com"} + with patch( + "litellm.litellm_core_utils.cli_token_utils.load_cli_token", + return_value=token_data, + ): + assert ( + get_stored_api_key(expected_base_url="https://real-proxy.com/") + == "sk-prod" + ) + + def test_get_stored_api_key_base_url_mismatch(self): + """Stored key is NOT returned when expected_base_url differs from stored origin""" + token_data = {"key": "sk-prod", "base_url": "https://real-proxy.com"} + with patch( + "litellm.litellm_core_utils.cli_token_utils.load_cli_token", + return_value=token_data, + ): + assert get_stored_api_key(expected_base_url="https://evil.com") is None + + def test_get_stored_api_key_old_token_no_base_url(self): + """Old tokens without a base_url field are rejected when origin check is requested""" + token_data = {"key": "sk-old-token"} + with patch( + "litellm.litellm_core_utils.cli_token_utils.load_cli_token", + return_value=token_data, + ): + assert ( + get_stored_api_key(expected_base_url="https://real-proxy.com") is None + ) + class TestLoginCommand: """Test login CLI command""" diff --git a/tests/test_litellm/proxy/common_utils/test_cache_codec.py b/tests/test_litellm/proxy/common_utils/test_cache_codec.py new file mode 100644 index 00000000000..044d4c2d1a7 --- /dev/null +++ b/tests/test_litellm/proxy/common_utils/test_cache_codec.py @@ -0,0 +1,126 @@ +import logging +from typing import Optional +from unittest.mock import patch + +import pytest +from pydantic import BaseModel, ValidationError + +from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec + + +class _SampleModel(BaseModel): + name: str + count: Optional[int] = None + + +class _SampleSubModel(_SampleModel): + pass + + +class TestCacheCodecSerialize: + def test_without_model_type_base_model_dumped_json_safe(self): + m = _SampleModel(name="a", count=1) + out = CacheCodec.serialize(m) + assert out == {"name": "a", "count": 1} + + def test_without_model_type_dict_unchanged(self): + d = {"name": "x"} + assert CacheCodec.serialize(d) is d + + def test_without_model_type_primitive_unchanged(self): + assert CacheCodec.serialize(42) == 42 + + def test_with_model_type_dict_validated_and_dumped(self): + out = CacheCodec.serialize({"name": "b", "count": 2}, model_type=_SampleModel) + assert out == {"name": "b", "count": 2} + + def test_with_model_type_base_model_validated_and_dumped(self): + m = _SampleModel(name="c", count=None) + out = CacheCodec.serialize(m, model_type=_SampleModel) + assert out == {"name": "c"} + + def test_with_model_type_exclude_none_on_dump(self): + out = CacheCodec.serialize({"name": "d"}, model_type=_SampleModel) + assert out == {"name": "d"} + assert "count" not in out + + def test_with_model_type_non_dict_non_model_passthrough(self): + assert CacheCodec.serialize("raw", model_type=_SampleModel) == "raw" + + def test_with_model_type_invalid_dict_raises(self): + with pytest.raises(ValidationError): + CacheCodec.serialize({"count": 1}, model_type=_SampleModel) + + def test_with_model_type_already_correct_instance_skips_revalidation(self): + """Fast-path: value is already model_type — model_validate must NOT be called.""" + m = _SampleModel(name="fast", count=7) + with patch.object(_SampleModel, "model_validate", wraps=_SampleModel.model_validate) as mock_validate: + out = CacheCodec.serialize(m, model_type=_SampleModel) + assert out == {"name": "fast", "count": 7} + mock_validate.assert_not_called() + + def test_with_model_type_subclass_instance_skips_revalidation(self): + """Subclass is isinstance of base → should also take the fast path.""" + sub = _SampleSubModel(name="sub", count=2) + with patch.object(_SampleModel, "model_validate", wraps=_SampleModel.model_validate) as mock_validate: + out = CacheCodec.serialize(sub, model_type=_SampleModel) + assert out == {"name": "sub", "count": 2} + mock_validate.assert_not_called() + + def test_with_model_type_dict_input_goes_through_model_validate(self): + """A dict value (not yet an instance) must still go through model_validate.""" + raw = {"name": "via-dict", "count": 5} + with patch.object( + _SampleModel, "model_validate", wraps=_SampleModel.model_validate + ) as mock_validate: + out = CacheCodec.serialize(raw, model_type=_SampleModel) + assert out == {"name": "via-dict", "count": 5} + mock_validate.assert_called_once() + + def test_with_model_type_incompatible_model_raises_validation_error(self): + """Passing a BaseModel whose fields don't satisfy model_type's required fields raises. + + _IncompatibleModel only has `foo: int`, so when Pydantic v2 extracts its + data and validates it against _SampleModel (which requires `name: str`), + a ValidationError is raised. + """ + + class _IncompatibleModel(BaseModel): + foo: int # missing required 'name' field of _SampleModel + + with pytest.raises(ValidationError): + CacheCodec.serialize(_IncompatibleModel(foo=1), model_type=_SampleModel) + + +class TestCacheCodecDeserialize: + def test_none_returns_none(self): + assert CacheCodec.deserialize(None, _SampleModel) is None + + def test_dict_validates_to_model(self): + m = CacheCodec.deserialize({"name": "e", "count": 3}, _SampleModel) + assert isinstance(m, _SampleModel) + assert m.name == "e" + assert m.count == 3 + + def test_instance_same_type_returned_as_is(self): + original = _SampleModel(name="f") + m = CacheCodec.deserialize(original, _SampleModel) + assert m is original + + def test_subclass_instance_accepted(self): + sub = _SampleSubModel(name="g") + m = CacheCodec.deserialize(sub, _SampleModel) + assert m is sub + + def test_wrong_type_returns_none(self): + assert CacheCodec.deserialize("not-a-dict", _SampleModel) is None + + def test_invalid_dict_returns_none_and_logs_warning(self, caplog): + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + out = CacheCodec.deserialize({"count": 1}, _SampleModel) + assert out is None + assert any( + "CacheCodec.deserialize" in r.message and "_SampleModel" in r.message + for r in caplog.records + if r.levelno >= logging.WARNING + ), f"Expected deserialize validation warning. Records: {[r.message for r in caplog.records]}" diff --git a/tests/test_litellm/proxy/common_utils/test_user_api_key_cache.py b/tests/test_litellm/proxy/common_utils/test_user_api_key_cache.py new file mode 100644 index 00000000000..8667348d223 --- /dev/null +++ b/tests/test_litellm/proxy/common_utils/test_user_api_key_cache.py @@ -0,0 +1,219 @@ +import json +from typing import Any + +import pytest + +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.caching.redis_cache import RedisCache +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache +from litellm.proxy.proxy_server import UserAPIKeyCacheTTLEnum + + +class CapturingInMemoryCache(InMemoryCache): + """Records ``ttl`` passed into ``set_cache`` (what DualCache injects).""" + + def __init__(self) -> None: + super().__init__() + self.last_ttl: Any = None + + def set_cache(self, key, value, **kwargs): # type: ignore[override] + self.last_ttl = kwargs.get("ttl") + super().set_cache(key, value, **kwargs) + + +class FakeRedisCache(RedisCache): + """ + In-memory fake that enforces the UserApiKeyCache Redis payload contract. + + For user_api_key_cache entries we expect Redis to store a JSON object (dict) + produced by `CacheCodec.serialize(..., model_type=...)`. + + This fake: + - raises TypeError if the value is not a dict + - raises TypeError if the dict is not JSON-serializable + + Records the ``ttl`` kwarg DualCache forwards on each Redis write for tests. + """ + + def __init__(self): # noqa: super().__init__ skipped intentionally + self._store: dict[str, str] = {} + self.last_ttl: Any = None + + def set_cache(self, key: str, value: Any, **kwargs): # type: ignore[override] + if not isinstance(value, dict): + raise TypeError("FakeRedisCache only accepts dict payloads") + self.last_ttl = kwargs.get("ttl") + self._store[key] = json.dumps(value) + return True + + def get_cache(self, key: str, **kwargs): # type: ignore[override] + raw = self._store.get(key) + if raw is None: + return None + return json.loads(raw) + + async def async_set_cache(self, key: str, value: Any, **kwargs): # type: ignore[override] + if not isinstance(value, dict): + raise TypeError("FakeRedisCache only accepts dict payloads") + self.last_ttl = kwargs.get("ttl") + self._store[key] = json.dumps(value) + return True + + async def async_get_cache(self, key: str, **kwargs): # type: ignore[override] + raw = self._store.get(key) + if raw is None: + return None + return json.loads(raw) + + def delete_cache(self, key: str): # type: ignore[override] + self._store.pop(key, None) + + async def async_delete_cache(self, key: str): # type: ignore[override] + self._store.pop(key, None) + + +def _make_key_obj(token: str = "tok") -> UserAPIKeyAuth: + # Minimal object (UserAPIKeyAuth inherits token from base view). + return UserAPIKeyAuth(token=token) + + +class TestUserApiKeyCache: + @pytest.mark.asyncio + async def test_async_set_in_memory_gets_enum_default_when_user_api_key_cache_ttl_omitted( + self, + ): + """ + If ``general_settings.user_api_key_cache_ttl`` is absent, the proxy never + calls ``update_cache_ttl``; ``user_api_key_cache`` keeps + ``default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl``. + DualCache must forward that as the in-memory ``ttl`` kwarg on each set. + """ + mem = CapturingInMemoryCache() + cache = UserApiKeyCache( + in_memory_cache=mem, + redis_cache=FakeRedisCache(), + default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value, + ) + await cache.async_set_cache( + "k", + _make_key_obj("t"), + model_type=UserAPIKeyAuth, + ) + expected = UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value + assert mem.last_ttl == expected + + def test_sync_set_in_memory_gets_enum_default_when_user_api_key_cache_ttl_omitted( + self, + ): + mem = CapturingInMemoryCache() + cache = UserApiKeyCache( + in_memory_cache=mem, + redis_cache=FakeRedisCache(), + default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value, + ) + cache.set_cache("sk", _make_key_obj("s"), model_type=UserAPIKeyAuth) + assert mem.last_ttl == UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value + + @pytest.mark.asyncio + async def test_async_set_forwards_default_in_memory_ttl_to_redis_layer(self): + """ + DualCache injects missing ``ttl`` from ``default_in_memory_ttl`` into kwargs + before calling ``redis_cache.async_set_cache`` — Redis should receive the same + TTL as memory (matches proxy defaults: enum 60s). + """ + fake = FakeRedisCache() + cache = UserApiKeyCache( + redis_cache=fake, + default_in_memory_ttl=60, + ) + + await cache.async_set_cache( + key="ttl-key", + value=_make_key_obj("ttl-tok"), + model_type=UserAPIKeyAuth, + ) + + assert fake.last_ttl == 60 + + @pytest.mark.asyncio + async def test_async_set_explicit_ttl_override_reaches_redis(self): + fake = FakeRedisCache() + cache = UserApiKeyCache( + redis_cache=fake, + default_in_memory_ttl=60, + ) + + await cache.async_set_cache( + key="k", + value=_make_key_obj("x"), + model_type=UserAPIKeyAuth, + ttl=900, + ) + + assert fake.last_ttl == 900 + + def test_sync_set_forwards_default_in_memory_ttl_to_redis_layer(self): + fake = FakeRedisCache() + cache = UserApiKeyCache( + redis_cache=fake, + default_in_memory_ttl=45, + ) + cache.set_cache( + "sk", + _make_key_obj("sync"), + model_type=UserAPIKeyAuth, + ) + assert fake.last_ttl == 45 + + @pytest.mark.asyncio + async def test_async_set_typed_stores_serialized_payload_in_memory_and_redis(self): + cache = UserApiKeyCache(redis_cache=FakeRedisCache()) + obj = _make_key_obj("abc") + + await cache.async_set_cache("k", obj, model_type=UserAPIKeyAuth) + + # In-memory hit should still be raw dict (not BaseModel) because wrapper + # stores the serialized payload into both layers. + raw = await cache.in_memory_cache.async_get_cache("k") # type: ignore[union-attr] + assert isinstance(raw, dict) + assert raw["token"] == "abc" + + # Redis should also hold the same serialized dict + redis_raw = await cache.redis_cache.async_get_cache("k") # type: ignore[union-attr] + assert redis_raw == raw + + @pytest.mark.asyncio + async def test_async_get_typed_returns_model_on_valid_hit(self): + cache = UserApiKeyCache(redis_cache=FakeRedisCache()) + await cache.async_set_cache("k", {"token": "abc"}, model_type=UserAPIKeyAuth) + + value = await cache.async_get_cache("k", model_type=UserAPIKeyAuth) + assert value is not None + assert isinstance(value, UserAPIKeyAuth) + assert value.token == "abc" + + @pytest.mark.asyncio + async def test_async_get_typed_returns_none_on_validation_failure_after_hit(self): + cache = UserApiKeyCache(redis_cache=FakeRedisCache()) + + # Bypass UserApiKeyCache.serialize: CacheCodec rejects non-dict cached values + # for dict-based models (deserialize returns None). + await cache.in_memory_cache.async_set_cache( + key="k", value="invalid-payload-not-a-dict" + ) + + value = await cache.async_get_cache("k", model_type=UserAPIKeyAuth) + assert value is None + + def test_fake_redis_cache_rejects_non_json_serializable_values(self): + fake = FakeRedisCache() + + class NotSerializable: + pass + + with pytest.raises(TypeError): + fake.set_cache("k", NotSerializable()) + + with pytest.raises(TypeError): + fake.set_cache("k2", {"ok": NotSerializable()}) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_permission.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_permission.py index 55d92e91416..716b4470d25 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_permission.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_permission.py @@ -220,6 +220,27 @@ class TestToolPermissionGuardrail: assert tool_calls[0].id == "call_123" assert tool_calls[0].function.name == "Read" + def test_extract_tool_calls_legacy_function_call_format(self): + response = ModelResponse( + choices=[ + Choices( + message={ + "function_call": { + "name": "Read", + "arguments": '{"file_path": "/test/file.txt"}', + }, + } + ) + ] + ) + + tool_calls = self.guardrail._extract_tool_calls_from_response(response) + assert len(tool_calls) == 1 + assert isinstance(tool_calls[0], ChatCompletionMessageToolCall) + assert tool_calls[0].id == "legacy_function_call_0" + assert tool_calls[0].function.name == "Read" + assert tool_calls[0].function.arguments == '{"file_path": "/test/file.txt"}' + def test_extract_tool_calls_empty_response(self): response = ModelResponse(choices=[]) tool_calls = self.guardrail._extract_tool_calls_from_response(response) @@ -271,6 +292,31 @@ class TestToolPermissionGuardrail: data=data, user_api_key_dict=user_api_key_dict, response=response ) + @pytest.mark.asyncio + async def test_async_post_call_success_hook_with_denied_legacy_function_call_raises( + self, + ): + response = ModelResponse( + choices=[ + Choices( + message={ + "function_call": { + "name": "Read", + "arguments": "{}", + }, + } + ) + ] + ) + user_api_key_dict = UserAPIKeyAuth() + data = {"guardrails": ["test-tool-permission"]} + + with patch.object(self.guardrail, "should_run_guardrail", return_value=True): + with pytest.raises(GuardrailRaisedException): + await self.guardrail.async_post_call_success_hook( + data=data, user_api_key_dict=user_api_key_dict, response=response + ) + @pytest.mark.asyncio async def test_async_post_call_success_hook_param_patterns_allow(self): guardrail = ToolPermissionGuardrail( @@ -379,7 +425,9 @@ class TestToolPermissionGuardrail: assert "berri" in choice.message.content @pytest.mark.asyncio - async def test_async_post_call_success_hook_missing_arguments_default_allows(self): + async def test_async_post_call_success_hook_missing_arguments_blocks_param_rule( + self, + ): guardrail = ToolPermissionGuardrail( guardrail_name="mail-guardrail", rules=[ @@ -405,9 +453,52 @@ class TestToolPermissionGuardrail: data = {"guardrails": ["mail-guardrail"]} with patch.object(guardrail, "should_run_guardrail", return_value=True): - await guardrail.async_post_call_success_hook( - data=data, user_api_key_dict=user_api_key_dict, response=response - ) + with pytest.raises(GuardrailRaisedException): + await guardrail.async_post_call_success_hook( + data=data, user_api_key_dict=user_api_key_dict, response=response + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "arguments", + [ + "{not-json", + '["owner@berri.ai"]', + ], + ) + async def test_async_post_call_success_hook_malformed_arguments_blocks_param_rule( + self, arguments + ): + guardrail = ToolPermissionGuardrail( + guardrail_name="mail-guardrail", + rules=[ + { + "id": "deny_gmail", + "tool_name": r"^mail_mcp-send_email$", + "decision": "deny", + "allowed_param_patterns": {"to[]": r"^.+@gmail\.com$"}, + } + ], + default_action="allow", + on_disallowed_action="block", + ) + + tool_call = { + "function": { + "name": "mail_mcp-send_email", + "arguments": arguments, + }, + "type": "function", + } + response = ModelResponse(choices=[Choices(message={"tool_calls": [tool_call]})]) + user_api_key_dict = UserAPIKeyAuth() + data = {"guardrails": ["mail-guardrail"]} + + with patch.object(guardrail, "should_run_guardrail", return_value=True): + with pytest.raises(GuardrailRaisedException): + await guardrail.async_post_call_success_hook( + data=data, user_api_key_dict=user_api_key_dict, response=response + ) @pytest.mark.asyncio async def test_async_pre_call_hook_block_mode(self): @@ -430,6 +521,65 @@ class TestToolPermissionGuardrail: ) assert excinfo.value.status_code == 400 + @pytest.mark.asyncio + async def test_async_pre_call_hook_blocks_legacy_functions(self): + data = { + "functions": [ + {"name": "Bash", "description": "allowed"}, + {"name": "Read", "description": "denied"}, + ] + } + user_api_key_dict = UserAPIKeyAuth() + cache = DualCache(default_in_memory_ttl=1) + + with patch.object(self.guardrail, "should_run_guardrail", return_value=True): + with pytest.raises(HTTPException) as excinfo: + await self.guardrail.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="completion", + ) + assert excinfo.value.status_code == 400 + + @pytest.mark.asyncio + async def test_async_pre_call_hook_blocks_named_legacy_function_call(self): + data = { + "functions": [{"name": "Bash"}], + "function_call": {"name": "Read"}, + } + user_api_key_dict = UserAPIKeyAuth() + cache = DualCache(default_in_memory_ttl=1) + + with patch.object(self.guardrail, "should_run_guardrail", return_value=True): + with pytest.raises(HTTPException) as excinfo: + await self.guardrail.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="completion", + ) + assert excinfo.value.status_code == 400 + + @pytest.mark.asyncio + async def test_async_pre_call_hook_blocks_named_tool_choice(self): + data = { + "tools": [{"type": "function", "function": {"name": "Bash"}}], + "tool_choice": {"type": "function", "function": {"name": "Read"}}, + } + user_api_key_dict = UserAPIKeyAuth() + cache = DualCache(default_in_memory_ttl=1) + + with patch.object(self.guardrail, "should_run_guardrail", return_value=True): + with pytest.raises(HTTPException) as excinfo: + await self.guardrail.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="completion", + ) + assert excinfo.value.status_code == 400 + @pytest.mark.asyncio async def test_async_pre_call_hook_uses_custom_template(self): guardrail = ToolPermissionGuardrail( @@ -491,6 +641,41 @@ class TestToolPermissionGuardrail: assert "Bash" in tool_names assert "Read" not in tool_names + @pytest.mark.asyncio + async def test_async_pre_call_hook_rewrite_mode_filters_legacy_functions(self): + guardrail = ToolPermissionGuardrail( + guardrail_name="test-tool-permission", + rules=self.test_rules, + default_action="deny", + on_disallowed_action="rewrite", + ) + data = { + "functions": [ + {"name": "Bash", "description": "allowed"}, + {"name": "Read", "description": "denied"}, + ], + "function_call": {"name": "Read"}, + "tools": [ + {"type": "function", "function": {"name": "Bash"}}, + ], + "tool_choice": {"type": "function", "function": {"name": "Read"}}, + } + user_api_key_dict = UserAPIKeyAuth() + cache = DualCache(default_in_memory_ttl=1) + + with patch.object(guardrail, "should_run_guardrail", return_value=True): + new_data = await guardrail.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="completion", + ) + + assert isinstance(new_data, dict) + assert [function["name"] for function in new_data["functions"]] == ["Bash"] + assert new_data["function_call"] == "none" + assert new_data["tool_choice"] == "none" + def test_modify_response_with_permission_errors(self): # Setup a response with one tool_call tool_call = ChatCompletionMessageToolCall( @@ -522,6 +707,40 @@ class TestToolPermissionGuardrail: assert isinstance(choice.message.content, str) assert "Permission denied" in choice.message.content + def test_modify_response_with_permission_errors_filters_legacy_function_call(self): + response = ModelResponse( + choices=[ + Choices( + message={ + "function_call": { + "name": "Read", + "arguments": "{}", + }, + "content": "", + } + ) + ] + ) + tool_call = self.guardrail._extract_tool_calls_from_response(response)[0] + denied_tools = [ + ( + tool_call, + PermissionError( + tool_name="Read", + rule_id="deny_read", + message="Tool 'Read' denied by rule 'deny_read'", + ), + ) + ] + + self.guardrail._modify_response_with_permission_errors(response, denied_tools) + + choice = response.choices[0] + assert isinstance(choice, Choices) + assert choice.message.function_call is None + assert isinstance(choice.message.content, str) + assert "Permission denied" in choice.message.content + class TestToolPermissionGuardrailIntegration: """Integration tests for Tool Permission Guardrail""" diff --git a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py index 59b6e24f430..e10258c0829 100644 --- a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py +++ b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py @@ -758,6 +758,60 @@ class TestDeferredStreamingClosure: apply_guardrail_called is False ), "apply_guardrail guardrails must be SKIPPED in deferred path" + @pytest.mark.asyncio + async def test_streaming_iterator_hook_skipped_in_deferred_path(self): + """regression test: guardrails that define async_post_call_streaming_iterator_hook must be SKIPPED in _run_deferred_stream_guardrails. + The iterator hook already scanned the assembled response in the streaming + pipeline""" + success_hook_called = False + + class IteratorHookGuardrail(CustomGuardrail): + def __init__(self): + super().__init__( + guardrail_name="iterator-hook", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + ) + + async def async_post_call_streaming_iterator_hook( + self, user_api_key_dict, response, request_data + ): + async for chunk in response: + yield chunk + + async def async_post_call_success_hook( + self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any + ) -> Any: + nonlocal success_hook_called + success_hook_called = True + return response + + mock_logging_obj = MagicMock() + mock_logging_obj.model_call_details = {"metadata": {}} + + async def track_async_success(*args, **kwargs): + pass + + mock_logging_obj.async_success_handler = track_async_success + + guardrail = IteratorHookGuardrail() + + with patch("litellm.callbacks", [guardrail]): + await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails( + captured_data={"model": "gpt-4", "metadata": {}}, + captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"), + captured_logging_obj=mock_logging_obj, + assembled_response=MagicMock(), + cache_hit=False, + ) + + await asyncio.sleep(0) + + assert success_hook_called is False, ( + "Guardrails that implement async_post_call_streaming_iterator_hook " + "must be SKIPPED in deferred path — the iterator hook already ran" + ) + @pytest.mark.asyncio async def test_hooks_receive_merged_guardrail_data(self): """Hooks must receive guardrail_data (the merged dict from diff --git a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py index ba260142351..353e67c9f77 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -778,3 +778,797 @@ def test_get_callback_identifier_custom_logger_registry_and_fallback(): result = get_callback_identifier(my_callback_function) # Should fall back to callback_name() which returns __name__ assert result == "my_callback_function" + + +# --------------------------------------------------------------------------- +# /health response shape: model-access scoping and display-field allowlist +# --------------------------------------------------------------------------- +# These tests pin the contract that the /health response (a) only includes +# deployments the calling key is allowed to see, and (b) does not return +# provider routing fields like api_base / api_version. They guard against +# regressions that would widen the response shape. + + +@pytest.mark.asyncio +async def test_health_endpoint_filters_model_list_by_user_access(): + """ + health_endpoint() should restrict _llm_model_list to deployments whose + model_name appears in user_api_key_dict.models before running the health + check. A key scoped to ["model-a"] should only see model-a in the result, + not other deployments configured on the proxy. + """ + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.health_endpoints._health_endpoints import health_endpoint + + full_model_list = [ + { + "model_name": "model-a", + "litellm_params": { + "model": "openai/gpt-4o", + "api_base": "https://example-a.test", + }, + "model_info": {"id": "id-a"}, + }, + { + "model_name": "model-b", + "litellm_params": { + "model": "openai/gpt-4o", + "api_base": "https://example-b.test", + "api_version": "2024-10-21", + }, + "model_info": {"id": "id-b"}, + }, + ] + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-test-key", + models=["model-a"], + ) + + captured: dict = {} + + async def fake_perform(**kwargs): + captured["model_list"] = kwargs["model_list"] + return { + "healthy_endpoints": [], + "unhealthy_endpoints": [], + "healthy_count": 0, + "unhealthy_count": 0, + } + + with ( + patch("litellm.proxy.proxy_server.llm_model_list", full_model_list), + patch("litellm.proxy.proxy_server.llm_router", None), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch("litellm.proxy.proxy_server.use_background_health_checks", False), + patch("litellm.proxy.proxy_server.user_model", None), + patch("litellm.proxy.proxy_server.health_check_results", {}), + patch("litellm.proxy.proxy_server.health_check_details", True), + patch("litellm.proxy.proxy_server.health_check_concurrency", 1), + patch( + "litellm.proxy.health_endpoints._health_endpoints._perform_health_check_and_save", + side_effect=fake_perform, + ), + ): + from fastapi import Response + + await health_endpoint(response=Response(), user_api_key_dict=user_api_key_dict) + + assert ( + "model_list" in captured + ), "health_endpoint did not call _perform_health_check_and_save" + returned_names = {m["model_name"] for m in captured["model_list"]} + assert returned_names == { + "model-a" + }, f"health_endpoint did not scope model_list to caller access: {returned_names}" + + +@pytest.mark.asyncio +async def test_health_endpoint_filters_background_cache_by_user_access(): + """ + When background_health_checks is enabled, health_endpoint() should also + scope the cached result to the caller's allowed models rather than + returning the cache verbatim. + """ + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.health_endpoints._health_endpoints import health_endpoint + + full_model_list = [ + { + "model_name": "model-a", + "litellm_params": { + "model": "openai/gpt-4o", + "api_base": "https://example-a.test", + }, + "model_info": {"id": "id-a"}, + }, + { + "model_name": "model-b", + "litellm_params": { + "model": "openai/gpt-4o", + "api_base": "https://example-b.test", + }, + "model_info": {"id": "id-b"}, + }, + ] + + cached_results = { + "healthy_endpoints": [ + { + "model": "openai/gpt-4o", + "model_id": "id-a", + "api_base": "https://example-a.test", + }, + { + "model": "openai/gpt-4o", + "model_id": "id-b", + "api_base": "https://example-b.test", + }, + ], + "unhealthy_endpoints": [], + "healthy_count": 2, + "unhealthy_count": 0, + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-test-key", + models=["model-a"], + ) + + with ( + patch("litellm.proxy.proxy_server.llm_model_list", full_model_list), + patch("litellm.proxy.proxy_server.llm_router", None), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch("litellm.proxy.proxy_server.use_background_health_checks", True), + patch("litellm.proxy.proxy_server.user_model", None), + patch("litellm.proxy.proxy_server.health_check_results", cached_results), + patch("litellm.proxy.proxy_server.health_check_details", True), + patch("litellm.proxy.proxy_server.health_check_concurrency", 1), + ): + from fastapi import Response + + # Pass model=None, model_id=None explicitly: direct calls to the + # handler skip FastAPI's Query() resolution, so unspecified params + # would otherwise carry the Query() sentinel (which is truthy). + result = await health_endpoint( + response=Response(), + user_api_key_dict=user_api_key_dict, + model=None, + model_id=None, + ) + + # Sanity: the source cache had two entries before scoping; the scoping + # step is what reduces it to one. (This guards against the test passing + # vacuously when the cache filter drops everything because cached + # entries lack the model_id key — both entries carry model_id above.) + assert len(cached_results["healthy_endpoints"]) == 2 + assert all( + ep.get("model_id") for ep in cached_results["healthy_endpoints"] + ), "test fixture invariant: every cached entry must carry a model_id" + + # The non-admin caller must not see api_base on the returned cache entries. + returned = result.get("healthy_endpoints", []) + assert ( + len(returned) == 1 + ), f"expected exactly one cached entry after scoping, got {len(returned)}" + assert returned[0]["model_id"] == "id-a" + assert "api_base" not in returned[0] + assert result["healthy_count"] == 1 + assert result["unhealthy_count"] == 0 + + +@pytest.mark.asyncio +async def test_health_endpoint_admin_sees_routing_fields_non_admin_does_not(): + """ + A proxy admin should still see ``api_base`` and ``api_version`` in the + /health response so they can tell which Vertex region / Azure resource + + API version is healthy. A non-admin caller must not — both fields + should be stripped, and the response should carry a notice header so + non-admin clients can detect the change programmatically. + """ + from fastapi import Response + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.health_endpoints._health_endpoints import health_endpoint + + full_model_list = [ + { + "model_name": "model-a", + "litellm_params": { + "model": "openai/gpt-4o", + "api_base": "https://example-a.test", + }, + "model_info": {"id": "id-a"}, + }, + ] + cached_results = { + "healthy_endpoints": [ + { + "model": "openai/gpt-4o", + "model_id": "id-a", + "api_base": "https://us-central1-aiplatform.googleapis.com/v1/projects/p", + "api_version": "2024-10-21", + }, + ], + "unhealthy_endpoints": [], + "healthy_count": 1, + "unhealthy_count": 0, + } + + admin_key = UserAPIKeyAuth( + api_key="hashed-admin-key", + models=["model-a"], + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + non_admin_key = UserAPIKeyAuth( + api_key="hashed-user-key", + models=["model-a"], + ) + + common_patches = [ + patch("litellm.proxy.proxy_server.llm_model_list", full_model_list), + patch("litellm.proxy.proxy_server.llm_router", None), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch("litellm.proxy.proxy_server.use_background_health_checks", True), + patch("litellm.proxy.proxy_server.user_model", None), + patch("litellm.proxy.proxy_server.health_check_results", cached_results), + patch("litellm.proxy.proxy_server.health_check_details", True), + patch("litellm.proxy.proxy_server.health_check_concurrency", 1), + ] + + for p in common_patches: + p.start() + try: + admin_response = Response() + non_admin_response = Response() + admin_result = await health_endpoint( + response=admin_response, + user_api_key_dict=admin_key, + model=None, + model_id=None, + ) + non_admin_result = await health_endpoint( + response=non_admin_response, + user_api_key_dict=non_admin_key, + model=None, + model_id=None, + ) + finally: + for p in common_patches: + p.stop() + + admin_eps = admin_result.get("healthy_endpoints", []) + non_admin_eps = non_admin_result.get("healthy_endpoints", []) + + assert len(admin_eps) == 1 + assert ( + admin_eps[0]["api_base"] + == "https://us-central1-aiplatform.googleapis.com/v1/projects/p" + ), "admin must see the full api_base so they can identify the region" + assert ( + admin_eps[0]["api_version"] == "2024-10-21" + ), "admin must see api_version so they can distinguish provider deployments" + + assert len(non_admin_eps) == 1 + assert "api_base" not in non_admin_eps[0] + assert "api_version" not in non_admin_eps[0] + + # Non-admin response must advertise that api_base/api_version were + # withheld so clients that previously parsed them can detect the change. + assert ( + non_admin_response.headers.get("Litellm-Health-Field-Notice") + == "api_base and api_version are admin-only on this endpoint" + ) + assert "Litellm-Health-Field-Notice" not in admin_response.headers + + # Stripping must produce a copy — the shared cache must still carry the + # routing fields so the next admin caller can read them. + cached_first = cached_results["healthy_endpoints"][0] + assert ( + cached_first["api_base"] + == "https://us-central1-aiplatform.googleapis.com/v1/projects/p" + ) + assert cached_first["api_version"] == "2024-10-21" + + +@pytest.mark.asyncio +async def test_health_endpoint_warns_when_scoped_models_lack_model_id(): + """ + When a scoped key's accessible models exist on the proxy but none of the + matching deployments expose a ``model_info.id``, the cache filter drops + everything. The response should include a structured ``warnings`` field + so the caller can distinguish "no deployments configured" from + "deployments excluded due to missing model_info.id". + """ + from fastapi import Response + + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.health_endpoints._health_endpoints import health_endpoint + + full_model_list = [ + { + "model_name": "model-a", + "litellm_params": { + "model": "openai/gpt-4o", + "api_base": "https://example-a.test", + }, + # Intentionally no model_info.id — this is the misconfiguration + # the warnings field is meant to flag. + "model_info": {}, + }, + ] + cached_results = { + "healthy_endpoints": [ + { + "model": "openai/gpt-4o", + "model_id": "id-a", + "api_base": "https://example-a.test", + }, + ], + "unhealthy_endpoints": [], + "healthy_count": 1, + "unhealthy_count": 0, + } + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-user-key", + models=["model-a"], + ) + + with ( + patch("litellm.proxy.proxy_server.llm_model_list", full_model_list), + patch("litellm.proxy.proxy_server.llm_router", None), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch("litellm.proxy.proxy_server.use_background_health_checks", True), + patch("litellm.proxy.proxy_server.user_model", None), + patch("litellm.proxy.proxy_server.health_check_results", cached_results), + patch("litellm.proxy.proxy_server.health_check_details", True), + patch("litellm.proxy.proxy_server.health_check_concurrency", 1), + ): + result = await health_endpoint( + response=Response(), + user_api_key_dict=user_api_key_dict, + model=None, + model_id=None, + ) + + assert result["healthy_count"] == 0 + assert result["unhealthy_count"] == 0 + assert "warnings" in result, ( + "empty cache result must surface a warnings field so the caller " + "can distinguish 'no deployments' from 'deployments excluded'" + ) + assert any("model_info.id" in w for w in result["warnings"]) + + +@pytest.mark.asyncio +async def test_health_endpoint_blocks_cross_scope_model_id_under_background_cache(): + """ + A non-admin scoped to model-a must not be able to read model-b's cached + health entry by guessing its model_id. Before the fix, + _resolve_targeted_model_ids returned {model_id} unconditionally, so the + cache filter was driven by an unvalidated ID and the global cache + leaked id-b's entry to the caller. + """ + from fastapi import Response + + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.health_endpoints._health_endpoints import health_endpoint + + full_model_list = [ + { + "model_name": "model-a", + "litellm_params": {"model": "openai/gpt-4o"}, + "model_info": {"id": "id-a"}, + }, + { + "model_name": "model-b", # caller has no access + "litellm_params": {"model": "openai/gpt-4o"}, + "model_info": {"id": "id-b"}, + }, + ] + + cached_results = { + "healthy_endpoints": [ + { + "model": "openai/gpt-4o", + "model_id": "id-b", + "api_base": "https://leaky-internal.test", + }, + ], + "unhealthy_endpoints": [], + "healthy_count": 1, + "unhealthy_count": 0, + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-scoped", + models=["model-a"], + ) + + response = Response() + with ( + patch("litellm.proxy.proxy_server.llm_model_list", full_model_list), + # llm_router None here means the model_id 404 lookup short-circuits; + # we patch _llm_model_list directly instead to drive the cache path. + patch("litellm.proxy.proxy_server.llm_router", None), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch("litellm.proxy.proxy_server.use_background_health_checks", True), + patch("litellm.proxy.proxy_server.user_model", None), + patch("litellm.proxy.proxy_server.health_check_results", cached_results), + patch("litellm.proxy.proxy_server.health_check_details", True), + patch("litellm.proxy.proxy_server.health_check_concurrency", 1), + ): + # Calling with model="model-b" rather than model_id="id-b" because + # the model_id branch raises 404 when llm_router is None. The bug + # being verified is the same: targeted resolver must drop entries + # not in the caller's scoped model_list. With the fix, the result + # has no leaked endpoints and the targeted-503 path fires. + result = await health_endpoint( + response=response, + user_api_key_dict=user_api_key_dict, + model="model-b", + model_id=None, + ) + + leaked_ids = {ep.get("model_id") for ep in result.get("healthy_endpoints", [])} + leaked_ids |= {ep.get("model_id") for ep in result.get("unhealthy_endpoints", [])} + assert ( + "id-b" not in leaked_ids + ), "background cache leaked an out-of-scope deployment to a scoped caller" + assert result["healthy_count"] == 0 + assert response.status_code == 503 + + +@pytest.mark.asyncio +async def test_health_endpoint_503_for_targeted_unhealthy_model_under_background_cache_admin(): + """ + With background_health_checks enabled, an admin calling /health?model=foo + must get 503 when foo specifically has zero healthy endpoints — even if + other unrelated models in the cache are healthy. Without the cache-path + filter, the global healthy_count would mask the targeted failure. + """ + from fastapi import Response + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.health_endpoints._health_endpoints import health_endpoint + + full_model_list = [ + { + "model_name": "model-a", # the unhealthy target + "litellm_params": {"model": "openai/gpt-4o"}, + "model_info": {"id": "id-a"}, + }, + { + "model_name": "model-b", # an unrelated healthy model + "litellm_params": {"model": "openai/gpt-4o"}, + "model_info": {"id": "id-b"}, + }, + ] + + cached_results = { + "healthy_endpoints": [ + {"model": "openai/gpt-4o", "model_id": "id-b"}, + ], + "unhealthy_endpoints": [ + {"model": "openai/gpt-4o", "model_id": "id-a", "error": "boom"}, + ], + "healthy_count": 1, + "unhealthy_count": 1, + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-admin", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + + response = Response() + with ( + patch("litellm.proxy.proxy_server.llm_model_list", full_model_list), + patch("litellm.proxy.proxy_server.llm_router", None), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch("litellm.proxy.proxy_server.use_background_health_checks", True), + patch("litellm.proxy.proxy_server.user_model", None), + patch("litellm.proxy.proxy_server.health_check_results", cached_results), + patch("litellm.proxy.proxy_server.health_check_details", True), + patch("litellm.proxy.proxy_server.health_check_concurrency", 1), + ): + result = await health_endpoint( + response=response, + user_api_key_dict=user_api_key_dict, + model="model-a", + model_id=None, + ) + + assert response.status_code == 503 + # Body must be scoped to the targeted model — not the global cache. + assert result["healthy_count"] == 0 + assert result["unhealthy_count"] == 1 + returned_ids = {ep["model_id"] for ep in result.get("unhealthy_endpoints", [])} + assert returned_ids == {"id-a"} + + +@pytest.mark.asyncio +async def test_health_endpoint_returns_503_when_requested_model_has_no_healthy_endpoints(): + """ + /health?model=foo must return 503 when the targeted model resolves but + has zero healthy endpoints. Body shape stays the same so existing + parsers still work; only the HTTP status changes. + """ + from fastapi import Response + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.health_endpoints._health_endpoints import health_endpoint + + full_model_list = [ + { + "model_name": "model-a", + "litellm_params": { + "model": "openai/gpt-4o", + "api_base": "https://example-a.test", + }, + "model_info": {"id": "id-a"}, + }, + ] + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-test-key", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + + async def fake_perform(**kwargs): + return { + "healthy_endpoints": [], + "unhealthy_endpoints": [ + { + "model": "openai/gpt-4o", + "model_id": "id-a", + "error": "boom", + } + ], + "healthy_count": 0, + "unhealthy_count": 1, + } + + response = Response() + with ( + patch("litellm.proxy.proxy_server.llm_model_list", full_model_list), + patch("litellm.proxy.proxy_server.llm_router", None), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch("litellm.proxy.proxy_server.use_background_health_checks", False), + patch("litellm.proxy.proxy_server.user_model", None), + patch("litellm.proxy.proxy_server.health_check_results", {}), + patch("litellm.proxy.proxy_server.health_check_details", True), + patch("litellm.proxy.proxy_server.health_check_concurrency", 1), + patch( + "litellm.proxy.health_endpoints._health_endpoints._perform_health_check_and_save", + side_effect=fake_perform, + ), + ): + result = await health_endpoint( + response=response, + user_api_key_dict=user_api_key_dict, + model="model-a", + ) + + assert response.status_code == 503 + assert result["healthy_count"] == 0 + assert result["unhealthy_count"] == 1 + + +@pytest.mark.asyncio +async def test_health_endpoint_returns_200_when_requested_model_has_healthy_endpoints(): + """ + /health?model=foo with a healthy endpoint must keep returning the + default 200. Verifies the 503 path doesn't fire when healthy_count > 0. + """ + from fastapi import Response + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.health_endpoints._health_endpoints import health_endpoint + + full_model_list = [ + { + "model_name": "model-a", + "litellm_params": {"model": "openai/gpt-4o"}, + "model_info": {"id": "id-a"}, + }, + ] + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-test-key", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + + async def fake_perform(**kwargs): + return { + "healthy_endpoints": [{"model": "openai/gpt-4o", "model_id": "id-a"}], + "unhealthy_endpoints": [], + "healthy_count": 1, + "unhealthy_count": 0, + } + + response = Response() + # Default Response() exposes status_code as None; the endpoint should + # leave it alone for the healthy path. + with ( + patch("litellm.proxy.proxy_server.llm_model_list", full_model_list), + patch("litellm.proxy.proxy_server.llm_router", None), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch("litellm.proxy.proxy_server.use_background_health_checks", False), + patch("litellm.proxy.proxy_server.user_model", None), + patch("litellm.proxy.proxy_server.health_check_results", {}), + patch("litellm.proxy.proxy_server.health_check_details", True), + patch("litellm.proxy.proxy_server.health_check_concurrency", 1), + patch( + "litellm.proxy.health_endpoints._health_endpoints._perform_health_check_and_save", + side_effect=fake_perform, + ), + ): + await health_endpoint( + response=response, + user_api_key_dict=user_api_key_dict, + model="model-a", + ) + + assert response.status_code == 200 + + +@pytest.mark.asyncio +async def test_health_endpoint_no_model_param_returns_200_even_when_zero_healthy(): + """ + The non-targeted /health (no model / model_id query) preserves the + legacy 200 behavior even when healthy_count == 0. Existing K8s probes + and dashboards depend on this; only the targeted call became 5xx-aware. + """ + from fastapi import Response + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.health_endpoints._health_endpoints import health_endpoint + + full_model_list = [ + { + "model_name": "model-a", + "litellm_params": {"model": "openai/gpt-4o"}, + "model_info": {"id": "id-a"}, + }, + ] + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-test-key", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + + async def fake_perform(**kwargs): + return { + "healthy_endpoints": [], + "unhealthy_endpoints": [ + {"model": "openai/gpt-4o", "model_id": "id-a", "error": "boom"} + ], + "healthy_count": 0, + "unhealthy_count": 1, + } + + response = Response() + with ( + patch("litellm.proxy.proxy_server.llm_model_list", full_model_list), + patch("litellm.proxy.proxy_server.llm_router", None), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch("litellm.proxy.proxy_server.use_background_health_checks", False), + patch("litellm.proxy.proxy_server.user_model", None), + patch("litellm.proxy.proxy_server.health_check_results", {}), + patch("litellm.proxy.proxy_server.health_check_details", True), + patch("litellm.proxy.proxy_server.health_check_concurrency", 1), + patch( + "litellm.proxy.health_endpoints._health_endpoints._perform_health_check_and_save", + side_effect=fake_perform, + ), + ): + # Pass model=None, model_id=None explicitly: when invoked through + # FastAPI, the Query(None) defaults resolve to None, but direct + # function calls in unit tests receive Query() sentinel objects + # (which are truthy). The explicit None mirrors production routing. + await health_endpoint( + response=response, + user_api_key_dict=user_api_key_dict, + model=None, + model_id=None, + ) + + assert response.status_code == 200 + + +@pytest.mark.asyncio +async def test_health_readiness_returns_503_when_db_disconnected(): + """ + When a Prisma client is configured but its health_check fails, the + readiness probe should mark the worker as unhealthy via the HTTP + status — not just a body field — so K8s removes the pod from the + Service endpoints. + """ + from fastapi import Response + + from litellm.proxy.health_endpoints._health_endpoints import health_readiness + + mock_prisma = MagicMock() + mock_prisma.health_check = AsyncMock(side_effect=PrismaError("nope")) + mock_prisma.attempt_db_reconnect = AsyncMock(side_effect=Exception("still nope")) + + _health_endpoints_module.db_health_cache = { + "status": "unknown", + "last_updated": datetime.now() - timedelta(seconds=60), + } + + response = Response() + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + result = await health_readiness(response=response) + + assert response.status_code == 503 + assert result["db"] == "disconnected" + assert result["status"] == "healthy" # body shape unchanged for back-compat + + +@pytest.mark.asyncio +async def test_health_readiness_returns_200_when_db_connected(): + """Happy path: connected DB keeps the legacy 200.""" + from fastapi import Response + + from litellm.proxy.health_endpoints._health_endpoints import health_readiness + + mock_prisma = MagicMock() + mock_prisma.health_check = AsyncMock() + + _health_endpoints_module.db_health_cache = { + "status": "unknown", + "last_updated": datetime.now() - timedelta(seconds=60), + } + + response = Response() + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + result = await health_readiness(response=response) + + assert response.status_code == 200 + assert result["db"] == "connected" + + +@pytest.mark.asyncio +async def test_health_readiness_returns_200_when_no_db_configured(): + """ + `prisma_client is None` means the operator chose not to use a DB. That + is a valid configuration — the worker should still report ready. We + only flip to 503 when a DB *was* configured but is unreachable. + """ + from fastapi import Response + + from litellm.proxy.health_endpoints._health_endpoints import health_readiness + + response = Response() + with patch("litellm.proxy.proxy_server.prisma_client", None): + result = await health_readiness(response=response) + + assert response.status_code == 200 + assert result["db"] == "Not connected" + + +def test_clean_endpoint_data_strips_credentials_keeps_routing_fields(): + """ + _clean_endpoint_data() drops credentials but leaves api_base / + api_version intact — the per-caller hide/show happens in the endpoint + layer based on user role, not in the cleaning helper. This guarantees + proxy admins continue to see those fields in the /health response. + """ + from litellm.proxy.health_check import _clean_endpoint_data + + raw = { + "model": "openai/gpt-4o", + "api_key": "sk-test", + "api_base": "https://example.test/v1", + "api_version": "2024-10-21", + "aws_access_key_id": "AKIAEXAMPLE", + } + + cleaned = _clean_endpoint_data(raw, details=True) + + assert "api_key" not in cleaned + assert "aws_access_key_id" not in cleaned + assert cleaned.get("api_base") == "https://example.test/v1" + assert cleaned.get("api_version") == "2024-10-21" diff --git a/tests/test_litellm/proxy/hooks/test_rate_limiter_toctou.py b/tests/test_litellm/proxy/hooks/test_rate_limiter_toctou.py new file mode 100644 index 00000000000..ceea5de7991 --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_rate_limiter_toctou.py @@ -0,0 +1,489 @@ +""" +Tests validating TOCTOU race condition in batch + dynamic rate limiters. + +Issue: rate-limit check (read_only=True) and counter increment happen as two +separate awaits. Concurrent requests all observe the same pre-increment state, +all pass validation, then all increment — bypassing the limit. + +Vulnerable code paths: +- litellm/proxy/hooks/batch_rate_limiter.py:181-248 + (_check_and_increment_batch_counters: should_rate_limit(read_only=True) + → validate → async_increment_tokens_with_ttl_preservation) +- litellm/proxy/hooks/dynamic_rate_limiter_v3.py:463-548 + (_check_rate_limits PHASE 1 read_only check → PHASE 3 increment) + +These tests EXPECTED to fail against current (vulnerable) code and pass once +check-and-increment becomes atomic. +""" + +import asyncio +import os +import sys +from typing import List + +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +import litellm +from litellm import DualCache, Router +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.hooks.batch_rate_limiter import BatchFileUsage +from litellm.proxy.hooks.dynamic_rate_limiter_v3 import ( + _PROXY_DynamicRateLimitHandlerV3 as DynamicRateLimitHandler, +) +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, +) +from litellm.proxy.utils import InternalUsageCache, hash_token + + +def _make_phase1_barrier(num_concurrent: int, timeout: float = 0.1): + """ + Sync primitive that, pre-fix, forces all N concurrent coroutines to finish + their read-only Phase 1 check before any proceeds to Phase 3 increment — + mimicking asyncio I/O scheduling under load on the vulnerable code. + + Wraps `should_rate_limit` so on `read_only=True` calls it waits until N + callers arrive (TOCTOU window opened) OR `timeout` elapses (post-fix path: + the limiter's serialization lock prevents N from ever reaching the + barrier; the timeout lets the holder proceed so the lock can do its job). + + Pre-fix: barrier fills before timeout → all see same state → bypass observed. + Post-fix: only lock-holder reaches barrier → times out → serial execution + enforces limit. + """ + arrived = 0 + all_arrived = asyncio.Event() + + def wrap(original): + async def patched(*args, **kwargs): + result = await original(*args, **kwargs) + if kwargs.get("read_only"): + nonlocal arrived + arrived += 1 + if arrived >= num_concurrent: + all_arrived.set() + try: + await asyncio.wait_for(all_arrived.wait(), timeout=timeout) + except asyncio.TimeoutError: + pass + return result + + return patched + + return wrap + + +@pytest.mark.asyncio +async def test_batch_limiter_concurrent_bypasses_tpm_via_toctou(): + """ + 5 concurrent batch submissions of 40 tokens each against TPM=100 limit. + + Sequential semantics: only 2 batches fit (2 * 40 = 80 ≤ 100, 3rd at 120 fails). + With TOCTOU: all 5 succeed → 200 tokens consumed, 100% over limit. + + Demonstrates batch_rate_limiter.py:183-248 multi-phase flaw. + """ + NUM_CONCURRENT = 5 + BATCH_TOKENS = 40 + TPM_LIMIT = 100 + + dual_cache = DualCache() + internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) + rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + internal_usage_cache=internal_usage_cache + ) + batch_limiter = rate_limiter._get_batch_rate_limiter() + assert batch_limiter is not None + + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("toctou-batch-key"), + tpm_limit=TPM_LIMIT, + rpm_limit=1000, + ) + batch_usage = BatchFileUsage(total_tokens=BATCH_TOKENS, request_count=1) + + barrier = _make_phase1_barrier(NUM_CONCURRENT) + rate_limiter.should_rate_limit = barrier(rate_limiter.should_rate_limit) + + results = await asyncio.gather( + *[ + batch_limiter._check_and_increment_batch_counters( + user_api_key_dict=user_api_key_dict, + data={}, + batch_usage=batch_usage, + ) + for _ in range(NUM_CONCURRENT) + ], + return_exceptions=True, + ) + + successes = [r for r in results if not isinstance(r, Exception)] + rejections = [r for r in results if isinstance(r, Exception)] + total_consumed = len(successes) * BATCH_TOKENS + max_allowed_successes = TPM_LIMIT // BATCH_TOKENS # 2 + + assert len(successes) <= max_allowed_successes, ( + f"TOCTOU bypass: {len(successes)}/{NUM_CONCURRENT} concurrent batches " + f"passed despite TPM={TPM_LIMIT}. Consumed {total_consumed} tokens " + f"({total_consumed - TPM_LIMIT} over limit). " + f"Atomic check-and-increment would allow ≤{max_allowed_successes}. " + f"Rejections: {len(rejections)}" + ) + + +@pytest.mark.asyncio +async def test_batch_limiter_uses_atomic_check_and_increment(): + """ + Regression test: batch limiter routes through + `atomic_check_and_increment_by_n` rather than the legacy two-phase + pattern (read_only=True check + separate async_increment_tokens_with_ttl_preservation). + + Ensures future refactors don't reintroduce the TOCTOU window. + """ + dual_cache = DualCache() + internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) + rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + internal_usage_cache=internal_usage_cache + ) + batch_limiter = rate_limiter._get_batch_rate_limiter() + assert batch_limiter is not None + + call_log: List[str] = [] + original_atomic = rate_limiter.atomic_check_and_increment_by_n + original_should = rate_limiter.should_rate_limit + + async def logging_atomic(*args, **kwargs): + call_log.append("atomic_check_and_increment_by_n") + return await original_atomic(*args, **kwargs) + + async def logging_should(*args, **kwargs): + call_log.append(f"should_rate_limit(read_only={kwargs.get('read_only')})") + return await original_should(*args, **kwargs) + + rate_limiter.atomic_check_and_increment_by_n = logging_atomic + rate_limiter.should_rate_limit = logging_should + + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("atomic-test-key"), + tpm_limit=10000, + rpm_limit=1000, + ) + + await batch_limiter._check_and_increment_batch_counters( + user_api_key_dict=user_api_key_dict, + data={}, + batch_usage=BatchFileUsage(total_tokens=50, request_count=1), + ) + + assert "atomic_check_and_increment_by_n" in call_log, ( + f"Batch limiter must route through atomic_check_and_increment_by_n. " + f"Calls observed: {call_log}" + ) + legacy_calls = [c for c in call_log if c.startswith("should_rate_limit(")] + assert not legacy_calls, ( + f"Batch limiter must not call should_rate_limit directly (legacy " + f"two-phase pattern). Observed: {legacy_calls}" + ) + + +@pytest.mark.asyncio +async def test_dynamic_rate_limiter_v3_concurrent_bypasses_model_capacity(): + """ + DynamicRateLimitHandler PHASE 1 (read_only check) → PHASE 3 (increment) + is non-atomic: dynamic_rate_limiter_v3.py:463-548. + + With TPM=100 model capacity and 5 concurrent priority="high" requests + each consuming the full model_saturation_check counter, all observe the + same Phase 1 state (counter=0), all pass, all proceed to Phase 3. + + Sequential atomic semantics would block requests once the model counter + reaches its limit. TOCTOU lets all pass Phase 1 simultaneously. + """ + NUM_CONCURRENT = 10 + MODEL_RPM = 2 + # Sequential bound: dynamic limiter rejects when `counter > current_limit` + # (strict `>`), so a request whose Phase 1 sees counter=RPM still passes + # (RPM is not strictly greater). Atomic execution therefore admits up to + # RPM + 1 successes before the next sees counter > RPM. + MAX_SEQUENTIAL_SUCCESSES = MODEL_RPM + 1 + + os.environ["LITELLM_LICENSE"] = "test-license-key" + litellm.priority_reservation = {"high": 0.9, "low": 0.1} + + dual_cache = DualCache() + handler = DynamicRateLimitHandler(internal_usage_cache=dual_cache) + + model = "toctou-dyn-model" + llm_router = Router( + model_list=[ + { + "model_name": model, + "litellm_params": { + "model": "gpt-3.5-turbo", + "api_key": "test-key", + "api_base": "test-base", + "rpm": MODEL_RPM, + }, + } + ] + ) + handler.update_variables(llm_router=llm_router) + + barrier = _make_phase1_barrier(NUM_CONCURRENT) + handler.v3_limiter.should_rate_limit = barrier(handler.v3_limiter.should_rate_limit) + + from litellm.types.router import ModelGroupInfo + + model_group_info = ModelGroupInfo( + model_group=model, + providers=["openai"], + rpm=MODEL_RPM, + tpm=None, + ) + + async def one_request(idx: int): + user = UserAPIKeyAuth(api_key=hash_token(f"dyn-key-{idx}")) + user.metadata = {"priority": "high"} + try: + await handler._check_rate_limits( + model=model, + model_group_info=model_group_info, + user_api_key_dict=user, + priority="high", + saturation=0.0, + data={}, + ) + return "OK" + except Exception as e: + return e + + results = await asyncio.gather( + *[one_request(i) for i in range(NUM_CONCURRENT)], + return_exceptions=True, + ) + successes = [r for r in results if r == "OK"] + + assert len(successes) <= MAX_SEQUENTIAL_SUCCESSES, ( + f"TOCTOU bypass in DynamicRateLimitHandler: {len(successes)}/{NUM_CONCURRENT} " + f"concurrent requests passed Phase 1 + Phase 3 despite model RPM={MODEL_RPM}. " + f"Atomic check-and-increment would block once counter > RPM " + f"(at most {MAX_SEQUENTIAL_SUCCESSES} sequential successes)." + ) + + +@pytest.mark.asyncio +async def test_dynamic_rate_limiter_v3_uses_atomic_check_and_increment(): + """ + Regression test: dynamic limiter's enforced descriptors flow through + `atomic_check_and_increment_by_n`, not the legacy + read_only=True check followed by a separate read_only=False increment. + + When priority is enforced (saturation >= threshold), priority_model is + bundled into the atomic call alongside model_saturation_check. When not + enforced, priority counter is incremented for tracking only. + """ + os.environ["LITELLM_LICENSE"] = "test-license-key" + litellm.priority_reservation = {"high": 0.9, "low": 0.1} + + dual_cache = DualCache() + handler = DynamicRateLimitHandler(internal_usage_cache=dual_cache) + + model = "atomic-dyn-model" + llm_router = Router( + model_list=[ + { + "model_name": model, + "litellm_params": { + "model": "gpt-3.5-turbo", + "api_key": "test-key", + "api_base": "test-base", + "tpm": 1000, + }, + } + ] + ) + handler.update_variables(llm_router=llm_router) + + atomic_descriptors_observed: List[List[str]] = [] + original_atomic = handler.v3_limiter.atomic_check_and_increment_by_n + + async def logging_atomic(*args, **kwargs): + ds = kwargs.get("descriptors") or (args[0] if args else []) + atomic_descriptors_observed.append([d["key"] for d in ds]) + return await original_atomic(*args, **kwargs) + + handler.v3_limiter.atomic_check_and_increment_by_n = logging_atomic + + from litellm.types.router import ModelGroupInfo + + user = UserAPIKeyAuth(api_key=hash_token("dyn-atomic-key")) + user.metadata = {"priority": "high"} + + await handler._check_rate_limits( + model=model, + model_group_info=ModelGroupInfo( + model_group=model, + providers=["openai"], + rpm=None, + tpm=1000, + ), + user_api_key_dict=user, + priority="high", + saturation=0.0, + data={}, + ) + + assert atomic_descriptors_observed, ( + "Dynamic limiter must route enforced descriptors through " + "atomic_check_and_increment_by_n (no legacy read_only=True / " + "separate-increment pattern)." + ) + assert "model_saturation_check" in atomic_descriptors_observed[0], ( + f"Expected model_saturation_check in atomic descriptor set. " + f"Got: {atomic_descriptors_observed}" + ) + + +@pytest.mark.asyncio +async def test_batch_zero_token_consumes_rpm_only(): + """ + Zero-token batch (e.g. metadata-only call) should still increment RPM + counter but NOT TPM counter. + + Edge case from review: `if inc_amount <= 0: continue` in + `atomic_check_and_increment_by_n` skips descriptor counters whose + increment is zero. Verifies asymmetric quota consumption is intentional + and observable: an RPM-bounded but TPM-free request path stays bounded + by RPM alone. + """ + dual_cache = DualCache() + internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) + rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + internal_usage_cache=internal_usage_cache + ) + batch_limiter = rate_limiter._get_batch_rate_limiter() + assert batch_limiter is not None + + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("zero-token-key"), + tpm_limit=100, + rpm_limit=3, + ) + zero_batch = BatchFileUsage(total_tokens=0, request_count=1) + + # 3 zero-token batches must succeed (RPM=3 allows). 4th must hit RPM cap, + # NOT TPM (because token counter never increments past 0). + for i in range(3): + await batch_limiter._check_and_increment_batch_counters( + user_api_key_dict=user_api_key_dict, + data={}, + batch_usage=zero_batch, + ) + + # Inspect counters: RPM key incremented to 3, TPM key absent (or 0). + rpm_key = rate_limiter.create_rate_limit_keys( + "api_key", user_api_key_dict.api_key or "", "requests" + ) + tpm_key = rate_limiter.create_rate_limit_keys( + "api_key", user_api_key_dict.api_key or "", "tokens" + ) + rpm_val = await internal_usage_cache.async_get_cache( + key=rpm_key, litellm_parent_otel_span=None, local_only=True + ) + tpm_val = await internal_usage_cache.async_get_cache( + key=tpm_key, litellm_parent_otel_span=None, local_only=True + ) + assert int(rpm_val or 0) == 3, f"RPM counter must reach 3, got {rpm_val}" + assert tpm_val in ( + None, + 0, + "0", + ), f"TPM counter must remain unset/0 for zero-token batches, got {tpm_val}" + + # 4th attempt: RPM exhausted -> 429. + from fastapi import HTTPException + + with pytest.raises(HTTPException) as exc: + await batch_limiter._check_and_increment_batch_counters( + user_api_key_dict=user_api_key_dict, + data={}, + batch_usage=zero_batch, + ) + assert exc.value.status_code == 429 + + +@pytest.mark.asyncio +async def test_dynamic_rate_limiter_v3_fails_closed_on_unknown_descriptor(): + """ + Fail-closed guard: when atomic_check_and_increment_by_n returns + overall_code=OVER_LIMIT but with a descriptor_key the dispatcher does + not recognize, the dynamic limiter must raise 429 rather than silently + fall through. + + Reproduces by patching atomic_check_and_increment_by_n to return an + OVER_LIMIT response carrying an unknown descriptor_key. + """ + from fastapi import HTTPException + + os.environ["LITELLM_LICENSE"] = "test-license-key" + litellm.priority_reservation = {"high": 0.9, "low": 0.1} + + dual_cache = DualCache() + handler = DynamicRateLimitHandler(internal_usage_cache=dual_cache) + + model = "fail-closed-model" + llm_router = Router( + model_list=[ + { + "model_name": model, + "litellm_params": { + "model": "gpt-3.5-turbo", + "api_key": "test-key", + "api_base": "test-base", + "tpm": 1000, + }, + } + ] + ) + handler.update_variables(llm_router=llm_router) + + async def fake_atomic(*args, **kwargs): + return { + "overall_code": "OVER_LIMIT", + "statuses": [ + { + "code": "OVER_LIMIT", + "current_limit": 100, + "limit_remaining": 0, + "rate_limit_type": "tokens", + "descriptor_key": "future_unrecognized_descriptor", + } + ], + } + + handler.v3_limiter.atomic_check_and_increment_by_n = fake_atomic + + from litellm.types.router import ModelGroupInfo + + user = UserAPIKeyAuth(api_key=hash_token("fail-closed-key")) + user.metadata = {"priority": "high"} + + with pytest.raises(HTTPException) as exc: + await handler._check_rate_limits( + model=model, + model_group_info=ModelGroupInfo( + model_group=model, + providers=["openai"], + rpm=None, + tpm=1000, + ), + user_api_key_dict=user, + priority="high", + saturation=0.0, + data={}, + ) + assert ( + exc.value.status_code == 429 + ), f"Expected 429 fail-closed on unknown descriptor; got {exc.value.status_code}" diff --git a/tests/test_litellm/proxy/management_endpoints/test_access_group_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_access_group_endpoints.py index cd2eb789589..016e10859b6 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_access_group_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_access_group_endpoints.py @@ -738,15 +738,32 @@ def test_delete_access_group_patches_cached_team_and_key( return_value=None ) - # Build cached key object (returned from user_api_key_cache) - if key_cache_group_ids is not None: - cached_key = UserAPIKeyAuth( - token="hashed-key-1", - access_group_ids=list(key_cache_group_ids), + # user_api_key_cache is queried both for teams (fallback after dual_cache) and + # hashed keys — return the right stub per ``key``. A single AsyncMock(return_value=key) + # would wrongly serve the key blob for ``team_id:team-1`` and trigger team patching. + # Use a synchronous side_effect (not async def): AsyncMock awaits coroutine side_effects + # inconsistently across Python/unittest versions; sync returns are awaited as immediate results. + def user_cache_get_side_effect(*args, **kwargs): + cache_key = ( + kwargs.get("key") if "key" in kwargs else (args[0] if args else None) ) - mock_cache.async_get_cache = AsyncMock(return_value=cached_key) - else: - mock_cache.async_get_cache = AsyncMock(return_value=None) + if cache_key == "team_id:team-1": + if team_cache_group_ids is None: + return None + return LiteLLM_TeamTableCachedObj( + team_id="team-1", + access_group_ids=list(team_cache_group_ids), + ) + if cache_key == "hashed-key-1": + if key_cache_group_ids is None: + return None + return UserAPIKeyAuth( + token="hashed-key-1", + access_group_ids=list(key_cache_group_ids), + ) + return None + + mock_cache.async_get_cache = AsyncMock(side_effect=user_cache_get_side_effect) resp = client.delete("/v1/access_group/ag-to-delete") assert resp.status_code == 204 @@ -803,7 +820,7 @@ def test_delete_access_group_patches_cached_team_and_key( def test_delete_access_group_patches_key_cached_as_dict(client_and_mocks): - """Delete correctly patches a key cached as a raw dict (not UserAPIKeyAuth).""" + """Delete patches key cache — mock returns UserAPIKeyAuth (what UserApiKeyCache emits after deserialize).""" client, mock_prisma, mock_access_group_table, mock_cache, mock_proxy_logging = ( client_and_mocks ) @@ -826,12 +843,24 @@ def test_delete_access_group_patches_key_cached_as_dict(client_and_mocks): return_value=None ) - # Key cached as a plain dict (as can happen with Redis serialization) + # Serialized shape from Redis dict; UserApiKeyCache.async_get_cache(model_type=...) yields a model — simulate that. + cached_key_payload = { + "token": "hashed-key-dict", + "access_group_ids": ["ag-to-delete", "ag-other"], + } + + def user_cache_get_dict_when_key_matches(*args, **kwargs): + cache_key = ( + kwargs.get("key") if "key" in kwargs else (args[0] if args else None) + ) + if cache_key == "team_id:team-1": + return None + if cache_key == "hashed-key-dict": + return UserAPIKeyAuth.model_validate(cached_key_payload) + return None + mock_cache.async_get_cache = AsyncMock( - return_value={ - "token": "hashed-key-dict", - "access_group_ids": ["ag-to-delete", "ag-other"], - } + side_effect=user_cache_get_dict_when_key_matches ) resp = client.delete("/v1/access_group/ag-to-delete") diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 0362d6f97d9..e668672dd2a 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -5512,6 +5512,9 @@ async def test_update_team_guardrails_with_org_id(): return_value=mock_updated_team ) mock_prisma.jsonify_team_object = MagicMock(side_effect=lambda db_data: db_data) + # async_get_cache must be an AsyncMock so `await` in get_org_object works + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() # Mock llm_router mock_router = MagicMock() diff --git a/tests/test_litellm/proxy/middleware/test_prometheus_auth_middleware.py b/tests/test_litellm/proxy/middleware/test_prometheus_auth_middleware.py index 9fd244d9c3f..310ee11573b 100644 --- a/tests/test_litellm/proxy/middleware/test_prometheus_auth_middleware.py +++ b/tests/test_litellm/proxy/middleware/test_prometheus_auth_middleware.py @@ -26,6 +26,15 @@ async def fake_valid_auth(request, api_key): return +async def fake_valid_auth_reads_body(request, api_key, **kwargs): + """ + Like real user_api_key_auth, consumes the ASGI body stream. Regression test + for successful auth passing a drained receive to the inner app (hang). + """ + await request.body() + return + + async def fake_invalid_auth(request, api_key): print("running fake invalid auth", request, api_key) # Simulate invalid auth by raising an exception. @@ -62,6 +71,28 @@ def app_with_middleware(): return app +def test_valid_auth_metrics_after_body_consumed(app_with_middleware, monkeypatch): + """ + Auth that reads the request body must not cause /metrics to hang on success. + """ + litellm.require_auth_for_metrics_endpoint = True + monkeypatch.setattr( + "litellm.proxy.middleware.prometheus_auth_middleware.user_api_key_auth", + fake_valid_auth_reads_body, + ) + + client = TestClient(app_with_middleware) + headers = {SpecialHeaders.openai_authorization.value: "valid"} + + response = client.get("/metrics", headers=headers) + assert response.status_code == 200, response.text + assert response.json() == {"msg": "metrics OK"} + + response = client.get("/metrics/", headers=headers) + assert response.status_code == 200, response.text + assert response.json() == {"msg": "metrics OK"} + + def test_valid_auth_metrics(app_with_middleware, monkeypatch): """ Test that a request to /metrics (and /metrics/) with valid auth headers passes. diff --git a/tests/test_litellm/proxy/test_filter_models_by_team_access_group.py b/tests/test_litellm/proxy/test_filter_models_by_team_access_group.py new file mode 100644 index 00000000000..2d8a9f30c1b --- /dev/null +++ b/tests/test_litellm/proxy/test_filter_models_by_team_access_group.py @@ -0,0 +1,236 @@ +""" +Tests for _filter_models_by_team_id resolving access group names. + +Verifies that when a team's `models` field contains an access group name +(e.g., "Group-A"), the filter resolves it to the member model names before +looking up deployments — matching the behavior of the auth path in +auth_checks.py:model_in_access_group(). +""" + +import os +import sys +from unittest.mock import AsyncMock, MagicMock + +import pytest + +sys.path.insert(0, os.path.abspath("../../..")) + +from litellm.proxy.proxy_server import _filter_models_by_team_id + + +def _make_model(model_name: str, model_id: str, access_groups: list[str] = None): + """Helper to build a model dict matching the router's format.""" + return { + "model_name": model_name, + "litellm_params": {"model": model_name}, + "model_info": { + "id": model_id, + "access_groups": access_groups or [], + }, + } + + +def _make_team(models: list[str], team_id: str = "team_alpha"): + """Helper to build a mock team DB object.""" + mock = MagicMock() + mock.model_dump.return_value = { + "team_id": team_id, + "team_alias": "Team Alpha", + "models": models, + "max_budget": None, + "spend": 0.0, + "blocked": False, + "members_with_roles": [], + "metadata": {}, + } + return mock + + +@pytest.mark.asyncio +async def test_filter_resolves_access_group_names(): + """ + When team.models contains an access group name, _filter_models_by_team_id + should resolve it to the member models and return only those deployments. + """ + # Models on the proxy + gpt4o = _make_model("gpt-4o", "id-1", ["Group-A"]) + gpt5 = _make_model("gpt-5", "id-2", ["Group-A"]) + claude = _make_model("claude-3", "id-3", ["Group-B"]) + + all_models = [gpt4o, gpt5, claude] + + # Router mock + mock_router = MagicMock() + # get_model_access_groups returns {group_name: [model_names]} + mock_router.get_model_access_groups.return_value = { + "Group-A": ["gpt-4o", "gpt-5"], + "Group-B": ["claude-3"], + } + + # get_model_list returns deployments matching a model_name + def fake_get_model_list(model_name=None, team_id=None): + return [m for m in all_models if m["model_name"] == model_name] + + mock_router.get_model_list = MagicMock(side_effect=fake_get_model_list) + + # Team has models: ["Group-A"] — an access group name, not a literal model + team_db = _make_team(models=["Group-A"]) + + # Prisma mock + mock_prisma = MagicMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_db) + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + + result = await _filter_models_by_team_id( + all_models=all_models, + team_id="team_alpha", + prisma_client=mock_prisma, + llm_router=mock_router, + ) + + result_ids = {m["model_info"]["id"] for m in result} + # Should include gpt-4o and gpt-5 (Group-A), but NOT claude-3 (Group-B) + assert result_ids == { + "id-1", + "id-2", + }, f"Expected Group-A models only, got {result_ids}" + + # Verify DB fallback query received resolved model names, not access group name + call_kwargs = mock_prisma.db.litellm_proxymodeltable.find_many.call_args[1] + assert set(call_kwargs["where"]["model_name"]["in"]) == { + "gpt-4o", + "gpt-5", + }, "find_many should receive resolved model names, not the access group name" + + +@pytest.mark.asyncio +async def test_filter_resolves_mix_of_access_groups_and_literal_names(): + """ + When team.models contains both an access group name and a literal model name, + both should be resolved correctly. + """ + gpt4o = _make_model("gpt-4o", "id-1", ["Group-A"]) + gpt5 = _make_model("gpt-5", "id-2", ["Group-A"]) + claude = _make_model("claude-3", "id-3", ["Group-B"]) + mistral = _make_model("mistral-large", "id-4", []) # no access group + + all_models = [gpt4o, gpt5, claude, mistral] + + mock_router = MagicMock() + mock_router.get_model_access_groups.return_value = { + "Group-A": ["gpt-4o", "gpt-5"], + "Group-B": ["claude-3"], + } + + def fake_get_model_list(model_name=None, team_id=None): + return [m for m in all_models if m["model_name"] == model_name] + + mock_router.get_model_list = MagicMock(side_effect=fake_get_model_list) + + # Team has access to Group-A (access group) + mistral-large (literal name) + team_db = _make_team(models=["Group-A", "mistral-large"]) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_db) + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + + result = await _filter_models_by_team_id( + all_models=all_models, + team_id="team_alpha", + prisma_client=mock_prisma, + llm_router=mock_router, + ) + + result_ids = {m["model_info"]["id"] for m in result} + # Group-A models + mistral-large, but NOT claude-3 + assert result_ids == { + "id-1", + "id-2", + "id-4", + }, f"Expected Group-A + mistral-large, got {result_ids}" + + +@pytest.mark.asyncio +async def test_filter_excludes_models_from_other_access_group(): + """ + Models belonging only to a different access group must not appear in results. + """ + gpt4o = _make_model("gpt-4o", "id-1", ["Group-A"]) + claude = _make_model("claude-3", "id-3", ["Group-B"]) + llama = _make_model("llama-4", "id-4", ["Group-B"]) + + all_models = [gpt4o, claude, llama] + + mock_router = MagicMock() + mock_router.get_model_access_groups.return_value = { + "Group-A": ["gpt-4o"], + "Group-B": ["claude-3", "llama-4"], + } + + def fake_get_model_list(model_name=None, team_id=None): + return [m for m in all_models if m["model_name"] == model_name] + + mock_router.get_model_list = MagicMock(side_effect=fake_get_model_list) + + team_db = _make_team(models=["Group-A"]) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_db) + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + + result = await _filter_models_by_team_id( + all_models=all_models, + team_id="team_alpha", + prisma_client=mock_prisma, + llm_router=mock_router, + ) + + result_names = {m["model_name"] for m in result} + assert "claude-3" not in result_names, "Group-B model should not be accessible" + assert "llama-4" not in result_names, "Group-B model should not be accessible" + assert "gpt-4o" in result_names, "Group-A model should be accessible" + + +@pytest.mark.asyncio +async def test_filter_db_fallback_receives_resolved_model_names(): + """ + When get_model_list returns no results (forcing the DB fallback path), + the DB query should receive resolved model names, not the raw access group name. + """ + gpt4o = _make_model("gpt-4o", "id-1", ["Group-A"]) + all_models = [gpt4o] + + mock_router = MagicMock() + mock_router.get_model_access_groups.return_value = { + "Group-A": ["gpt-4o", "gpt-5"], + } + # get_model_list returns nothing — forces reliance on the DB fallback + mock_router.get_model_list = MagicMock(return_value=[]) + + team_db = _make_team(models=["Group-A"]) + + # DB returns a model that the router didn't find + mock_db_model = MagicMock() + mock_db_model.model_id = "id-db-1" + + mock_prisma = MagicMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_db) + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock( + return_value=[mock_db_model] + ) + + result = await _filter_models_by_team_id( + all_models=all_models, + team_id="team_alpha", + prisma_client=mock_prisma, + llm_router=mock_router, + ) + + # Verify DB query received resolved names, not "Group-A" + call_kwargs = mock_prisma.db.litellm_proxymodeltable.find_many.call_args[1] + queried_names = set(call_kwargs["where"]["model_name"]["in"]) + assert queried_names == { + "gpt-4o", + "gpt-5", + }, f"DB query should receive resolved model names, got {queried_names}" + assert "Group-A" not in queried_names, "Raw access group name should not be in DB query" diff --git a/tests/test_litellm/proxy/test_lazy_openapi_snapshot.py b/tests/test_litellm/proxy/test_lazy_openapi_snapshot.py new file mode 100644 index 00000000000..8bc39c93eeb --- /dev/null +++ b/tests/test_litellm/proxy/test_lazy_openapi_snapshot.py @@ -0,0 +1,35 @@ +from litellm.proxy._lazy_openapi_snapshot import _normalize_operation_ids + + +def test_normalize_operation_ids_uses_each_http_method(): + paths = { + "/proxy/{endpoint}": { + "delete": {"operationId": "proxy_route_proxy__endpoint__put"}, + "get": {"operationId": "proxy_route_proxy__endpoint__put"}, + "post": {"operationId": "proxy_route_proxy__endpoint__put"}, + "put": {"operationId": "proxy_route_proxy__endpoint__put"}, + } + } + + _normalize_operation_ids(paths) + + operations = paths["/proxy/{endpoint}"] + assert operations["delete"]["operationId"] == "proxy_route_proxy__endpoint__delete" + assert operations["get"]["operationId"] == "proxy_route_proxy__endpoint__get" + assert operations["post"]["operationId"] == "proxy_route_proxy__endpoint__post" + assert operations["put"]["operationId"] == "proxy_route_proxy__endpoint__put" + + +def test_normalize_operation_ids_preserves_custom_ids(): + paths = { + "/proxy/{endpoint}": { + "get": {"operationId": "custom_operation"}, + "post": {"operationId": "custom_operation"}, + } + } + + _normalize_operation_ids(paths) + + operations = paths["/proxy/{endpoint}"] + assert operations["get"]["operationId"] == "custom_operation" + assert operations["post"]["operationId"] == "custom_operation" diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 37e53005650..465ce579e09 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -5505,6 +5505,202 @@ async def test_reseed_warms_cache_even_on_zero_db_spend(): ps.prisma_client = orig_prisma +# ----------------------------------------------------------------------------- +# /config/update — critical paths only. +# +# These exercise the four behaviors that broke or changed in the rewrite of +# update_config (litellm/proxy/proxy_server.py): targeted per-section writes, +# the removal of the store_model_in_db gate, env var encryption, and the +# success_callback / litellm_settings merge semantics. All other branches +# (auth, missing-DB, slack auto-enable, router_settings merge) are covered +# implicitly or by upstream tests. +# ----------------------------------------------------------------------------- + + +class _FakeRow: + def __init__(self, param_name, param_value): + self.param_name = param_name + self.param_value = param_value + + +class _FakeLitellmConfig: + def __init__(self, initial_rows=None): + self.rows = dict(initial_rows or {}) + self.upsert_calls: list = [] + self.find_first = AsyncMock(side_effect=self._find_first) + self.upsert = AsyncMock(side_effect=self._upsert) + + async def _find_first(self, where=None): + if where and "param_name" in where: + name = where["param_name"] + if name in self.rows: + return _FakeRow(name, self.rows[name]) + return None + + async def _upsert(self, where=None, data=None): + name = where["param_name"] + raw = data["update"]["param_value"] + value = json.loads(raw) if isinstance(raw, str) else raw + self.rows[name] = value + self.upsert_calls.append((name, value)) + + +class _FakePrismaClient: + def __init__(self, initial_rows=None): + self.db = mock.MagicMock() + self.db.litellm_config = _FakeLitellmConfig(initial_rows=initial_rows) + self.jsonify_object = lambda obj: obj + + +@pytest.fixture +def _update_config_setup(monkeypatch): + """Install fakes for the /config/update endpoint and return (client, prisma).""" + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth as auth_dep + + def _install(initial_rows=None, store_model_in_db=True): + prisma = _FakePrismaClient(initial_rows=initial_rows) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma) + monkeypatch.setattr( + "litellm.proxy.proxy_server.store_model_in_db", store_model_in_db + ) + monkeypatch.setattr( + "litellm.proxy.proxy_server.encrypt_value_helper", + lambda value, **_: f"enc:{value}", + ) + monkeypatch.setattr( + "litellm.proxy.proxy_server.invalidate_config_param", + AsyncMock(return_value=None), + ) + from litellm.proxy.proxy_server import proxy_config as real_proxy_config + + monkeypatch.setattr( + real_proxy_config, "add_deployment", AsyncMock(return_value=None) + ) + + original_overrides = app.dependency_overrides.copy() + app.dependency_overrides[auth_dep] = lambda: UserAPIKeyAuth( + user_id="test_admin", + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-1234", + ) + client = TestClient(app) + + def _restore(): + app.dependency_overrides = original_overrides + + return client, prisma, _restore + + return _install + + +def test_update_config_writes_only_sent_section(_update_config_setup): + """A request that only touches general_settings must not write any other + section row, and must leave previously-written rows byte-identical.""" + client, prisma, restore = _update_config_setup( + initial_rows={ + "litellm_settings": {"drop_params": True}, + "environment_variables": {"FOO": "enc:bar"}, + } + ) + try: + resp = client.post( + "/config/update", + json={"general_settings": {"store_prompts_in_spend_logs": True}}, + ) + assert resp.status_code == 200 + written = {name for name, _ in prisma.db.litellm_config.upsert_calls} + assert written == {"general_settings"} + assert prisma.db.litellm_config.rows["litellm_settings"] == { + "drop_params": True + } + assert prisma.db.litellm_config.rows["environment_variables"] == { + "FOO": "enc:bar" + } + finally: + restore() + + +def test_update_config_can_flip_store_model_in_db_when_currently_false( + _update_config_setup, +): + """The endpoint used to refuse all writes when store_model_in_db was + False, blocking the very request that would flip it to True.""" + client, prisma, restore = _update_config_setup(store_model_in_db=False) + try: + resp = client.post( + "/config/update", json={"general_settings": {"store_model_in_db": True}} + ) + assert resp.status_code == 200 + assert ( + prisma.db.litellm_config.rows["general_settings"]["store_model_in_db"] + is True + ) + finally: + restore() + + +def test_update_config_environment_variables_encrypted_before_write( + _update_config_setup, +): + """env var values must be encrypted before they hit the DB row.""" + client, prisma, restore = _update_config_setup() + try: + resp = client.post( + "/config/update", + json={"environment_variables": {"OPENAI_API_KEY": "sk-secret"}}, + ) + assert resp.status_code == 200 + stored = prisma.db.litellm_config.rows["environment_variables"] + assert stored == {"OPENAI_API_KEY": "enc:sk-secret"} + finally: + restore() + + +def test_update_config_litellm_settings_request_wins_for_non_callback_keys( + _update_config_setup, +): + """Sending {"drop_params": False} when the row holds drop_params: True + must persist False (request wins). Untouched keys preserved.""" + client, prisma, restore = _update_config_setup( + initial_rows={ + "litellm_settings": {"drop_params": True, "set_verbose": True}, + } + ) + try: + resp = client.post( + "/config/update", json={"litellm_settings": {"drop_params": False}} + ) + assert resp.status_code == 200 + stored = prisma.db.litellm_config.rows["litellm_settings"] + assert stored["drop_params"] is False + assert stored["set_verbose"] is True + finally: + restore() + + +def test_update_config_success_callback_normalizes_existing_mixed_case( + _update_config_setup, +): + """Existing mixed-case callback names (written elsewhere) must be + normalized to lowercase before union, otherwise the union dedup misses + against the lowercase incoming entry and delete_callback (lowercase + lookup) cannot find the original.""" + client, prisma, restore = _update_config_setup( + initial_rows={"litellm_settings": {"success_callback": ["Langfuse", "SQS"]}} + ) + try: + resp = client.post( + "/config/update", + json={"litellm_settings": {"success_callback": ["langfuse"]}}, + ) + assert resp.status_code == 200 + stored = prisma.db.litellm_config.rows["litellm_settings"]["success_callback"] + assert set(stored) == {"langfuse", "sqs"} + finally: + restore() + + # --------------------------------------------------------------------------- # Lazy feature loading (LazyFeatureMiddleware) — verifies that optional # routers are NOT imported at module load and ARE imported on first request @@ -5513,9 +5709,6 @@ async def test_reseed_warms_cache_even_on_zero_db_spend(): # --------------------------------------------------------------------------- -import sys - - class TestLazyFeatureRegistry: """Sanity checks on the registry shape — guards against accidental edits.""" diff --git a/tests/test_litellm/proxy/test_redis_auth_cache_flag.py b/tests/test_litellm/proxy/test_redis_auth_cache_flag.py new file mode 100644 index 00000000000..d0cb5ec5465 --- /dev/null +++ b/tests/test_litellm/proxy/test_redis_auth_cache_flag.py @@ -0,0 +1,145 @@ +""" +Tests for the enable_redis_auth_cache litellm_settings flag. + +Verifies that _init_cache attaches Redis to user_api_key_cache only when +the flag is explicitly set to True, and leaves it in-memory-only otherwise. +""" + +from contextlib import contextmanager +import json +from unittest.mock import MagicMock, patch + +import pytest + +import litellm +import litellm.proxy.proxy_server as ps +from litellm.caching.caching import RedisCache +from litellm.caching.dual_cache import DualCache + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +class _FakeRedisCache(RedisCache): + """ + Minimal RedisCache subclass that passes isinstance checks without + requiring a real Redis connection. __init__ is bypassed so no + network calls are made. + """ + + def __init__(self): # noqa: super().__init__ skipped intentionally + self._store = {} + + def set_cache(self, key, value, **kwargs): # type: ignore[override] + # Enforce Redis JSON-serializable payload contract. + self._store[key] = json.dumps(value) + return True + + def get_cache(self, key, **kwargs): # type: ignore[override] + raw = self._store.get(key) + if raw is None: + return None + return json.loads(raw) + + +@contextmanager +def _patched_init_cache(litellm_settings: dict, cache_params: dict): + """ + Context manager that: + 1. Replaces the module-level globals with fresh DualCache instances. + 2. Patches ``litellm.Cache`` (locally imported inside _init_cache) so + it returns a fake cache whose ``.cache`` attribute is a + _FakeRedisCache (passes the isinstance guard in _init_cache). + 3. Extracts enable_redis_auth_cache from litellm_settings and passes it + as the second argument to _init_cache (matching production behaviour). + 4. Yields (user_api_key_cache, spend_counter_cache) after calling + _init_cache, then restores everything. + """ + fake_redis = _FakeRedisCache() + + mock_litellm_cache = MagicMock() + mock_litellm_cache.cache = fake_redis + + fresh_user_cache = DualCache() + fresh_spend_cache = DualCache() + + enable_redis_auth_cache = litellm_settings.get("enable_redis_auth_cache", False) + + with ( + patch.object(ps, "user_api_key_cache", fresh_user_cache), + patch.object(ps, "spend_counter_cache", fresh_spend_cache), + patch.object(ps, "llm_router", None), + # Cache is locally imported inside _init_cache: patch it at source. + patch("litellm.Cache", return_value=mock_litellm_cache), + ): + litellm.cache = None + ps.ProxyConfig()._init_cache(cache_params, enable_redis_auth_cache) + yield fresh_user_cache, fresh_spend_cache + + +# --------------------------------------------------------------------------- +# Tests +# --------------------------------------------------------------------------- + + +class TestRedisAuthCacheFlag: + def test_flag_true_attaches_redis_to_user_api_key_cache(self): + """When enable_redis_auth_cache=True, user_api_key_cache.redis_cache must be set.""" + with _patched_init_cache( + litellm_settings={"enable_redis_auth_cache": True}, + cache_params={"type": "redis", "host": "localhost", "port": 6379}, + ) as (user_cache, _): + assert user_cache.redis_cache is not None, ( + "Redis should be attached to user_api_key_cache when " + "enable_redis_auth_cache=True" + ) + + def test_flag_false_leaves_user_api_key_cache_in_memory_only(self): + """When enable_redis_auth_cache=False, user_api_key_cache must stay in-memory.""" + with _patched_init_cache( + litellm_settings={"enable_redis_auth_cache": False}, + cache_params={"type": "redis", "host": "localhost", "port": 6379}, + ) as (user_cache, _): + assert user_cache.redis_cache is None, ( + "user_api_key_cache must remain in-memory-only when " + "enable_redis_auth_cache=False" + ) + + def test_flag_absent_leaves_user_api_key_cache_in_memory_only(self): + """When enable_redis_auth_cache is not set at all, default is in-memory-only.""" + with _patched_init_cache( + litellm_settings={}, + cache_params={"type": "redis", "host": "localhost", "port": 6379}, + ) as (user_cache, _): + assert user_cache.redis_cache is None, ( + "user_api_key_cache must remain in-memory-only when " + "enable_redis_auth_cache is absent from litellm_settings" + ) + + def test_spend_counter_cache_always_gets_redis_regardless_of_flag(self): + """spend_counter_cache must receive Redis regardless of the auth-cache flag.""" + for flag_value in (True, False, None): + ls = ( + {"enable_redis_auth_cache": flag_value} + if flag_value is not None + else {} + ) + with _patched_init_cache( + litellm_settings=ls, + cache_params={"type": "redis", "host": "localhost", "port": 6379}, + ) as (_, spend_cache): + assert spend_cache.redis_cache is not None, ( + f"spend_counter_cache must always get Redis " + f"(enable_redis_auth_cache={flag_value!r})" + ) + + def test_flag_false_spend_gets_redis_but_user_cache_does_not(self): + """Explicit False: spend cache wired, auth cache left in-memory.""" + with _patched_init_cache( + litellm_settings={"enable_redis_auth_cache": False}, + cache_params={"type": "redis", "host": "localhost", "port": 6379}, + ) as (user_cache, spend_cache): + assert spend_cache.redis_cache is not None + assert user_cache.redis_cache is None diff --git a/tests/test_litellm/router_utils/test_router_utils_common_utils.py b/tests/test_litellm/router_utils/test_router_utils_common_utils.py index 02241d4bc92..465c6669ceb 100644 --- a/tests/test_litellm/router_utils/test_router_utils_common_utils.py +++ b/tests/test_litellm/router_utils/test_router_utils_common_utils.py @@ -6,6 +6,7 @@ import pytest from litellm import Router from litellm.router_utils.common_utils import ( _deployment_supports_web_search, + add_model_file_id_mappings, filter_team_based_models, filter_web_search_deployments, ) @@ -362,3 +363,112 @@ def test_invalidate_model_group_info_cache(): # Invalidate and verify cache is cleared router._invalidate_model_group_info_cache() assert router._cached_get_model_group_info.cache_info().currsize == 0 + + +class TestAddModelFileIdMappings: + """Test cases for add_model_file_id_mappings. + + The router may pass either a list of deployment dicts (multiple matched + deployments) or a single deployment dict (when a specific deployment was + resolved, e.g. because the requested model matched a `model_info.id`). + Both shapes must produce a `{model_id: file_id}` mapping by extracting + `model_info.id` from each deployment. + """ + + @staticmethod + def _make_response(file_id: str): + response = Mock() + response.id = file_id + return response + + def test_should_map_each_deployment_id_when_given_list(self): + deployments = [ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4"}, + "model_info": {"id": "deployment-1"}, + }, + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4"}, + "model_info": {"id": "deployment-2"}, + }, + ] + responses = [self._make_response("file-1"), self._make_response("file-2")] + + result = add_model_file_id_mappings(deployments, responses) + + assert result == {"deployment-1": "file-1", "deployment-2": "file-2"} + + def test_should_extract_model_info_id_when_given_single_deployment_dict(self): + """Regression test: when `_common_checks_available_deployment` resolves + a specific deployment (returned as a dict, not a list), the function + must still extract `model_info.id` rather than iterate over the + deployment's own keys (`model_name`, `litellm_params`, `model_info`). + """ + deployment = { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4", "api_key": "sk-test"}, + "model_info": {"id": "deployment-1", "mode": "chat"}, + } + responses = [self._make_response("file-1")] + + result = add_model_file_id_mappings(deployment, responses) + + assert result == {"deployment-1": "file-1"} + assert all(isinstance(v, str) for v in result.values()) + + def test_should_handle_batch_model_when_id_matches_model_name(self): + """Regression test for the batch-model case: when `model_info.id` is + intentionally set equal to `model_name`, the router resolves a single + deployment via `has_model_id` and returns it as a dict. The mapping + must contain only `{id: file_id}` with string values so the resulting + `LiteLLM_ManagedFileTable` Pydantic validation passes. + """ + deployment = { + "model_name": "openai/openai/gpt-5.5-batch", + "litellm_params": { + "model": "openai/gpt-5.5", + "api_key": "sk-test", + "tpm": 40000000, + "rpm": 15000, + }, + "model_info": { + "id": "openai/openai/gpt-5.5-batch", + "mode": "batch", + "base_model": "gpt-5.5", + "access_groups": ["default-models"], + }, + } + responses = [self._make_response("file-batch-1")] + + result = add_model_file_id_mappings(deployment, responses) + + # Bug case would have produced keys ["model_name", "litellm_params", + # "model_info"] with non-string values. + assert result == {"openai/openai/gpt-5.5-batch": "file-batch-1"} + assert "litellm_params" not in result + assert "model_info" not in result + + def test_should_skip_deployment_when_model_info_id_missing(self): + deployments = [ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4"}, + "model_info": {}, + }, + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4"}, + "model_info": {"id": "deployment-2"}, + }, + ] + responses = [self._make_response("file-1"), self._make_response("file-2")] + + result = add_model_file_id_mappings(deployments, responses) + + assert result == {"deployment-2": "file-2"} + + def test_should_return_empty_mapping_when_given_empty_list(self): + result = add_model_file_id_mappings([], []) + assert result == {} diff --git a/tests/test_litellm/test_anthropic_skills_transformation.py b/tests/test_litellm/test_anthropic_skills_transformation.py index 6761d671806..1b917f08ca9 100644 --- a/tests/test_litellm/test_anthropic_skills_transformation.py +++ b/tests/test_litellm/test_anthropic_skills_transformation.py @@ -70,6 +70,15 @@ class TestAnthropicSkillsConfigURLConstruction: ) assert url == f"{FAKE_API_BASE}/v1/skills/skill_abc123" + def test_url_with_skill_id_encodes_path_segment(self): + url = self.config.get_complete_url( + api_base=FAKE_API_BASE, + endpoint="skills", + skill_id="../../files?x=1#frag", + ) + + assert url == f"{FAKE_API_BASE}/v1/skills/..%2F..%2Ffiles%3Fx%3D1%23frag" + def test_url_falls_back_to_anthropic_default(self): with patch( "litellm.llms.anthropic.common_utils.AnthropicModelInfo.get_api_base", diff --git a/tests/test_litellm/test_openai_embedding_encoding_format_default.py b/tests/test_litellm/test_openai_embedding_encoding_format_default.py new file mode 100644 index 00000000000..94e4e3c81e5 --- /dev/null +++ b/tests/test_litellm/test_openai_embedding_encoding_format_default.py @@ -0,0 +1,124 @@ +from unittest.mock import MagicMock, patch + +import pytest + +from litellm import embedding + + +@pytest.mark.parametrize( + "set_env, env_value, expected", + [ + (False, None, "float"), + (True, "base64", "base64"), + ], +) +def test_openai_embedding_encoding_format_default( + monkeypatch, set_env, env_value, expected +): + monkeypatch.delenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", raising=False) + if set_env: + monkeypatch.setenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", env_value) + + mock_response = MagicMock() + mock_response.parse.return_value = MagicMock( + model_dump=lambda: { + "data": [{"embedding": [0.1, 0.2, 0.3], "index": 0}], + "model": "text-embedding-ada-002", + "object": "list", + "usage": {"prompt_tokens": 1, "total_tokens": 1}, + } + ) + mock_response.headers = {} + + with patch( + "litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client" + ) as mock_get_client: + mock_client_instance = MagicMock() + mock_get_client.return_value = mock_client_instance + mock_client_instance.embeddings.with_raw_response.create.return_value = ( + mock_response + ) + + embedding( + model="text-embedding-ada-002", + input="Hello world", + ) + + call_kwargs = ( + mock_client_instance.embeddings.with_raw_response.create.call_args[1] + ) + assert call_kwargs["encoding_format"] == expected + + +@pytest.mark.parametrize("env_none", ["none", "NONE", " none "]) +def test_openai_embedding_encoding_format_env_none_omits_param( + monkeypatch, env_none +): + """LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT=none omits encoding_format (provider default).""" + monkeypatch.setenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", env_none) + + mock_response = MagicMock() + mock_response.parse.return_value = MagicMock( + model_dump=lambda: { + "data": [{"embedding": [0.1, 0.2, 0.3], "index": 0}], + "model": "text-embedding-ada-002", + "object": "list", + "usage": {"prompt_tokens": 1, "total_tokens": 1}, + } + ) + mock_response.headers = {} + + with patch( + "litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client" + ) as mock_get_client: + mock_client_instance = MagicMock() + mock_get_client.return_value = mock_client_instance + mock_client_instance.embeddings.with_raw_response.create.return_value = ( + mock_response + ) + + embedding( + model="text-embedding-ada-002", + input="Hello world", + ) + + call_kwargs = ( + mock_client_instance.embeddings.with_raw_response.create.call_args[1] + ) + assert "encoding_format" not in call_kwargs + + +def test_openai_embedding_encoding_format_explicit_overrides_env(monkeypatch): + """Request `encoding_format` wins over LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT.""" + monkeypatch.setenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", "float") + + mock_response = MagicMock() + mock_response.parse.return_value = MagicMock( + model_dump=lambda: { + "data": [{"embedding": [0.1, 0.2, 0.3], "index": 0}], + "model": "text-embedding-ada-002", + "object": "list", + "usage": {"prompt_tokens": 1, "total_tokens": 1}, + } + ) + mock_response.headers = {} + + with patch( + "litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client" + ) as mock_get_client: + mock_client_instance = MagicMock() + mock_get_client.return_value = mock_client_instance + mock_client_instance.embeddings.with_raw_response.create.return_value = ( + mock_response + ) + + embedding( + model="text-embedding-ada-002", + input="Hello world", + encoding_format="base64", + ) + + call_kwargs = ( + mock_client_instance.embeddings.with_raw_response.create.call_args[1] + ) + assert call_kwargs["encoding_format"] == "base64" diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/projects/useProjects.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/projects/useProjects.ts index 85c8b25645c..79976f54626 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/projects/useProjects.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/projects/useProjects.ts @@ -6,8 +6,8 @@ import { deriveErrorMessage, handleError, } from "@/components/networking"; -import { all_admin_roles } from "@/utils/roles"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { all_admin_roles } from "@/utils/roles"; // ── Types ──────────────────────────────────────────────────────────────────── @@ -81,7 +81,6 @@ export const useProjects = () => { return useQuery({ queryKey: projectKeys.list({}), queryFn: async () => fetchProjects(accessToken!), - enabled: - Boolean(accessToken) && all_admin_roles.includes(userRole || ""), + enabled: Boolean(accessToken) && all_admin_roles.includes(userRole!), }); }; diff --git a/ui/litellm-dashboard/src/components/UsagePage/components/UsagePageView.tsx b/ui/litellm-dashboard/src/components/UsagePage/components/UsagePageView.tsx index 69b29564d83..809f1d4e17b 100644 --- a/ui/litellm-dashboard/src/components/UsagePage/components/UsagePageView.tsx +++ b/ui/litellm-dashboard/src/components/UsagePage/components/UsagePageView.tsx @@ -169,8 +169,8 @@ const UsagePage: React.FC = ({ teams, organizations }) => { } }, [isAdmin, userID]); - // For non-admins, always pass their own user_id - const effectiveUserId = isAdmin ? selectedUserId : userID || null; + // For non-admins or "my-usage" view, always pass their own user_id + const effectiveUserId = usageView === "my-usage" || !isAdmin ? userID || null : selectedUserId; const startTime = useMemo(() => (dateValue.from ? new Date(dateValue.from) : null), [dateValue.from]); const endTime = useMemo(() => (dateValue.to ? new Date(dateValue.to) : null), [dateValue.to]); @@ -477,10 +477,10 @@ const UsagePage: React.FC = ({ teams, organizations }) => { } /> )} - {/* Your Usage Panel */} - {usageView === "global" && ( + {/* Your Usage / Global Usage Panel */} + {(usageView === "global" || usageView === "my-usage") && ( <> - {isAdmin && ( + {isAdmin && usageView === "global" && (
Filter by user