diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index fa7b73f6c45..544ace9063a 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -914,6 +914,7 @@ router_settings: | MODEL_COST_MAP_MAX_SHRINK_RATIO | Maximum allowed shrinkage ratio when validating a fetched model cost map against the local backup. Rejects the fetched map if it is smaller than this fraction of the backup. Default is 0.5 | MODEL_COST_MAP_MIN_MODEL_COUNT | Minimum number of models a fetched cost map must contain to be considered valid. Default is 50 | NO_DOCS | Flag to disable Swagger UI documentation +| NO_OPENAPI | Flag to disable the /openapi.json endpoint | NO_REDOC | Flag to disable Redoc documentation | NO_PROXY | List of addresses to bypass proxy | NON_LLM_CONNECTION_TIMEOUT | Timeout in seconds for non-LLM service connections. Default is 15 diff --git a/docs/my-website/docs/proxy/guardrails/hiddenlayer.md b/docs/my-website/docs/proxy/guardrails/hiddenlayer.md index 1ec892972d0..2aab139cd24 100644 --- a/docs/my-website/docs/proxy/guardrails/hiddenlayer.md +++ b/docs/my-website/docs/proxy/guardrails/hiddenlayer.md @@ -174,6 +174,7 @@ guardrails: - **`default_on`**: Automatically attach the guardrail to every request unless the client opts out. - **`hl-project-id` header**: Routes scans to a specific HiddenLayer project. - **`hl-requester-id` header**: Sets `metadata.requester_id` for auditing. +- **`hl-session-id` header**: Groups related requests into a session for contextual analysis and tracing in the HiddenLayer console. ## Environment variables diff --git a/litellm/caching/in_memory_cache.py b/litellm/caching/in_memory_cache.py index 5239fa1f4b0..ba446dd4f60 100644 --- a/litellm/caching/in_memory_cache.py +++ b/litellm/caching/in_memory_cache.py @@ -161,9 +161,10 @@ class InMemoryCache(BaseCache): if self.max_size_in_memory == 0: return # Don't cache anything if max size is 0 - if len(self.cache_dict) >= self.max_size_in_memory: - # only evict when cache is full - self.evict_cache() + # Always prune expired/outdated heap roots before inserting. + # This keeps expiration_heap bounded even when the live cache stays + # below max_size_in_memory and keys are reinserted after TTL expiry. + self.evict_cache() if not self.check_value_size(value): return diff --git a/litellm/integrations/datadog/datadog.py b/litellm/integrations/datadog/datadog.py index 4de3644b581..c3e555f6e89 100644 --- a/litellm/integrations/datadog/datadog.py +++ b/litellm/integrations/datadog/datadog.py @@ -16,6 +16,7 @@ For batching specific details see CustomBatchLogger class import asyncio import datetime import os +import time import traceback from datetime import datetime as datetimeObj from typing import Any, Dict, List, Optional, Union @@ -301,7 +302,7 @@ class DataDogLogger( self.log_queue.append(dd_payload) if len(self.log_queue) >= self.batch_size: - await self.async_send_batch() + await self.flush_queue() except Exception as e: verbose_logger.exception( f"Datadog: async_post_call_failure_hook - {str(e)}\n{traceback.format_exc()}" @@ -324,9 +325,12 @@ class DataDogLogger( verbose_logger.exception("Datadog: log_queue does not exist") return + batch_to_send = self.log_queue[:] + self.log_queue = [] + verbose_logger.debug( "Datadog - about to flush %s events on %s", - len(self.log_queue), + len(batch_to_send), self.intake_url, ) @@ -335,9 +339,10 @@ class DataDogLogger( "[DATADOG MOCK] Mock mode enabled - API calls will be intercepted" ) - response = await self.async_send_compressed_data(self.log_queue) + response = await self.async_send_compressed_data(batch_to_send) if response.status_code == 413: verbose_logger.exception(DD_ERRORS.DATADOG_413_ERROR.value) + self.log_queue = batch_to_send + self.log_queue return response.raise_for_status() @@ -348,7 +353,7 @@ class DataDogLogger( if self.is_mock_mode: verbose_logger.debug( - f"[DATADOG MOCK] Batch of {len(self.log_queue)} events successfully mocked" + f"[DATADOG MOCK] Batch of {len(batch_to_send)} events successfully mocked" ) else: verbose_logger.debug( @@ -356,11 +361,26 @@ class DataDogLogger( response.status_code, response.text, ) + except Exception as e: + self.log_queue = batch_to_send + self.log_queue verbose_logger.exception( f"Datadog Error sending batch API - {str(e)}\n{traceback.format_exc()}" ) + async def flush_queue(self): + if self.flush_lock is None: + return + + async with self.flush_lock: + if self.log_queue: + verbose_logger.debug( + "Datadog: Flushing batch of %s events", len(self.log_queue) + ) + await self.async_send_batch() + if not self.log_queue: + self.last_flush_time = time.time() + def log_success_event(self, kwargs, response_obj, start_time, end_time): """ Sync Log success events to Datadog @@ -429,7 +449,7 @@ class DataDogLogger( ) if len(self.log_queue) >= self.batch_size: - await self.async_send_batch() + await self.flush_queue() def _create_datadog_logging_payload_helper( self, diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py index 6bddad09f21..799e8ab9a0a 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py @@ -129,14 +129,22 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): if should_start_new_block and not self.sent_content_block_finish: # Queue the sequence: content_block_stop -> content_block_start - # The trigger chunk itself is not emitted as a delta since the - # content_block_start already carries the relevant information. + # For text blocks the trigger chunk is not emitted as a separate + # delta because content_block_start carries the information. + # For tool_use blocks we must also emit the trigger chunk's delta + # when it carries input_json_delta data, because some providers + # (e.g. xAI, Gemini) include tool arguments in the same streaming + # chunk as the function name/id. + + # 1. Stop current content block self.chunk_queue.append( { "type": "content_block_stop", "index": max(self.current_content_block_index - 1, 0), } ) + + # 2. Start new content block self.chunk_queue.append( { "type": "content_block_start", @@ -144,6 +152,17 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): "content_block": self.current_content_block_start, } ) + + # 3. If the trigger chunk carries tool argument data, queue it + # so the input_json_delta is not silently dropped. + if ( + processed_chunk.get("type") == "content_block_delta" + and isinstance(processed_chunk.get("delta"), dict) + and processed_chunk["delta"].get("type") == "input_json_delta" + and processed_chunk["delta"].get("partial_json") + ): + self.chunk_queue.append(processed_chunk) + self.sent_content_block_finish = False return self.chunk_queue.popleft() @@ -282,16 +301,16 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): hasattr(chunk.usage, "_cache_creation_input_tokens") and chunk.usage._cache_creation_input_tokens > 0 ): - usage_dict[ - "cache_creation_input_tokens" - ] = chunk.usage._cache_creation_input_tokens + usage_dict["cache_creation_input_tokens"] = ( + chunk.usage._cache_creation_input_tokens + ) if ( hasattr(chunk.usage, "_cache_read_input_tokens") and chunk.usage._cache_read_input_tokens > 0 ): - usage_dict[ - "cache_read_input_tokens" - ] = chunk.usage._cache_read_input_tokens + usage_dict["cache_read_input_tokens"] = ( + chunk.usage._cache_read_input_tokens + ) merged_chunk["usage"] = usage_dict # Queue the merged chunk and reset @@ -305,8 +324,12 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): if not self.queued_usage_chunk: if should_start_new_block and not self.sent_content_block_finish: # Queue the sequence: content_block_stop -> content_block_start - # The trigger chunk itself is not emitted as a delta since the - # content_block_start already carries the relevant information. + # For text blocks the trigger chunk is not emitted as a separate + # delta because content_block_start carries the information. + # For tool_use blocks we must also emit the trigger chunk's delta + # when it carries input_json_delta data, because some providers + # (e.g. xAI, Gemini) include tool arguments in the same streaming + # chunk as the function name/id. # 1. Stop current content block self.chunk_queue.append( @@ -325,6 +348,17 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): } ) + # 3. If the trigger chunk carries tool argument data, queue it + # so the input_json_delta is not silently dropped. + if ( + processed_chunk.get("type") == "content_block_delta" + and isinstance(processed_chunk.get("delta"), dict) + and processed_chunk["delta"].get("type") + == "input_json_delta" + and processed_chunk["delta"].get("partial_json") + ): + self.chunk_queue.append(processed_chunk) + # Reset state for new block self.sent_content_block_finish = False diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index e6e548ab98a..cd27b4c362a 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -480,6 +480,62 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): else: return None + @staticmethod + def _resolve_search_tool_conflict( + gtool_func_declarations: list, + googleSearch: Optional[dict], + googleSearchRetrieval: Optional[dict], + enterpriseWebSearch: Optional[dict], + urlContext: Optional[dict], + optional_params: dict, + ) -> tuple: + """ + Resolve Vertex AI constraint: multiple Tool objects in a request must + ALL be search tools. When function declarations are mixed with search + tools, drop search tools to avoid 400 error. + + Skip when include_server_side_tool_invocations is enabled (Gemini 3+ + supports tool combination natively). + + Note: code_execution, computerUse, and googleMaps are NOT search tools + and CAN coexist with function declarations, so they are preserved. + + Ref: https://github.com/BerriAI/litellm/issues/23337 + + Returns: + tuple of (googleSearch, googleSearchRetrieval, enterpriseWebSearch, urlContext) + """ + has_search_tools = any( + v is not None + for v in [ + googleSearch, + googleSearchRetrieval, + enterpriseWebSearch, + urlContext, + ] + ) + server_side_tool_invocations = optional_params.get( + "include_server_side_tool_invocations", False + ) + if ( + gtool_func_declarations + and has_search_tools + and not server_side_tool_invocations + ): + verbose_logger.warning( + "Vertex AI does not support mixing function declarations with " + "search tools (googleSearch, enterpriseWebSearch, urlContext, " + "googleSearchRetrieval) in the same request. Dropping search " + "tools and keeping function declarations. To use search tools, " + "send a request without function calling tools." + ) + googleSearch = None + googleSearchRetrieval = None + enterpriseWebSearch = None + urlContext = None + + return googleSearch, googleSearchRetrieval, enterpriseWebSearch, urlContext + def _map_function( # noqa: PLR0915 self, value: List[dict], optional_params: dict ) -> List[Tools]: @@ -512,9 +568,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): value = _remove_strict_from_schema(value) for tool in value: - openai_function_object: Optional[ - ChatCompletionToolParamFunctionChunk - ] = None + openai_function_object: Optional[ChatCompletionToolParamFunctionChunk] = ( + None + ) if "function" in tool: # tools list _openai_function_object = ChatCompletionToolParamFunctionChunk( # type: ignore **tool["function"] @@ -633,6 +689,20 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): # per Vertex AI API spec: "A Tool object should contain exactly one type of Tool" _tools_list: List[Tools] = [] + ( + googleSearch, + googleSearchRetrieval, + enterpriseWebSearch, + urlContext, + ) = self._resolve_search_tool_conflict( + gtool_func_declarations=gtool_func_declarations, + googleSearch=googleSearch, + googleSearchRetrieval=googleSearchRetrieval, + enterpriseWebSearch=enterpriseWebSearch, + urlContext=urlContext, + optional_params=optional_params, + ) + # Function declarations can be grouped together in one Tool if gtool_func_declarations: func_tool = Tools() @@ -646,15 +716,15 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): _tools_list.append(search_tool) if googleSearchRetrieval is not None: retrieval_tool = Tools() - retrieval_tool[ - VertexToolName.GOOGLE_SEARCH_RETRIEVAL.value - ] = googleSearchRetrieval + retrieval_tool[VertexToolName.GOOGLE_SEARCH_RETRIEVAL.value] = ( + googleSearchRetrieval + ) _tools_list.append(retrieval_tool) if enterpriseWebSearch is not None: enterprise_tool = Tools() - enterprise_tool[ - VertexToolName.ENTERPRISE_WEB_SEARCH.value - ] = enterpriseWebSearch + enterprise_tool[VertexToolName.ENTERPRISE_WEB_SEARCH.value] = ( + enterpriseWebSearch + ) _tools_list.append(enterprise_tool) if code_execution is not None: code_tool = Tools() @@ -1101,16 +1171,16 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): param_description="thinking_budget", ) if VertexGeminiConfig._is_gemini_3_or_newer(model): - optional_params[ - "thinkingConfig" - ] = VertexGeminiConfig._map_reasoning_effort_to_thinking_level( - effort_value, model + optional_params["thinkingConfig"] = ( + VertexGeminiConfig._map_reasoning_effort_to_thinking_level( + effort_value, model + ) ) else: - optional_params[ - "thinkingConfig" - ] = VertexGeminiConfig._map_reasoning_effort_to_thinking_budget( - effort_value, model + optional_params["thinkingConfig"] = ( + VertexGeminiConfig._map_reasoning_effort_to_thinking_budget( + effort_value, model + ) ) elif param == "thinking": # Validate no conflict with thinking_level @@ -1119,11 +1189,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): param_name="thinking", param_description="thinking_budget", ) - optional_params[ - "thinkingConfig" - ] = VertexGeminiConfig._map_thinking_param( - cast(AnthropicThinkingParam, value), - model=model, + optional_params["thinkingConfig"] = ( + VertexGeminiConfig._map_thinking_param( + cast(AnthropicThinkingParam, value), + model=model, + ) ) elif param == "modalities" and isinstance(value, list): response_modalities = self.map_response_modalities(value) @@ -1547,10 +1617,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): _tool_response_chunk["provider_specific_fields"] = { # type: ignore "thought_signature": thought_signature } - _tool_response_chunk[ - "id" - ] = _encode_tool_call_id_with_signature( - _tool_response_chunk["id"] or "", thought_signature + _tool_response_chunk["id"] = ( + _encode_tool_call_id_with_signature( + _tool_response_chunk["id"] or "", thought_signature + ) ) _tools.append(_tool_response_chunk) cumulative_tool_call_idx += 1 @@ -2397,28 +2467,28 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ## ADD METADATA TO RESPONSE ## setattr(model_response, "vertex_ai_grounding_metadata", grounding_metadata) - model_response._hidden_params[ - "vertex_ai_grounding_metadata" - ] = grounding_metadata + model_response._hidden_params["vertex_ai_grounding_metadata"] = ( + grounding_metadata + ) setattr( model_response, "vertex_ai_url_context_metadata", url_context_metadata ) - model_response._hidden_params[ - "vertex_ai_url_context_metadata" - ] = url_context_metadata + model_response._hidden_params["vertex_ai_url_context_metadata"] = ( + url_context_metadata + ) setattr(model_response, "vertex_ai_safety_results", safety_ratings) - model_response._hidden_params[ - "vertex_ai_safety_results" - ] = safety_ratings # older approach - maintaining to prevent regressions + model_response._hidden_params["vertex_ai_safety_results"] = ( + safety_ratings # older approach - maintaining to prevent regressions + ) ## ADD CITATION METADATA ## setattr(model_response, "vertex_ai_citation_metadata", citation_metadata) - model_response._hidden_params[ - "vertex_ai_citation_metadata" - ] = citation_metadata # older approach - maintaining to prevent regressions + model_response._hidden_params["vertex_ai_citation_metadata"] = ( + citation_metadata # older approach - maintaining to prevent regressions + ) ## ADD TRAFFIC TYPE ## traffic_type = completion_response.get("usageMetadata", {}).get( @@ -3126,7 +3196,12 @@ class ModelResponseIterator: setattr(model_response, "vertex_ai_safety_ratings", safety_ratings) # type: ignore setattr(model_response, "vertex_ai_citation_metadata", citation_metadata) # type: ignore - return grounding_metadata, url_context_metadata, safety_ratings, citation_metadata + return ( + grounding_metadata, + url_context_metadata, + safety_ratings, + citation_metadata, + ) def _apply_stream_usage_metadata( self, @@ -3151,9 +3226,9 @@ class ModelResponseIterator: traffic_type = processed_chunk.get("usageMetadata", {}).get("trafficType") if traffic_type: - model_response._hidden_params.setdefault( - "provider_specific_fields", {} - )["traffic_type"] = traffic_type + model_response._hidden_params.setdefault("provider_specific_fields", {})[ + "traffic_type" + ] = traffic_type service_tier = self.response_headers.get("x-gemini-service-tier") if service_tier: 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 08831a8215f..389a3a85f56 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 @@ -292,10 +292,10 @@ def process_response( _predictions: VertexAIBatchEmbeddingsResponseObject, ) -> EmbeddingResponse: openai_embeddings: List[Embedding] = [] - for embedding in _predictions["embeddings"]: + for idx, embedding in enumerate(_predictions["embeddings"]): openai_embedding = Embedding( embedding=embedding["values"], - index=0, + index=idx, object="embedding", ) openai_embeddings.append(openai_embedding) diff --git a/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/__init__.py index 065ba2e12d0..cd71d55991e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/__init__.py @@ -2,7 +2,7 @@ from typing import TYPE_CHECKING from litellm.types.guardrails import SupportedGuardrailIntegrations -from .hiddenlayer import HiddenlayerGuardrail +from .hiddenlayer import HiddenlayerGuardrail, HiddenlayerGuardrailV2 if TYPE_CHECKING: from litellm.types.guardrails import Guardrail, LitellmParams @@ -13,17 +13,32 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" api_id = litellm_params.api_id if hasattr(litellm_params, "api_id") else None auth_url = litellm_params.auth_url if hasattr(litellm_params, "auth_url") else None - - _hiddenlayer_callback = HiddenlayerGuardrail( - api_base=litellm_params.api_base, - api_id=api_id, - api_key=litellm_params.api_key, - auth_url=auth_url, - guardrail_name=guardrail.get("guardrail_name", ""), - event_hook=litellm_params.mode, - default_on=litellm_params.default_on, + version: int | None = ( + litellm_params.version if hasattr(litellm_params, "version") else None ) + _hiddenlayer_callback: HiddenlayerGuardrail | HiddenlayerGuardrailV2 + if not version or version < 2: + _hiddenlayer_callback = HiddenlayerGuardrail( + api_base=litellm_params.api_base, + api_id=api_id, + api_key=litellm_params.api_key, + auth_url=auth_url, + guardrail_name=guardrail.get("guardrail_name", ""), + event_hook=litellm_params.mode, + default_on=litellm_params.default_on, + ) + else: + _hiddenlayer_callback = HiddenlayerGuardrailV2( + api_base=litellm_params.api_base, + api_id=api_id, + api_key=litellm_params.api_key, + auth_url=auth_url, + guardrail_name=guardrail.get("guardrail_name", ""), + event_hook=litellm_params.mode, + default_on=litellm_params.default_on, + ) + litellm.logging_callback_manager.add_litellm_callback(_hiddenlayer_callback) return _hiddenlayer_callback diff --git a/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/hiddenlayer.py b/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/hiddenlayer.py index b907fbbcbda..091187983a2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/hiddenlayer.py +++ b/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/hiddenlayer.py @@ -1,4 +1,6 @@ from __future__ import annotations +from uuid import uuid4 +import httpx import os from typing import TYPE_CHECKING, Any, Literal, Optional, Type @@ -151,14 +153,19 @@ class HiddenlayerGuardrail(CustomGuardrail): project_id = headers.get("hl-project-id") if scan_params := inputs.get("structured_messages"): - # Convert AllMessageValues to simple dict format for HiddenLayer API - messages = [ - {"role": msg.get("role", "user"), "content": msg.get("content", "")} - for msg in scan_params - if isinstance(msg, dict) - ] + last_msg = scan_params[-1] result = await self._call_hiddenlayer( - project_id, hl_request_metadata, {"messages": messages}, input_type + project_id, + hl_request_metadata, + { + "messages": [ + { + "role": last_msg.get("role", "user"), + "content": str(last_msg.get("content", "")), + } + ] + }, + input_type, ) elif text := inputs.get("texts"): result = await self._call_hiddenlayer( @@ -171,22 +178,48 @@ class HiddenlayerGuardrail(CustomGuardrail): result = {} if result.get("evaluation", {}).get("action") == HiddenlayerAction.BLOCK: + detected_reasons = [ + entry.get("name", "unknown") + for entry in result.get("analysis", []) + if entry.get("detected") + ] + threat_level = result.get("evaluation", {}).get("threat_level") raise HTTPException( status_code=400, detail={ "error": "Violated guardrail policy", - "hiddenlayer_guardrail_response": HiddenlayerMessages.BLOCK_MESSAGE, + "hiddenlayer_guardrail_response": HiddenlayerMessages.BLOCK_MESSAGE.value, + "block_reasons": detected_reasons, + "threat_level": threat_level, }, ) if result.get("evaluation", {}).get("action") == HiddenlayerAction.REDACT: modified_data = result.get("modified_data", {}) if modified_data.get("input") and input_type == "request": - inputs["texts"] = [modified_data["input"]["messages"][-1]["content"]] + last_content = modified_data["input"]["messages"][-1]["content"] + if isinstance(last_content, list): + texts = [ + item["text"] + for item in last_content + if isinstance(item, dict) and item.get("type") == "text" + ] + inputs["texts"] = texts if texts else [""] + else: + inputs["texts"] = [last_content] inputs["structured_messages"] = modified_data["input"]["messages"] if modified_data.get("output") and input_type == "response": - inputs["texts"] = [modified_data["output"]["messages"][-1]["content"]] + last_content = modified_data["output"]["messages"][-1]["content"] + if isinstance(last_content, list): + texts = [ + item["text"] + for item in last_content + if isinstance(item, dict) and item.get("type") == "text" + ] + inputs["texts"] = texts if texts else [""] + else: + inputs["texts"] = [last_content] return inputs @@ -206,6 +239,8 @@ class HiddenlayerGuardrail(CustomGuardrail): headers = { "Content-Type": "application/json", + "hl-runtime-edge-provider": "litellm", + "hl-runtime-edge-provider-version": "1", } if project_id: @@ -257,3 +292,229 @@ class HiddenlayerGuardrail(CustomGuardrail): ) return HiddenlayerGuardrailConfigModel + + +class HiddenlayerGuardrailV2(CustomGuardrail): + """Custom guardrail wrapper for HiddenLayer's safety checks.""" + + def __init__( + self, + api_id: Optional[str] = None, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + auth_url: Optional[str] = None, + **kwargs: Any, + ) -> None: + self.hiddenlayer_client_id = api_id or os.getenv("HIDDENLAYER_CLIENT_ID") + self.hiddenlayer_client_secret = api_key or os.getenv( + "HIDDENLAYER_CLIENT_SECRET" + ) + self.api_base = ( + api_base + or os.getenv("HIDDENLAYER_API_BASE") + or "https://api.hiddenlayer.ai" + ) + self.jwt_token = None + + auth_url = ( + auth_url + or os.getenv("HIDDENLAYER_AUTH_URL") + or "https://auth.hiddenlayer.ai" + ) + + if is_saas(self.api_base): + if not self.hiddenlayer_client_id: + raise RuntimeError( + "`api_id` cannot be None when using the SaaS version of HiddenLayer." + ) + + if not self.hiddenlayer_client_secret: + raise RuntimeError( + "`api_key` cannot be None when using the SaaS version of HiddenLayer." + ) + + self.jwt_token = _get_jwt( + auth_url=auth_url, + api_id=self.hiddenlayer_client_id, + api_key=self.hiddenlayer_client_secret, + ) + self.refresh_jwt_func = lambda: _get_jwt( + auth_url=auth_url, + api_id=self.hiddenlayer_client_id, + api_key=self.hiddenlayer_client_secret, + ) + + self._http_client = get_async_httpx_client( + llm_provider=httpxSpecialProvider.GuardrailCallback + ) + super().__init__(**kwargs) + + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional["LiteLLMLoggingObj"] = None, + ) -> GenericGuardrailAPIInputs: + """Validate (and optionally redact) text via HiddenLayer before/after LLM calls.""" + + # We need the hiddenlayer project id and requester id on both the input and output + # Since headers aren't available on the response back from the model, we get them + # from the logging object. It ends up working out that on the request, we parse the + # hiddenlayer params from the raw request and then retrieve those same headers + # from the logger object on the response from the model. + headers = request_data.get("proxy_server_request", {}).get("headers", {}) + if not headers and logging_obj and logging_obj.model_call_details: + headers = ( + logging_obj.model_call_details.get("litellm_params", {}) + .get("metadata", {}) + .get("headers", {}) + ) + + # put our roundtrip id in the header to the model so we get it on the way back from the model + if "hl-roundtrip-id" not in headers: + proxy_req = request_data.get("proxy_server_request") + if proxy_req is not None and "headers" in proxy_req: + proxy_req["headers"]["hl-roundtrip-id"] = str(uuid4()) + headers["hl-roundtrip-id"] = proxy_req["headers"]["hl-roundtrip-id"] + + hl_headers = { + h.lower(): v for h, v in headers.items() if h.lower().startswith("hl-") + } + + if "hl-requester-id" not in hl_headers: + hl_headers["hl-requester-id"] = "LiteLLM" + + payload: Any + if input_type == "request": + payload = { + "messages": inputs.get("structured_messages"), + "model": inputs.get("model"), + "tools": inputs.get("tools"), + } + else: + if inputs.get("texts"): + payload = { + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": inputs["texts"][0] + if inputs.get("texts") + else "", + }, + "finish_reason": "stop", + } + ] + } + elif tool_calls := inputs.get("tool_calls"): + payload = tool_calls + else: + payload = {} + + response = await self._call_hiddenlayer( + payload, input_type, hl_headers + ) + output = response.json() + + if response.headers.get("hl-runtime-action", "").lower() == "block": + raise HTTPException( + status_code=400, + detail={ + "error": "Violated guardrail policy", + "hiddenlayer_guardrail_response": HiddenlayerMessages.BLOCK_MESSAGE.value, + }, + ) + + new_texts = [] + if input_type == "request": + inputs["structured_messages"] = output + + for message in output.get("messages", []): + content = message.get("content", "") + if isinstance(content, list): + text_parts = [ + item["text"] + for item in content + if isinstance(item, dict) and item.get("type") == "text" + ] + if text_parts: + new_texts.append(" ".join(text_parts)) + elif content: + new_texts.append(content) + + inputs["texts"] = new_texts + + elif input_type == "response" and inputs.get("texts"): + inputs["texts"] = [ + output.get("choices", [{}])[-1].get("message", {}).get("content", "") + ] + elif input_type == "response" and inputs.get("tool_calls"): + inputs["tool_calls"] = output + + return inputs + + async def _call_hiddenlayer( + self, + payload: Any, + input_type: Literal["request", "response"], + hl_headers: dict[str, str], + ) -> httpx.Response: + if input_type == "request": + path = "detection/v2/request-evaluations" + else: + path = "detection/v2/response-evaluations" + + headers = { + "Content-Type": "application/json", + "hl-runtime-edge-provider": "litellm", + "hl-runtime-edge-provider-version": "2", + } + if self.jwt_token: + headers["Authorization"] = f"Bearer {self.jwt_token}" + + headers.update(hl_headers) + + try: + response = await self._http_client.post( + f"{self.api_base}/{path}", + json=payload, + headers=headers, + ) + response.raise_for_status() + + verbose_proxy_logger.debug(f"Hiddenlayer reponse: {response}") + + return response + except HTTPStatusError as e: + # Try the request again by refreshing the jwt if we get 401 + # since the Hiddenlayer jwt timeout is an hour and this is + # a long lived session application + if e.response.status_code == 401 and self.jwt_token is not None: + verbose_proxy_logger.debug( + "Unable to authenticate to Hiddenlayer, JWT token is invalid or expired, trying to refresh the token." + ) + self.jwt_token = self.refresh_jwt_func() + headers["Authorization"] = f"Bearer {self.jwt_token}" + response = await self._http_client.post( + f"{self.api_base}/{path}", + json=payload, + headers=headers, + ) + else: + raise e + + response.raise_for_status() + + verbose_proxy_logger.debug(f"Hiddenlayer reponse: {response}") + return response + + @staticmethod + def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + from litellm.types.proxy.guardrails.guardrail_hooks.hiddenlayer import ( + HiddenlayerGuardrailConfigModel, + ) + + return HiddenlayerGuardrailConfigModel diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index fcc224e848a..138469312e1 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -112,6 +112,14 @@ from litellm.types.proxy.management_endpoints.team_endpoints import ( router = APIRouter() +def _sanitize_for_log(value: Any) -> str: + """Strip CR/LF from user-controlled values to prevent log injection.""" + try: + text = str(value) + except Exception: + text = repr(value) + return text.replace("\r", "").replace("\n", "") + async def _verify_team_access( team_obj: LiteLLM_TeamTable, user_api_key_dict: UserAPIKeyAuth, @@ -285,6 +293,61 @@ class TeamMemberBudgetHandler: data_dict.pop("team_member_rpm_limit", None) data_dict.pop("team_member_tpm_limit", None) + @staticmethod + async def backfill_team_member_budget_entries( + team_id: str, + members_with_roles: List[Union[Member, dict]], + team_member_budget_id: str, + prisma_client: PrismaClient, + ) -> None: + """ + Create team_memberships entries for existing members that don't have one. + + Called after team_member_budget is set/updated on a team to ensure + members who joined before the budget was configured also get budget + enforcement. + + Only creates missing entries — does not touch existing memberships + (which may carry individual per-member budgets). + """ + if not members_with_roles: + return + + # Batch-fetch existing memberships for this team (avoids N+1 queries) + existing_memberships = ( + await prisma_client.db.litellm_teammembership.find_many( + where={"team_id": team_id} + ) + ) + existing_user_ids = {m.user_id for m in existing_memberships} + + # Identify members with no existing membership row. + # members_with_roles may contain Member instances or raw dicts depending + # on how the team was fetched/deserialized. + missing = [] + for m in members_with_roles: + user_id = m.get("user_id") if isinstance(m, dict) else m.user_id + if user_id is not None and user_id not in existing_user_ids: + missing.append( + { + "team_id": team_id, + "user_id": user_id, + "budget_id": team_member_budget_id, + } + ) + + if missing: + await prisma_client.db.litellm_teammembership.create_many( + data=missing, + skip_duplicates=True, # safety net against concurrent races + ) + verbose_proxy_logger.info( + "Backfilled %d team_memberships for team %s with budget %s", + len(missing), + _sanitize_for_log(team_id), + _sanitize_for_log(team_member_budget_id), + ) + def _get_default_team_param(field: str) -> Any: """ @@ -1551,6 +1614,18 @@ async def update_team( # noqa: PLR0915 team_member_tpm_limit=data.team_member_tpm_limit, team_member_budget_duration=data.team_member_budget_duration, ) + # Backfill team_memberships for members who joined before the + # budget was configured — they won't have a membership row yet. + _backfill_budget_id = (updated_kv.get("metadata") or {}).get( + "team_member_budget_id" + ) + if _backfill_budget_id and existing_team_row.members_with_roles: + await TeamMemberBudgetHandler.backfill_team_member_budget_entries( + team_id=data.team_id, + members_with_roles=existing_team_row.members_with_roles, + team_member_budget_id=_backfill_budget_id, + prisma_client=prisma_client, + ) else: TeamMemberBudgetHandler._clean_team_member_fields(updated_kv) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index db9468356dd..d0ccfb40dba 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -493,6 +493,7 @@ from litellm.proxy.utils import ( ProxyUpdateSpend, _cache_user_row, _get_docs_url, + _get_openapi_url, _get_projected_spend_over_limit, _get_redoc_url, _is_projected_spend_over_limit, @@ -1000,6 +1001,7 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915 app = FastAPI( docs_url=_get_docs_url(), redoc_url=_get_redoc_url(), + openapi_url=_get_openapi_url(), title=_title, description=_description, version=version, diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index e15b48577de..a6f81986a6f 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -5321,6 +5321,19 @@ def get_error_message_str(e: Exception) -> str: return error_message +def _get_openapi_url() -> Optional[str]: + """ + Get the OpenAPI schema URL from the environment variables. + + - If NO_OPENAPI is True, return None. + - Otherwise, default to "/openapi.json". + """ + if str_to_bool(os.getenv("NO_OPENAPI")) is True: + return None + + return "/openapi.json" + + def _get_redoc_url() -> Optional[str]: """ Get the Redoc URL from the environment variables. diff --git a/litellm/router_strategy/lowest_latency.py b/litellm/router_strategy/lowest_latency.py index 20db28fa10e..870b3f29d48 100644 --- a/litellm/router_strategy/lowest_latency.py +++ b/litellm/router_strategy/lowest_latency.py @@ -143,7 +143,7 @@ class LowestLatencyLoggingHandler(CustomLogger): else: request_count_dict[id]["latency"] = request_count_dict[id][ "latency" - ][: self.routing_args.max_latency_list_size - 1] + [final_value] + ][1:] + [final_value] ## Time to first token if time_to_first_token is not None: @@ -155,13 +155,10 @@ class LowestLatencyLoggingHandler(CustomLogger): "time_to_first_token", [] ).append(time_to_first_token) else: - request_count_dict[id][ - "time_to_first_token" - ] = request_count_dict[id]["time_to_first_token"][ - : self.routing_args.max_latency_list_size - 1 - ] + [ - time_to_first_token - ] + request_count_dict[id]["time_to_first_token"] = ( + request_count_dict[id]["time_to_first_token"][1:] + + [time_to_first_token] + ) if precise_minute not in request_count_dict[id]: request_count_dict[id][precise_minute] = {} @@ -244,7 +241,7 @@ class LowestLatencyLoggingHandler(CustomLogger): else: request_count_dict[id]["latency"] = request_count_dict[id][ "latency" - ][: self.routing_args.max_latency_list_size - 1] + [1000.0] + ][1:] + [1000.0] await self.router_cache.async_set_cache( key=latency_key, @@ -371,7 +368,7 @@ class LowestLatencyLoggingHandler(CustomLogger): else: request_count_dict[id]["latency"] = request_count_dict[id][ "latency" - ][: self.routing_args.max_latency_list_size - 1] + [final_value] + ][1:] + [final_value] ## Time to first token if time_to_first_token is not None: @@ -383,13 +380,10 @@ class LowestLatencyLoggingHandler(CustomLogger): "time_to_first_token", [] ).append(time_to_first_token) else: - request_count_dict[id][ - "time_to_first_token" - ] = request_count_dict[id]["time_to_first_token"][ - : self.routing_args.max_latency_list_size - 1 - ] + [ - time_to_first_token - ] + request_count_dict[id]["time_to_first_token"] = ( + request_count_dict[id]["time_to_first_token"][1:] + + [time_to_first_token] + ) if precise_minute not in request_count_dict[id]: request_count_dict[id][precise_minute] = {} diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index f2319942c14..2a9995a4e59 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -32,6 +32,9 @@ from litellm.types.proxy.guardrails.guardrail_hooks.qualifire import ( from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import ( ToolPermissionGuardrailConfigModel, ) +from litellm.types.proxy.guardrails.guardrail_hooks.hiddenlayer import ( + HiddenlayerGuardrailConfigModel +) """ Pydantic object defining how to set guardrails on litellm proxy @@ -763,6 +766,7 @@ class LitellmParams( IBMGuardrailsBaseConfigModel, QualifireGuardrailConfigModel, BlockCodeExecutionGuardrailConfigModel, + HiddenlayerGuardrailConfigModel ): guardrail: str = Field(description="The type of guardrail integration to use") mode: Union[str, List[str], Mode] = Field( diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/hiddenlayer.py b/litellm/types/proxy/guardrails/guardrail_hooks/hiddenlayer.py index c3132846ada..4a0e5a23389 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/hiddenlayer.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/hiddenlayer.py @@ -32,6 +32,8 @@ class HiddenlayerGuardrailConfigModel(GuardrailConfigModel): description="The Hiddenlayer Secret Key for the Hiddenlayer API.. If not provided, the `HIDDENLAYER_CLIENT_SECRET` environment variable is checked.", ) + version: Optional[int] = Field(default=2, description="Hiddenlayer guardrail version to use.") + @staticmethod def ui_friendly_name() -> str: return "Hiddenlayer Guardrail" diff --git a/tests/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py b/tests/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py index a8e427d3bc1..d814f8ec97f 100644 --- a/tests/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py +++ b/tests/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py @@ -22,6 +22,7 @@ from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation _is_multimodal_input, _parse_data_url, process_embed_content_response, + process_response, transform_openai_input_gemini_content, transform_openai_input_gemini_embed_content, ) @@ -563,3 +564,32 @@ def test_vertex_ai_text_only_embedding_uses_embed_content(): assert data["content"]["parts"][0]["text"] == "Hello, world!" assert len(response.data) == 1 + +def test_batch_embeddings_response_has_correct_indices_and_order(): + """Test that process_response assigns sequential indices and preserves order.""" + response_json = { + "embeddings": [ + {"values": [0.1, 0.2, 0.3]}, + {"values": [0.4, 0.5, 0.6]}, + {"values": [0.7, 0.8, 0.9]}, + ] + } + expected_values = [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6], [0.7, 0.8, 0.9]] + + model_response = EmbeddingResponse() + result = process_response( + input=["first", "second", "third"], + model_response=model_response, + model="text-embedding-004", + _predictions=response_json, + ) + + assert len(result.data) == 3 + for i, embedding in enumerate(result.data): + assert ( + embedding.index == i + ), f"embedding {i} has index={embedding.index}, expected {i}" + assert ( + embedding.embedding == expected_values[i] + ), f"embedding {i} has wrong values: {embedding.embedding}" + diff --git a/tests/local_testing/test_lowest_latency_routing.py b/tests/local_testing/test_lowest_latency_routing.py index 429aae88b87..194c35d6642 100644 --- a/tests/local_testing/test_lowest_latency_routing.py +++ b/tests/local_testing/test_lowest_latency_routing.py @@ -964,3 +964,390 @@ async def test_lowest_latency_routing_time_to_first_token(sync_mode): assert len(selected_deployments.keys()) == 1 assert "1" in list(selected_deployments.keys()) + + +def test_latency_list_trimming_discards_oldest_entry(): + """ + When the latency list reaches max_latency_list_size, the oldest entry is + discarded to make room for new entries. The newest entry is appended at + the end of the list. + """ + max_size = 3 + test_cache = DualCache() + lowest_latency_logger = LowestLatencyLoggingHandler( + router_cache=test_cache, routing_args={"max_latency_list_size": max_size} + ) + + model_group = "gpt-3.5-turbo" + deployment_id = "test-deployment" + kwargs = { + "litellm_params": { + "metadata": { + "model_group": model_group, + "deployment": "azure/gpt-4.1-mini", + }, + "model_info": {"id": deployment_id}, + } + } + + # With 1 completion token, the logged latency value equals the raw + # response time, so we can use distinct, identifiable values. + latencies_to_add = [] + for i in range(max_size + 1): # One more than max to trigger trimming + start_time = time.time() + response_obj = {"usage": {"total_tokens": 1, "completion_tokens": 1}} + expected_latency = float(i + 1) # 1.0, 2.0, 3.0, 4.0 + end_time = start_time + expected_latency + latencies_to_add.append(expected_latency) + + lowest_latency_logger.log_success_event( + response_obj=response_obj, + kwargs=kwargs, + start_time=start_time, + end_time=end_time, + ) + + latency_key = f"{model_group}_map" + cached_data = test_cache.get_cache(key=latency_key) + latency_list = cached_data[deployment_id]["latency"] + + assert ( + len(latency_list) == max_size + ), f"Expected {max_size} entries, got {len(latency_list)}" + + newest_latency = latencies_to_add[-1] # 4.0 + oldest_latency = latencies_to_add[0] # 1.0 + tolerance = 0.1 + + # Newest entry is at the end of the list. + assert ( + abs(latency_list[-1] - newest_latency) < tolerance + ), f"Newest latency {newest_latency} should be at end, got {latency_list[-1]}" + + # Oldest entry is no longer in the list. + for latency in latency_list: + assert ( + abs(latency - oldest_latency) > tolerance + ), f"Oldest latency {oldest_latency} should have been discarded, found {latency}" + + +@pytest.mark.asyncio +async def test_latency_list_trimming_discards_oldest_entry_async(): + """ + Async counterpart: the oldest entry is discarded when the latency list is + trimmed. + """ + max_size = 3 + test_cache = DualCache() + lowest_latency_logger = LowestLatencyLoggingHandler( + router_cache=test_cache, routing_args={"max_latency_list_size": max_size} + ) + + model_group = "gpt-3.5-turbo" + deployment_id = "test-deployment" + kwargs = { + "litellm_params": { + "metadata": { + "model_group": model_group, + "deployment": "azure/gpt-4.1-mini", + }, + "model_info": {"id": deployment_id}, + } + } + + latencies_to_add = [] + for i in range(max_size + 1): + start_time = time.time() + response_obj = {"usage": {"total_tokens": 1, "completion_tokens": 1}} + expected_latency = float(i + 1) + end_time = start_time + expected_latency + latencies_to_add.append(expected_latency) + + await lowest_latency_logger.async_log_success_event( + response_obj=response_obj, + kwargs=kwargs, + start_time=start_time, + end_time=end_time, + ) + + latency_key = f"{model_group}_map" + cached_data = await test_cache.async_get_cache(key=latency_key) + latency_list = cached_data[deployment_id]["latency"] + + assert len(latency_list) == max_size + + newest_latency = latencies_to_add[-1] + oldest_latency = latencies_to_add[0] + tolerance = 0.1 + + assert ( + abs(latency_list[-1] - newest_latency) < tolerance + ), f"Newest latency {newest_latency} should be at end of list" + + for latency in latency_list: + assert ( + abs(latency - oldest_latency) > tolerance + ), f"Oldest latency {oldest_latency} should have been discarded" + + +def test_ttft_list_trimming_discards_oldest_entry(): + """ + The time_to_first_token list trims the oldest entry when full, matching + the behavior of the latency list. + """ + max_size = 3 + test_cache = DualCache() + lowest_latency_logger = LowestLatencyLoggingHandler( + router_cache=test_cache, routing_args={"max_latency_list_size": max_size} + ) + + model_group = "gpt-3.5-turbo" + deployment_id = "test-deployment" + + ttft_values = [] + for i in range(max_size + 1): + start_time = time.time() + expected_ttft = float(i + 1) * 0.1 # 0.1, 0.2, 0.3, 0.4 + completion_start_time = start_time + expected_ttft + end_time = start_time + float(i + 1) + ttft_values.append(expected_ttft) + + kwargs = { + "litellm_params": { + "metadata": { + "model_group": model_group, + "deployment": "azure/gpt-4.1-mini", + }, + "model_info": {"id": deployment_id}, + }, + "stream": True, + "completion_start_time": completion_start_time, + } + # TTFT is only recorded when response_obj is a ModelResponse. + response_obj = litellm.ModelResponse( + usage=litellm.Usage(completion_tokens=1, total_tokens=1) + ) + + lowest_latency_logger.log_success_event( + response_obj=response_obj, + kwargs=kwargs, + start_time=start_time, + end_time=end_time, + ) + + latency_key = f"{model_group}_map" + cached_data = test_cache.get_cache(key=latency_key) + ttft_list = cached_data[deployment_id].get("time_to_first_token", []) + + assert ( + len(ttft_list) == max_size + ), f"Expected {max_size} entries, got {len(ttft_list)}" + + newest_ttft = ttft_values[-1] + oldest_ttft = ttft_values[0] + tolerance = 0.05 + + assert ( + abs(ttft_list[-1] - newest_ttft) < tolerance + ), f"Newest TTFT {newest_ttft} should be at end of list" + + for ttft in ttft_list: + assert ( + abs(ttft - oldest_ttft) > tolerance + ), f"Oldest TTFT {oldest_ttft} should have been discarded" + + +@pytest.mark.asyncio +async def test_timeout_penalty_discards_oldest_entry(): + """ + Timeout penalties (1000.0) are appended to the latency list and, when the + list is full, the oldest entry is discarded. + """ + max_size = 3 + test_cache = DualCache() + lowest_latency_logger = LowestLatencyLoggingHandler( + router_cache=test_cache, routing_args={"max_latency_list_size": max_size} + ) + + model_group = "gpt-3.5-turbo" + deployment_id = "test-deployment" + kwargs = { + "litellm_params": { + "metadata": { + "model_group": model_group, + "deployment": "azure/gpt-4.1-mini", + }, + "model_info": {"id": deployment_id}, + } + } + + # Fill the list with max_size normal latency entries first. + for i in range(max_size): + start_time = time.time() + response_obj = {"usage": {"total_tokens": 1, "completion_tokens": 1}} + end_time = start_time + float(i + 1) + + await lowest_latency_logger.async_log_success_event( + response_obj=response_obj, + kwargs=kwargs, + start_time=start_time, + end_time=end_time, + ) + + # Trigger a timeout failure: this appends 1000.0 and should discard the + # oldest normal entry (1.0). + timeout_kwargs = { + **kwargs, + "exception": litellm.Timeout( + message="Request timed out", model="test-model", llm_provider="test" + ), + } + + await lowest_latency_logger.async_log_failure_event( + kwargs=timeout_kwargs, + response_obj=None, + start_time=time.time(), + end_time=time.time() + 30, + ) + + latency_key = f"{model_group}_map" + cached_data = await test_cache.async_get_cache(key=latency_key) + latency_list = cached_data[deployment_id]["latency"] + + assert len(latency_list) == max_size + + # Timeout penalty is the newest entry. + assert ( + latency_list[-1] == 1000.0 + ), f"Timeout penalty should be at end of list, got {latency_list[-1]}" + + # Oldest normal entry (1.0) has been discarded. + tolerance = 0.1 + for latency in latency_list[:-1]: + assert ( + abs(latency - 1.0) > tolerance + ), f"Oldest latency 1.0 should have been discarded, found {latency}" + + +def test_list_order_preserved_after_multiple_trims(): + """ + After many trims, the list still holds the most recent `max_size` entries + in insertion order (oldest at index 0, newest at index -1). + """ + max_size = 3 + test_cache = DualCache() + lowest_latency_logger = LowestLatencyLoggingHandler( + router_cache=test_cache, routing_args={"max_latency_list_size": max_size} + ) + + model_group = "gpt-3.5-turbo" + deployment_id = "test-deployment" + kwargs = { + "litellm_params": { + "metadata": { + "model_group": model_group, + "deployment": "azure/gpt-4.1-mini", + }, + "model_info": {"id": deployment_id}, + } + } + + # Add 10 entries (7 more than max) to trigger multiple trims. + all_latencies = [] + for i in range(10): + start_time = time.time() + response_obj = {"usage": {"total_tokens": 1, "completion_tokens": 1}} + expected_latency = float(i + 1) + end_time = start_time + expected_latency + all_latencies.append(expected_latency) + + lowest_latency_logger.log_success_event( + response_obj=response_obj, + kwargs=kwargs, + start_time=start_time, + end_time=end_time, + ) + + latency_key = f"{model_group}_map" + cached_data = test_cache.get_cache(key=latency_key) + latency_list = cached_data[deployment_id]["latency"] + + assert len(latency_list) == max_size + + # After inserting 1..10 with max_size=3, the list should be [8, 9, 10]. + expected_remaining = all_latencies[-max_size:] + tolerance = 0.1 + + for i, expected in enumerate(expected_remaining): + assert ( + abs(latency_list[i] - expected) < tolerance + ), f"At index {i}, expected ~{expected}, got {latency_list[i]}" + + +@pytest.mark.asyncio +async def test_ttft_list_trimming_discards_oldest_entry_async(): + """ + Async counterpart: the time_to_first_token list trims the oldest entry + when full. Exercises the async_log_success_event TTFT path, which only + runs when response_obj is a ModelResponse and the call is marked as + streaming with a completion_start_time. + """ + max_size = 3 + test_cache = DualCache() + lowest_latency_logger = LowestLatencyLoggingHandler( + router_cache=test_cache, routing_args={"max_latency_list_size": max_size} + ) + + model_group = "gpt-3.5-turbo" + deployment_id = "test-deployment" + + ttft_values = [] + for i in range(max_size + 1): + start_time = time.time() + expected_ttft = float(i + 1) * 0.1 # 0.1, 0.2, 0.3, 0.4 + completion_start_time = start_time + expected_ttft + end_time = start_time + float(i + 1) + ttft_values.append(expected_ttft) + + kwargs = { + "litellm_params": { + "metadata": { + "model_group": model_group, + "deployment": "azure/gpt-4.1-mini", + }, + "model_info": {"id": deployment_id}, + }, + "stream": True, + "completion_start_time": completion_start_time, + } + response_obj = litellm.ModelResponse( + usage=litellm.Usage(completion_tokens=1, total_tokens=1) + ) + + await lowest_latency_logger.async_log_success_event( + response_obj=response_obj, + kwargs=kwargs, + start_time=start_time, + end_time=end_time, + ) + + latency_key = f"{model_group}_map" + cached_data = await test_cache.async_get_cache(key=latency_key) + ttft_list = cached_data[deployment_id].get("time_to_first_token", []) + + assert ( + len(ttft_list) == max_size + ), f"Expected {max_size} entries, got {len(ttft_list)}" + + newest_ttft = ttft_values[-1] + oldest_ttft = ttft_values[0] + tolerance = 0.05 + + assert ( + abs(ttft_list[-1] - newest_ttft) < tolerance + ), f"Newest TTFT {newest_ttft} should be at end of list" + + for ttft in ttft_list: + assert ( + abs(ttft - oldest_ttft) > tolerance + ), f"Oldest TTFT {oldest_ttft} should have been discarded" diff --git a/tests/test_litellm/caching/test_in_memory_cache.py b/tests/test_litellm/caching/test_in_memory_cache.py index e7cc7f80ab3..8828ebf207e 100644 --- a/tests/test_litellm/caching/test_in_memory_cache.py +++ b/tests/test_litellm/caching/test_in_memory_cache.py @@ -97,26 +97,26 @@ def test_in_memory_cache_max_size_with_ttl(): """ in_memory_cache = InMemoryCache(max_size_in_memory=3) long_ttl = 86400 # 1 day - + # Fill the cache to max capacity for i in range(3): in_memory_cache.set_cache(key=f"key_{i}", value=f"value_{i}", ttl=long_ttl) time.sleep(0.01) # Small delay to ensure different timestamps - + assert len(in_memory_cache.cache_dict) == 3 assert len(in_memory_cache.ttl_dict) == 3 - + # Add another item - should evict the earliest item in_memory_cache.set_cache(key="key_3", value="value_3", ttl=long_ttl) - + # Cache should still be at max size, not larger assert len(in_memory_cache.cache_dict) == 3 assert len(in_memory_cache.ttl_dict) == 3 - + # key_0 should have been evicted (it was added first) assert "key_0" not in in_memory_cache.cache_dict assert "key_0" not in in_memory_cache.ttl_dict - + # Other keys should still be present assert "key_1" in in_memory_cache.cache_dict assert "key_2" in in_memory_cache.cache_dict @@ -128,26 +128,26 @@ def test_in_memory_cache_expired_items_evicted_first(): Test that expired items are evicted before non-expired items when cache is full. """ in_memory_cache = InMemoryCache(max_size_in_memory=3) - + # Add items with short TTL that will expire in_memory_cache.set_cache(key="expired_1", value="value_1", ttl=1) in_memory_cache.set_cache(key="expired_2", value="value_2", ttl=1) - + # Add item with long TTL in_memory_cache.set_cache(key="long_lived", value="value_long", ttl=86400) - + assert len(in_memory_cache.cache_dict) == 3 - + # Wait for short TTL items to expire time.sleep(2) - + # Add new item - should evict expired items first, not the long-lived one in_memory_cache.set_cache(key="new_item", value="new_value", ttl=86400) - + # Long-lived item should still be present assert "long_lived" in in_memory_cache.cache_dict assert "new_item" in in_memory_cache.cache_dict - + # Expired items should be gone assert "expired_1" not in in_memory_cache.cache_dict assert "expired_2" not in in_memory_cache.cache_dict @@ -160,29 +160,33 @@ def test_in_memory_cache_eviction_order(): Test that when non-expired items need to be evicted, those with earliest expiration times are evicted first. """ in_memory_cache = InMemoryCache(max_size_in_memory=2) - + # Add items with different TTLs now = time.time() - in_memory_cache.set_cache(key="early_expire", value="value_1", ttl=100) # expires in 100 seconds + in_memory_cache.set_cache( + key="early_expire", value="value_1", ttl=100 + ) # expires in 100 seconds time.sleep(0.01) - in_memory_cache.set_cache(key="late_expire", value="value_2", ttl=200) # expires in 200 seconds - + in_memory_cache.set_cache( + key="late_expire", value="value_2", ttl=200 + ) # expires in 200 seconds + # Verify TTL order early_ttl = in_memory_cache.ttl_dict["early_expire"] late_ttl = in_memory_cache.ttl_dict["late_expire"] assert early_ttl < late_ttl, "early_expire should have earlier expiration time" - + assert len(in_memory_cache.cache_dict) == 2 - + # Add third item - should evict the one with earliest expiration time in_memory_cache.set_cache(key="new_item", value="value_3", ttl=300) - + assert len(in_memory_cache.cache_dict) == 2 - + # Item with earliest expiration should be evicted assert "early_expire" not in in_memory_cache.cache_dict assert "early_expire" not in in_memory_cache.ttl_dict - + # Items with later expiration should remain assert "late_expire" in in_memory_cache.cache_dict assert "new_item" in in_memory_cache.cache_dict @@ -199,3 +203,23 @@ def test_in_memory_cache_heap_size_staus_bounded(): # Expiration heap should only have 1 entry assert len(in_memory_cache.expiration_heap) == 1 + + +def test_in_memory_cache_prunes_expired_heap_entries_below_capacity(): + """ + Re-inserting expired keys below capacity should not grow expiration_heap + without bound. + """ + in_memory_cache = InMemoryCache(max_size_in_memory=200, default_ttl=1) + + for cycle in range(3): + for i in range(5): + in_memory_cache.set_cache(key=f"key_{i}", value=f"value_{cycle}_{i}", ttl=1) + time.sleep(1.1) + + for i in range(5): + in_memory_cache.set_cache(key=f"key_{i}", value=f"value_final_{i}", ttl=1) + + assert len(in_memory_cache.cache_dict) == 5 + assert len(in_memory_cache.ttl_dict) == 5 + assert len(in_memory_cache.expiration_heap) == 5 diff --git a/tests/test_litellm/integrations/datadog/test_datadog_logger_batching.py b/tests/test_litellm/integrations/datadog/test_datadog_logger_batching.py new file mode 100644 index 00000000000..e4d7227cc88 --- /dev/null +++ b/tests/test_litellm/integrations/datadog/test_datadog_logger_batching.py @@ -0,0 +1,267 @@ +from unittest.mock import AsyncMock, Mock, patch + +import pytest +from httpx import Request, Response + +from litellm.integrations.datadog.datadog import DataDogLogger +from litellm.types.integrations.datadog import DatadogPayload + + +@pytest.fixture +def datadog_env(monkeypatch): + monkeypatch.setenv("DD_API_KEY", "test_api_key") + monkeypatch.setenv("DD_SITE", "test.datadoghq.com") + + +@pytest.mark.asyncio +async def test_async_send_batch_keeps_events_appended_during_send(datadog_env): + with patch("asyncio.create_task"): + logger = DataDogLogger() + + logger.log_queue = [ + DatadogPayload( + ddsource="litellm", + ddtags="env:test", + hostname="host", + message=f'{{"event": {i}}}', + service="svc", + status="info", + ) + for i in range(2) + ] + + async def _mock_send(data): + logger.log_queue.append( + DatadogPayload( + ddsource="litellm", + ddtags="env:test", + hostname="host", + message='{"event": 2}', + service="svc", + status="info", + ) + ) + return Response( + 202, request=Request("POST", "https://example.com"), text="Accepted" + ) + + logger.async_send_compressed_data = AsyncMock(side_effect=_mock_send) + + await logger.async_send_batch() + + assert logger.async_send_compressed_data.await_count == 1 + sent_batch = logger.async_send_compressed_data.await_args.args[0] + assert len(sent_batch) == 2 + assert len(logger.log_queue) == 1 + assert logger.log_queue[0]["message"] == '{"event": 2}' + + +@pytest.mark.asyncio +async def test_failure_hook_threshold_flush_uses_flush_queue(datadog_env): + with patch("asyncio.create_task"): + logger = DataDogLogger() + + logger.batch_size = 1 + logger.flush_queue = AsyncMock() + + await logger.async_post_call_failure_hook( + request_data={}, + original_exception=Exception("boom"), + user_api_key_dict=type("UserKey", (), {})(), + traceback_str="trace", + ) + + logger.flush_queue.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_async_send_batch_requeues_events_on_413(datadog_env): + with patch("asyncio.create_task"): + logger = DataDogLogger() + + logger.log_queue = [ + DatadogPayload( + ddsource="litellm", + ddtags="env:test", + hostname="host", + message=f'{{"event": {i}}}', + service="svc", + status="info", + ) + for i in range(2) + ] + + logger.async_send_compressed_data = AsyncMock( + return_value=Response( + 413, + request=Request("POST", "https://example.com"), + text="Payload Too Large", + ) + ) + + await logger.async_send_batch() + + assert logger.async_send_compressed_data.await_count == 1 + assert len(logger.log_queue) == 2 + assert [event["message"] for event in logger.log_queue] == [ + '{"event": 0}', + '{"event": 1}', + ] + + +@pytest.mark.asyncio +async def test_async_send_batch_handles_empty_queue(datadog_env): + with patch("asyncio.create_task"): + logger = DataDogLogger() + + logger.log_queue = [] + logger.async_send_compressed_data = AsyncMock() + + await logger.async_send_batch() + + logger.async_send_compressed_data.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_async_send_batch_requeues_events_on_exception(datadog_env): + with patch("asyncio.create_task"): + logger = DataDogLogger() + + logger.log_queue = [ + DatadogPayload( + ddsource="litellm", + ddtags="env:test", + hostname="host", + message=f'{{"event": {i}}}', + service="svc", + status="info", + ) + for i in range(2) + ] + + logger.async_send_compressed_data = AsyncMock(side_effect=RuntimeError("boom")) + + await logger.async_send_batch() + + assert [event["message"] for event in logger.log_queue] == [ + '{"event": 0}', + '{"event": 1}', + ] + + +@pytest.mark.asyncio +async def test_log_async_event_threshold_flush_uses_flush_queue(datadog_env): + with patch("asyncio.create_task"): + logger = DataDogLogger() + + logger.batch_size = 1 + logger.flush_queue = AsyncMock() + logger.create_datadog_logging_payload = Mock( + return_value=DatadogPayload( + ddsource="litellm", + ddtags="env:test", + hostname="host", + message='{"event": 0}', + service="svc", + status="info", + ) + ) + + await logger._log_async_event( + kwargs={}, + response_obj={}, + start_time=None, + end_time=None, + ) + + logger.flush_queue.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_flush_queue_updates_last_flush_time(datadog_env): + with patch("asyncio.create_task"): + logger = DataDogLogger() + + logger.log_queue = [ + DatadogPayload( + ddsource="litellm", + ddtags="env:test", + hostname="host", + message='{"event": 0}', + service="svc", + status="info", + ) + ] + logger.last_flush_time = 0 + + async def _successful_send(): + logger.log_queue = [] + + logger.async_send_batch = AsyncMock(side_effect=_successful_send) + + await logger.flush_queue() + + logger.async_send_batch.assert_awaited_once() + assert logger.last_flush_time > 0 + + +@pytest.mark.asyncio +async def test_flush_queue_does_not_update_last_flush_time_when_send_requeues( + datadog_env, +): + with patch("asyncio.create_task"): + logger = DataDogLogger() + + logger.log_queue = [ + DatadogPayload( + ddsource="litellm", + ddtags="env:test", + hostname="host", + message='{"event": 0}', + service="svc", + status="info", + ) + ] + logger.last_flush_time = 123.0 + + async def _requeue_batch(): + logger.log_queue = [ + DatadogPayload( + ddsource="litellm", + ddtags="env:test", + hostname="host", + message='{"event": 0}', + service="svc", + status="info", + ) + ] + + logger.async_send_batch = AsyncMock(side_effect=_requeue_batch) + + await logger.flush_queue() + + logger.async_send_batch.assert_awaited_once() + assert logger.last_flush_time == 123.0 + + +@pytest.mark.asyncio +async def test_flush_queue_returns_without_lock(datadog_env): + with patch("asyncio.create_task"): + logger = DataDogLogger() + + logger.flush_lock = None + logger.log_queue = [ + DatadogPayload( + ddsource="litellm", + ddtags="env:test", + hostname="host", + message='{"event": 0}', + service="svc", + status="info", + ) + ] + logger.async_send_batch = AsyncMock() + + await logger.flush_queue() + + logger.async_send_batch.assert_not_awaited() diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_tool_args.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_tool_args.py new file mode 100644 index 00000000000..bd39e420607 --- /dev/null +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_tool_args.py @@ -0,0 +1,383 @@ +""" +Test that AnthropicStreamWrapper emits input_json_delta when tool arguments +are bundled in the same streaming chunk as the function name/id. + +Providers like xAI and Gemini include tool_call function arguments in +the first chunk rather than streaming them separately (OpenAI-style). +Without the fix, the AnthropicStreamWrapper silently dropped these +arguments, causing tool_use blocks to arrive with empty input {}. +""" + +import os +import sys +from typing import List +from unittest.mock import MagicMock + +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.anthropic.experimental_pass_through.adapters.streaming_iterator import ( + AnthropicStreamWrapper, +) +from litellm.types.utils import ( + ChatCompletionDeltaToolCall, + Delta, + Function, + StreamingChoices, +) + + +def _make_chunk( + delta: Delta, + finish_reason: str = None, +) -> MagicMock: + """Create a minimal streaming chunk with the given delta and finish_reason.""" + chunk = MagicMock() + chunk.choices = [ + StreamingChoices( + finish_reason=finish_reason, + index=0, + delta=delta, + logprobs=None, + ) + ] + chunk.usage = None + chunk._hidden_params = {} + return chunk + + +def _collect_events_sync(wrapper: AnthropicStreamWrapper) -> List[dict]: + """Drain all events from a sync AnthropicStreamWrapper.""" + events = [] + for event in wrapper: + events.append(event) + return events + + +async def _collect_events_async(wrapper: AnthropicStreamWrapper) -> List[dict]: + """Drain all events from an async AnthropicStreamWrapper.""" + events = [] + async for event in wrapper: + events.append(event) + return events + + +@pytest.mark.asyncio +async def test_async_stream_emits_input_json_delta_for_bundled_tool_args(): + """ + When a provider bundles tool_call arguments in the first streaming chunk + (same chunk as name/id), the async wrapper must emit an input_json_delta + content_block_delta after the tool_use content_block_start. + """ + # Chunk 1: text content + text_chunk = _make_chunk(Delta(content="Hello", role="assistant", tool_calls=None)) + + # Chunk 2: tool call with name AND arguments in the same chunk (xAI/Gemini style) + tool_chunk = _make_chunk( + Delta( + content=None, + role="assistant", + tool_calls=[ + ChatCompletionDeltaToolCall( + id="call_abc123", + function=Function( + name="get_weather", + arguments='{"location": "Boston"}', + ), + type="function", + index=0, + ) + ], + ) + ) + + # Chunk 3: finish + finish_chunk = _make_chunk( + Delta(content=None, role="assistant", tool_calls=None), + finish_reason="tool_calls", + ) + + async def mock_stream(): + for c in [text_chunk, tool_chunk, finish_chunk]: + yield c + + wrapper = AnthropicStreamWrapper( + completion_stream=mock_stream(), + model="test-model", + ) + + events = await _collect_events_async(wrapper) + event_types = [e.get("type") if isinstance(e, dict) else str(e) for e in events] + + # Find the tool_use content_block_start and subsequent input_json_delta + tool_start_idx = None + input_json_delta_idx = None + + for i, event in enumerate(events): + if not isinstance(event, dict): + continue + if ( + event.get("type") == "content_block_start" + and isinstance(event.get("content_block"), dict) + and event["content_block"].get("type") == "tool_use" + ): + tool_start_idx = i + if ( + event.get("type") == "content_block_delta" + and isinstance(event.get("delta"), dict) + and event["delta"].get("type") == "input_json_delta" + ): + input_json_delta_idx = i + + assert ( + tool_start_idx is not None + ), f"Expected content_block_start with type=tool_use; events: {event_types}" + assert ( + input_json_delta_idx is not None + ), f"Expected content_block_delta with input_json_delta; events: {event_types}" + assert ( + input_json_delta_idx == tool_start_idx + 1 + ), "input_json_delta should immediately follow the tool_use content_block_start" + + # Verify the delta carries the tool arguments + delta_event = events[input_json_delta_idx] + assert delta_event["delta"][ + "partial_json" + ], "input_json_delta should have non-empty partial_json" + + +@pytest.mark.asyncio +async def test_async_stream_no_extra_delta_when_tool_args_empty(): + """ + When a provider sends tool name/id WITHOUT arguments in the first chunk + (OpenAI-style), the wrapper should NOT emit an extra input_json_delta + after content_block_start. This verifies backward compatibility. + """ + # Chunk 1: text + text_chunk = _make_chunk(Delta(content="Hi", role="assistant", tool_calls=None)) + + # Chunk 2: tool call with name but NO arguments (OpenAI-style) + tool_name_chunk = _make_chunk( + Delta( + content=None, + role="assistant", + tool_calls=[ + ChatCompletionDeltaToolCall( + id="call_xyz789", + function=Function(name="get_weather", arguments=""), + type="function", + index=0, + ) + ], + ) + ) + + # Chunk 3: arguments streamed separately + tool_args_chunk = _make_chunk( + Delta( + content=None, + role="assistant", + tool_calls=[ + ChatCompletionDeltaToolCall( + id=None, + function=Function(name=None, arguments='{"location": "NYC"}'), + type="function", + index=0, + ) + ], + ) + ) + + # Chunk 4: finish + finish_chunk = _make_chunk( + Delta(content=None, role="assistant", tool_calls=None), + finish_reason="tool_calls", + ) + + async def mock_stream(): + for c in [text_chunk, tool_name_chunk, tool_args_chunk, finish_chunk]: + yield c + + wrapper = AnthropicStreamWrapper( + completion_stream=mock_stream(), + model="test-model", + ) + + events = await _collect_events_async(wrapper) + + # Find tool_use content_block_start + tool_start_idx = None + for i, event in enumerate(events): + if not isinstance(event, dict): + continue + if ( + event.get("type") == "content_block_start" + and isinstance(event.get("content_block"), dict) + and event["content_block"].get("type") == "tool_use" + ): + tool_start_idx = i + break + + assert tool_start_idx is not None + + # Count how many input_json_delta events appear after the tool_use block start. + # With empty args in the trigger chunk, only the subsequent tool_args_chunk + # should produce one — not the trigger chunk itself. + input_json_deltas = [ + e + for e in events[tool_start_idx + 1 :] + if isinstance(e, dict) + and e.get("type") == "content_block_delta" + and isinstance(e.get("delta"), dict) + and e["delta"].get("type") == "input_json_delta" + ] + assert len(input_json_deltas) == 1, ( + f"Expected exactly 1 input_json_delta (from the follow-up chunk), " + f"got {len(input_json_deltas)}" + ) + assert input_json_deltas[0]["delta"]["partial_json"] == '{"location": "NYC"}' + + +def test_sync_stream_emits_input_json_delta_for_bundled_tool_args(): + """ + Sync counterpart: when a provider bundles tool_call arguments in the first + streaming chunk, the sync wrapper must also emit the input_json_delta. + """ + text_chunk = _make_chunk(Delta(content="Hello", role="assistant", tool_calls=None)) + tool_chunk = _make_chunk( + Delta( + content=None, + role="assistant", + tool_calls=[ + ChatCompletionDeltaToolCall( + id="call_abc123", + function=Function( + name="get_weather", + arguments='{"location": "Boston"}', + ), + type="function", + index=0, + ) + ], + ) + ) + finish_chunk = _make_chunk( + Delta(content=None, role="assistant", tool_calls=None), + finish_reason="tool_calls", + ) + + wrapper = AnthropicStreamWrapper( + completion_stream=iter([text_chunk, tool_chunk, finish_chunk]), + model="test-model", + ) + + events = _collect_events_sync(wrapper) + event_types = [e.get("type") if isinstance(e, dict) else str(e) for e in events] + + tool_start_idx = None + input_json_delta_idx = None + + for i, event in enumerate(events): + if not isinstance(event, dict): + continue + if ( + event.get("type") == "content_block_start" + and isinstance(event.get("content_block"), dict) + and event["content_block"].get("type") == "tool_use" + ): + tool_start_idx = i + if ( + event.get("type") == "content_block_delta" + and isinstance(event.get("delta"), dict) + and event["delta"].get("type") == "input_json_delta" + ): + input_json_delta_idx = i + + assert ( + tool_start_idx is not None + ), f"Expected content_block_start with type=tool_use; events: {event_types}" + assert ( + input_json_delta_idx is not None + ), f"Expected content_block_delta with input_json_delta; events: {event_types}" + assert ( + input_json_delta_idx == tool_start_idx + 1 + ), "input_json_delta should immediately follow the tool_use content_block_start" + assert events[input_json_delta_idx]["delta"]["partial_json"] + + +def test_sync_stream_no_extra_delta_when_tool_args_empty(): + """ + Sync counterpart: empty args (OpenAI-style) should not emit an extra + input_json_delta from the trigger chunk. + """ + text_chunk = _make_chunk(Delta(content="Hi", role="assistant", tool_calls=None)) + tool_name_chunk = _make_chunk( + Delta( + content=None, + role="assistant", + tool_calls=[ + ChatCompletionDeltaToolCall( + id="call_xyz789", + function=Function(name="get_weather", arguments=""), + type="function", + index=0, + ) + ], + ) + ) + tool_args_chunk = _make_chunk( + Delta( + content=None, + role="assistant", + tool_calls=[ + ChatCompletionDeltaToolCall( + id=None, + function=Function(name=None, arguments='{"location": "NYC"}'), + type="function", + index=0, + ) + ], + ) + ) + finish_chunk = _make_chunk( + Delta(content=None, role="assistant", tool_calls=None), + finish_reason="tool_calls", + ) + + wrapper = AnthropicStreamWrapper( + completion_stream=iter( + [text_chunk, tool_name_chunk, tool_args_chunk, finish_chunk] + ), + model="test-model", + ) + + events = _collect_events_sync(wrapper) + + tool_start_idx = None + for i, event in enumerate(events): + if not isinstance(event, dict): + continue + if ( + event.get("type") == "content_block_start" + and isinstance(event.get("content_block"), dict) + and event["content_block"].get("type") == "tool_use" + ): + tool_start_idx = i + break + + assert tool_start_idx is not None + + input_json_deltas = [ + e + for e in events[tool_start_idx + 1 :] + if isinstance(e, dict) + and e.get("type") == "content_block_delta" + and isinstance(e.get("delta"), dict) + and e["delta"].get("type") == "input_json_delta" + ] + assert len(input_json_deltas) == 1, ( + f"Expected exactly 1 input_json_delta (from the follow-up chunk), " + f"got {len(input_json_deltas)}" + ) + assert input_json_deltas[0]["delta"]["partial_json"] == '{"location": "NYC"}' diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index ddc404cb8c7..a0979664943 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -237,7 +237,9 @@ def test_vertex_ai_response_json_schema_preserves_refs_for_gemini_2(): # $defs and $ref should be preserved (not unpacked) assert "response_json_schema" in transformed_request result_schema = transformed_request["response_json_schema"] - assert "$defs" in result_schema, "responseJsonSchema should preserve $defs for Gemini 2.0+" + assert ( + "$defs" in result_schema + ), "responseJsonSchema should preserve $defs for Gemini 2.0+" def test_vertex_ai_get_json_schema_preserves_refs_for_nested_pydantic(): @@ -317,14 +319,22 @@ def test_vertex_ai_response_json_schema_for_gemini_2(): # Types should be lowercase (standard JSON Schema format) assert transformed_request["response_json_schema"]["type"] == "object" - assert transformed_request["response_json_schema"]["properties"]["name"]["type"] == "string" - assert transformed_request["response_json_schema"]["properties"]["age"]["type"] == "integer" + assert ( + transformed_request["response_json_schema"]["properties"]["name"]["type"] + == "string" + ) + assert ( + transformed_request["response_json_schema"]["properties"]["age"]["type"] + == "integer" + ) # Should NOT have propertyOrdering (not needed for responseJsonSchema) assert "propertyOrdering" not in transformed_request["response_json_schema"] # additionalProperties should be preserved (supported by responseJsonSchema) - assert transformed_request["response_json_schema"].get("additionalProperties") == False + assert ( + transformed_request["response_json_schema"].get("additionalProperties") == False + ) def test_vertex_ai_response_schema_for_old_models(): @@ -581,7 +591,7 @@ def test_streaming_chunk_with_tool_calls_and_thought_includes_reasoning_content( "args": {"timezone": "America/New_York"}, }, "thoughtSignature": "EsEDCr4DAdHtim...", # Just a token, not reasoning - } + }, ] }, "finishReason": "STOP", @@ -600,12 +610,18 @@ def test_streaming_chunk_with_tool_calls_and_thought_includes_reasoning_content( streaming_chunk = iterator.chunk_parser(chunk) # Verify reasoning_content comes from the thought: true part - assert streaming_chunk.choices[0].delta.reasoning_content == "Let me think about how to get the time..." + assert ( + streaming_chunk.choices[0].delta.reasoning_content + == "Let me think about how to get the time..." + ) # Verify tool calls are also present assert streaming_chunk.choices[0].delta.tool_calls is not None assert len(streaming_chunk.choices[0].delta.tool_calls) == 1 - assert streaming_chunk.choices[0].delta.tool_calls[0].function.name == "get_current_time" + assert ( + streaming_chunk.choices[0].delta.tool_calls[0].function.name + == "get_current_time" + ) def test_streaming_chunk_with_tool_calls_no_thought_no_reasoning_content(): @@ -653,12 +669,15 @@ def test_streaming_chunk_with_tool_calls_no_thought_no_reasoning_content(): streaming_chunk = iterator.chunk_parser(chunk) # reasoning_content should be None - thoughtSignature alone does NOT mean reasoning - assert getattr(streaming_chunk.choices[0].delta, 'reasoning_content', None) is None + assert getattr(streaming_chunk.choices[0].delta, "reasoning_content", None) is None # Tool calls should still work assert streaming_chunk.choices[0].delta.tool_calls is not None assert len(streaming_chunk.choices[0].delta.tool_calls) == 1 - assert streaming_chunk.choices[0].delta.tool_calls[0].function.name == "get_current_time" + assert ( + streaming_chunk.choices[0].delta.tool_calls[0].function.name + == "get_current_time" + ) def test_check_finish_reason(): @@ -711,7 +730,10 @@ def test_vertex_ai_usage_metadata_response_token_count(): "promptTokenCount": 66, "responseTokenCount": 74, "totalTokenCount": 131, - "promptTokensDetails": [{"modality": "TEXT", "tokenCount": 57}, {"modality": "IMAGE", "tokenCount": 9}], + "promptTokensDetails": [ + {"modality": "TEXT", "tokenCount": 57}, + {"modality": "IMAGE", "tokenCount": 9}, + ], "responseTokensDetails": [{"modality": "TEXT", "tokenCount": 74}], } usage_metadata = UsageMetadata(**usage_metadata) @@ -741,9 +763,9 @@ def test_vertex_ai_usage_metadata_with_image_tokens(): "promptTokensDetails": [{"modality": "TEXT", "tokenCount": 14}], "candidatesTokensDetails": [ {"modality": "IMAGE", "tokenCount": 1120}, - {"modality": "TEXT", "tokenCount": 322} # 1442 - 1120 = 322 + {"modality": "TEXT", "tokenCount": 322}, # 1442 - 1120 = 322 ], - "thoughtsTokenCount": 158 + "thoughtsTokenCount": 158, } usage_metadata = UsageMetadata(**usage_metadata) result = v._calculate_usage(completion_response={"usageMetadata": usage_metadata}) @@ -785,7 +807,7 @@ def test_vertex_ai_usage_metadata_with_image_tokens_auto_calculated_text(): {"modality": "IMAGE", "tokenCount": 1120} # TEXT modality omitted - should be auto-calculated ], - "thoughtsTokenCount": 158 + "thoughtsTokenCount": 158, } usage_metadata = UsageMetadata(**usage_metadata) result = v._calculate_usage(completion_response={"usageMetadata": usage_metadata}) @@ -809,13 +831,13 @@ def test_vertex_ai_usage_metadata_with_image_tokens_auto_calculated_text(): def test_vertex_ai_usage_metadata_with_image_tokens_in_prompt(): """Test promptTokensDetails with IMAGE modality for multimodal inputs - + This test verifies the fix for issue #18182 where image_tokens were missing from prompt_tokens_details when calling Gemini models with image inputs. - + Example scenario: User sends a text prompt + image, and Gemini generates an image response. The promptTokensDetails should include both TEXT and IMAGE token counts. - + In this test case, candidatesTokenCount is INCLUSIVE of thoughtsTokenCount because: promptTokenCount (533) + candidatesTokenCount (1337) = totalTokenCount (1870) """ @@ -826,31 +848,29 @@ def test_vertex_ai_usage_metadata_with_image_tokens_in_prompt(): "totalTokenCount": 1870, "promptTokensDetails": [ {"modality": "IMAGE", "tokenCount": 527}, - {"modality": "TEXT", "tokenCount": 6} + {"modality": "TEXT", "tokenCount": 6}, ], - "candidatesTokensDetails": [ - {"modality": "IMAGE", "tokenCount": 1120} - ], - "thoughtsTokenCount": 217 + "candidatesTokensDetails": [{"modality": "IMAGE", "tokenCount": 1120}], + "thoughtsTokenCount": 217, } usage_metadata = UsageMetadata(**usage_metadata) result = v._calculate_usage(completion_response={"usageMetadata": usage_metadata}) print("result", result) - + # Verify basic token counts assert result.prompt_tokens == 533 # candidatesTokenCount is INCLUSIVE, so completion_tokens = candidatesTokenCount assert result.completion_tokens == 1337 assert result.total_tokens == 1870 - + # Verify prompt_tokens_details includes both text and image tokens assert result.prompt_tokens_details.text_tokens == 6 assert result.prompt_tokens_details.image_tokens == 527 - + # Verify completion_tokens_details assert result.completion_tokens_details.image_tokens == 1120 assert result.completion_tokens_details.reasoning_tokens == 217 - + # Verify the math: prompt_tokens = text + image # 533 = 6 (text) + 527 (image) assert ( @@ -916,13 +936,17 @@ def test_vertex_ai_map_thinking_param_with_budget_tokens_0(): def test_vertex_ai_map_tools(): v = VertexGeminiConfig() optional_params = {} - tools = v._map_function(value=[{"code_execution": {}}], optional_params=optional_params) + tools = v._map_function( + value=[{"code_execution": {}}], optional_params=optional_params + ) assert len(tools) == 1 assert tools[0]["code_execution"] == {} print(tools) new_optional_params = {} - new_tools = v._map_function(value=[{"codeExecution": {}}], optional_params=new_optional_params) + new_tools = v._map_function( + value=[{"codeExecution": {}}], optional_params=new_optional_params + ) assert len(new_tools) == 1 print("new_tools", new_tools) assert new_tools[0]["code_execution"] == {} @@ -1088,7 +1112,13 @@ def test_vertex_ai_streaming_usage_web_search_calculation(): { "content": {"parts": [{"text": "Hello"}]}, "groundingMetadata": [ - {"webSearchQueries": ["", "What is the capital of France?", "Capital of France"]} + { + "webSearchQueries": [ + "", + "What is the capital of France?", + "Capital of France", + ] + } ], } ], @@ -1432,7 +1462,7 @@ def test_vertex_ai_process_candidates_with_grounding_metadata(): def test_vertex_ai_tool_call_id_format(): """ Test that tool call IDs have the correct format and length. - + The ID should be in format 'call_' + 28 hex characters (total 33 characters). This test verifies the fix for keeping the code line under 40 characters. """ @@ -1449,12 +1479,7 @@ def test_vertex_ai_tool_call_id_format(): "args": {"location": "San Francisco", "unit": "celsius"}, } ), - HttpxPartType( - functionCall={ - "name": "get_time", - "args": {"timezone": "PST"} - } - ), + HttpxPartType(functionCall={"name": "get_time", "args": {"timezone": "PST"}}), ] function, tools, updated_idx = VertexGeminiConfig._transform_parts( @@ -1469,19 +1494,27 @@ def test_vertex_ai_tool_call_id_format(): # Test ID format for both tool calls for tool in tools: tool_id = tool["id"] - + # Should start with 'call_' - assert tool_id.startswith("call_"), f"ID should start with 'call_', got: {tool_id}" - + assert tool_id.startswith( + "call_" + ), f"ID should start with 'call_', got: {tool_id}" + # Should have exactly 33 total characters (call_ + 28 hex chars) - assert len(tool_id) == 33, f"ID should be 33 characters long, got {len(tool_id)}: {tool_id}" - + assert ( + len(tool_id) == 33 + ), f"ID should be 33 characters long, got {len(tool_id)}: {tool_id}" + # The part after 'call_' should be 28 hex characters hex_part = tool_id[5:] # Remove 'call_' prefix - assert len(hex_part) == 28, f"Hex part should be 28 characters, got {len(hex_part)}: {hex_part}" - + assert ( + len(hex_part) == 28 + ), f"Hex part should be 28 characters, got {len(hex_part)}: {hex_part}" + # Should only contain valid hex characters - assert re.match(r'^[0-9a-f]{28}$', hex_part), f"Should contain only lowercase hex chars, got: {hex_part}" + assert re.match( + r"^[0-9a-f]{28}$", hex_part + ), f"Should contain only lowercase hex chars, got: {hex_part}" # Verify IDs are unique assert tools[0]["id"] != tools[1]["id"], "Tool call IDs should be unique" @@ -1496,15 +1529,17 @@ def test_vertex_ai_tool_call_id_format(): ) if test_tools: ids_generated.add(test_tools[0]["id"]) - + # All generated IDs should be unique - assert len(ids_generated) == 10, f"All 10 IDs should be unique, got {len(ids_generated)} unique IDs" + assert ( + len(ids_generated) == 10 + ), f"All 10 IDs should be unique, got {len(ids_generated)} unique IDs" def test_vertex_ai_code_line_length(): """ Test that the specific code line generating tool call IDs is within character limit. - + This is a meta-test to ensure the code change meets the 40-character requirement. """ import inspect @@ -1514,45 +1549,49 @@ def test_vertex_ai_code_line_length(): ) # Get the source code of the _transform_parts method - source_lines = inspect.getsource(VertexGeminiConfig._transform_parts).split('\n') - + source_lines = inspect.getsource(VertexGeminiConfig._transform_parts).split("\n") + # Find the line that generates the ID id_line = None for line in source_lines: - if '"id": f"call_' in line and 'uuid.uuid4().hex[:28]' in line: + if '"id": f"call_' in line and "uuid.uuid4().hex[:28]" in line: id_line = line.strip() # Remove indentation for length check break - + assert id_line is not None, "Could not find the ID generation line in source code" - + # Check that the line is 40 characters or less (excluding indentation) line_length = len(id_line) - assert line_length <= 40, f"ID generation line is {line_length} characters, should be ≤40: {id_line}" - + assert ( + line_length <= 40 + ), f"ID generation line is {line_length} characters, should be ≤40: {id_line}" + # Verify it contains the expected UUID format - assert 'uuid.uuid4().hex[:28]' in id_line, f"Line should contain shortened UUID format: {id_line}" + assert ( + "uuid.uuid4().hex[:28]" in id_line + ), f"Line should contain shortened UUID format: {id_line}" def test_vertex_ai_map_google_maps_tool_simple(): """ Test googleMaps tool transformation without location data. - + Input: value=[{"googleMaps": {"enableWidget": "ENABLE_WIDGET"}}] optional_params={} - + Expected Output: tools=[{"googleMaps": {"enableWidget": "ENABLE_WIDGET"}}] optional_params={} (unchanged) """ v = VertexGeminiConfig() optional_params = {} - + tools = v._map_function( value=[{"googleMaps": {"enableWidget": "ENABLE_WIDGET"}}], - optional_params=optional_params + optional_params=optional_params, ) - + assert len(tools) == 1 assert "googleMaps" in tools[0] assert tools[0]["googleMaps"]["enableWidget"] == "ENABLE_WIDGET" @@ -1563,7 +1602,7 @@ def test_vertex_ai_map_google_maps_tool_with_location(): """ Test googleMaps tool transformation with location data. Verifies latitude/longitude/languageCode are extracted to toolConfig.retrievalConfig. - + Input: value=[{ "googleMaps": { @@ -1574,7 +1613,7 @@ def test_vertex_ai_map_google_maps_tool_with_location(): } }] optional_params={} - + Expected Output: tools=[{ "googleMaps": {"enableWidget": "ENABLE_WIDGET"} @@ -1593,40 +1632,43 @@ def test_vertex_ai_map_google_maps_tool_with_location(): """ v = VertexGeminiConfig() optional_params = {} - + tools = v._map_function( - value=[{ - "googleMaps": { - "enableWidget": "ENABLE_WIDGET", - "latitude": 37.7749, - "longitude": -122.4194, - "languageCode": "en_US" + value=[ + { + "googleMaps": { + "enableWidget": "ENABLE_WIDGET", + "latitude": 37.7749, + "longitude": -122.4194, + "languageCode": "en_US", + } } - }], - optional_params=optional_params + ], + optional_params=optional_params, ) - + assert len(tools) == 1 assert "googleMaps" in tools[0] - + google_maps_tool = tools[0]["googleMaps"] assert google_maps_tool["enableWidget"] == "ENABLE_WIDGET" assert "latitude" not in google_maps_tool assert "longitude" not in google_maps_tool assert "languageCode" not in google_maps_tool - + assert "toolConfig" in optional_params assert "retrievalConfig" in optional_params["toolConfig"] - + retrieval_config = optional_params["toolConfig"]["retrievalConfig"] assert retrieval_config["latLng"]["latitude"] == 37.7749 assert retrieval_config["latLng"]["longitude"] == -122.4194 assert retrieval_config["languageCode"] == "en_US" + def test_vertex_ai_penalty_parameters_validation(): """ Test that penalty parameters are properly validated for different Gemini models. - + This test ensures that: 1. Models that don't support penalty parameters (like preview models) filter them out 2. Models that support penalty parameters include them in the request @@ -1641,14 +1683,19 @@ def test_vertex_ai_penalty_parameters_validation(): for model, should_support in test_cases: # Test _supports_penalty_parameters method - assert v._supports_penalty_parameters(model) == should_support, \ - f"Model {model} penalty support should be {should_support}" + assert ( + v._supports_penalty_parameters(model) == should_support + ), f"Model {model} penalty support should be {should_support}" # Test get_supported_openai_params method supported_params = v.get_supported_openai_params(model) - has_penalty_params = "frequency_penalty" in supported_params and "presence_penalty" in supported_params - assert has_penalty_params == should_support, \ - f"Model {model} should {'include' if should_support else 'exclude'} penalty params in supported list" + has_penalty_params = ( + "frequency_penalty" in supported_params + and "presence_penalty" in supported_params + ) + assert ( + has_penalty_params == should_support + ), f"Model {model} should {'include' if should_support else 'exclude'} penalty params in supported list" # Test parameter mapping for unsupported model model = "gemini-2.5-pro-preview-06-05" @@ -1656,7 +1703,7 @@ def test_vertex_ai_penalty_parameters_validation(): "temperature": 0.7, "frequency_penalty": 0.5, "presence_penalty": 0.3, - "max_tokens": 100 + "max_tokens": 100, } optional_params = {} @@ -1664,12 +1711,16 @@ def test_vertex_ai_penalty_parameters_validation(): non_default_params=non_default_params, optional_params=optional_params, model=model, - drop_params=False + drop_params=False, ) # Penalty parameters should be filtered out for unsupported models - assert "frequency_penalty" not in result, "frequency_penalty should be filtered out for unsupported model" - assert "presence_penalty" not in result, "presence_penalty should be filtered out for unsupported model" + assert ( + "frequency_penalty" not in result + ), "frequency_penalty should be filtered out for unsupported model" + assert ( + "presence_penalty" not in result + ), "presence_penalty should be filtered out for unsupported model" # Other parameters should still be included assert "temperature" in result, "temperature should still be included" @@ -1681,7 +1732,7 @@ def test_vertex_ai_penalty_parameters_validation(): def test_vertex_ai_gemini_3_penalty_parameters_unsupported(): """ Test that penalty parameters are not supported for Gemini 3 models. - + This test ensures that: 1. Gemini 3 models do not support penalty parameters 2. Penalty parameters are excluded from supported params list for Gemini 3 models @@ -1698,22 +1749,25 @@ def test_vertex_ai_gemini_3_penalty_parameters_unsupported(): for model in gemini_3_models: # Test _supports_penalty_parameters method - assert v._supports_penalty_parameters(model) == False, \ - f"Gemini 3 model {model} should not support penalty parameters" + assert ( + v._supports_penalty_parameters(model) == False + ), f"Gemini 3 model {model} should not support penalty parameters" # Test get_supported_openai_params method supported_params = v.get_supported_openai_params(model) - assert "frequency_penalty" not in supported_params, \ - f"frequency_penalty should not be in supported params for {model}" - assert "presence_penalty" not in supported_params, \ - f"presence_penalty should not be in supported params for {model}" + assert ( + "frequency_penalty" not in supported_params + ), f"frequency_penalty should not be in supported params for {model}" + assert ( + "presence_penalty" not in supported_params + ), f"presence_penalty should not be in supported params for {model}" # Test parameter mapping - penalty params should be filtered out non_default_params = { "temperature": 0.7, "frequency_penalty": 0.5, "presence_penalty": 0.3, - "max_tokens": 100 + "max_tokens": 100, } optional_params = {} @@ -1721,39 +1775,46 @@ def test_vertex_ai_gemini_3_penalty_parameters_unsupported(): non_default_params=non_default_params, optional_params=optional_params, model=model, - drop_params=False + drop_params=False, ) # Penalty parameters should be filtered out for Gemini 3 models - assert "frequency_penalty" not in result, \ - f"frequency_penalty should be filtered out for Gemini 3 model {model}" - assert "presence_penalty" not in result, \ - f"presence_penalty should be filtered out for Gemini 3 model {model}" + assert ( + "frequency_penalty" not in result + ), f"frequency_penalty should be filtered out for Gemini 3 model {model}" + assert ( + "presence_penalty" not in result + ), f"presence_penalty should be filtered out for Gemini 3 model {model}" # Other parameters should still be included - assert "temperature" in result, \ - f"temperature should still be included for Gemini 3 model {model}" - assert "max_output_tokens" in result, \ - f"max_output_tokens should still be included for Gemini 3 model {model}" + assert ( + "temperature" in result + ), f"temperature should still be included for Gemini 3 model {model}" + assert ( + "max_output_tokens" in result + ), f"max_output_tokens should still be included for Gemini 3 model {model}" assert result["temperature"] == 0.7 assert result["max_output_tokens"] == 100 # Test that non-Gemini 3 models still support penalty parameters (if they're not in the unsupported list) non_gemini_3_model = "gemini-2.5-pro" - assert v._supports_penalty_parameters(non_gemini_3_model) == True, \ - f"Non-Gemini 3 model {non_gemini_3_model} should support penalty parameters" - + assert ( + v._supports_penalty_parameters(non_gemini_3_model) == True + ), f"Non-Gemini 3 model {non_gemini_3_model} should support penalty parameters" + supported_params = v.get_supported_openai_params(non_gemini_3_model) - assert "frequency_penalty" in supported_params, \ - f"frequency_penalty should be in supported params for {non_gemini_3_model}" - assert "presence_penalty" in supported_params, \ - f"presence_penalty should be in supported params for {non_gemini_3_model}" + assert ( + "frequency_penalty" in supported_params + ), f"frequency_penalty should be in supported params for {non_gemini_3_model}" + assert ( + "presence_penalty" in supported_params + ), f"presence_penalty should be in supported params for {non_gemini_3_model}" def test_vertex_ai_annotation_streaming_events(): """ Test that annotation events are properly emitted during streaming for Vertex AI Gemini. - + This test verifies: 1. Grounding metadata is converted to annotations in streaming chunks 2. Annotations are included in the delta of streaming chunks @@ -1776,7 +1837,7 @@ def test_vertex_ai_annotation_streaming_events(): "groundingMetadata": { "webSearchQueries": ["weather San Francisco today"], "searchEntryPoint": { - "renderedContent": '