From 1b3c8fec83eae73b6caf05d41d1552e5224056f8 Mon Sep 17 00:00:00 2001 From: ".mobo" <1997638+dotmobo@users.noreply.github.com> Date: Fri, 16 Jan 2026 16:27:09 +0100 Subject: [PATCH 01/77] put logfile and pidfile in /tmp to avoid permission denied error in non root environment (#17267) --- docker/supervisord.conf | 2 ++ 1 file changed, 2 insertions(+) diff --git a/docker/supervisord.conf b/docker/supervisord.conf index c6855fe652b..877335804fe 100644 --- a/docker/supervisord.conf +++ b/docker/supervisord.conf @@ -1,6 +1,8 @@ [supervisord] nodaemon=true loglevel=info +logfile=/tmp/supervisord.log +pidfile=/tmp/supervisord.pid [group:litellm] programs=main,health From 2a0f87bde048e406e6dbf1d5465e77d7e0282bc3 Mon Sep 17 00:00:00 2001 From: Dushyant <157586156+dushyantzz@users.noreply.github.com> Date: Sat, 17 Jan 2026 01:00:16 +0530 Subject: [PATCH 02/77] Fix audio cost per second override (#19158) --- docs/my-website/docs/sdk_custom_pricing.md | 5 +++- litellm/llms/openai/cost_calculation.py | 22 ++++++-------- tests/test_litellm/test_cost_calculator.py | 34 ++++++++++++++++++++++ 3 files changed, 47 insertions(+), 14 deletions(-) diff --git a/docs/my-website/docs/sdk_custom_pricing.md b/docs/my-website/docs/sdk_custom_pricing.md index c8577115109..ac956db472a 100644 --- a/docs/my-website/docs/sdk_custom_pricing.md +++ b/docs/my-website/docs/sdk_custom_pricing.md @@ -2,7 +2,10 @@ Register custom pricing for sagemaker completion model. -For cost per second pricing, you **just** need to register `input_cost_per_second`. +For cost per second pricing, register `input_cost_per_second`. If your provider +charges for audio output duration (e.g., TTS), also set `output_cost_per_second`. +Values of `0` are treated as not billable, so `output_cost_per_second: 0` will +not override `input_cost_per_second`. ```python # !pip install boto3 diff --git a/litellm/llms/openai/cost_calculation.py b/litellm/llms/openai/cost_calculation.py index e5349db3af7..339ce9caedd 100644 --- a/litellm/llms/openai/cost_calculation.py +++ b/litellm/llms/openai/cost_calculation.py @@ -105,25 +105,21 @@ def cost_per_second( prompt_cost = 0.0 completion_cost = 0.0 ## Speech / Audio cost calculation - if ( - "output_cost_per_second" in model_info - and model_info["output_cost_per_second"] is not None - ): + output_cost_per_second = model_info.get("output_cost_per_second") + input_cost_per_second = model_info.get("input_cost_per_second") + + if output_cost_per_second is not None and output_cost_per_second > 0: verbose_logger.debug( - f"For model={model} - output_cost_per_second: {model_info.get('output_cost_per_second')}; duration: {duration}" + f"For model={model} - output_cost_per_second: {output_cost_per_second}; duration: {duration}" ) ## COST PER SECOND ## - completion_cost = model_info["output_cost_per_second"] * duration - elif ( - "input_cost_per_second" in model_info - and model_info["input_cost_per_second"] is not None - ): + completion_cost = output_cost_per_second * duration + if input_cost_per_second is not None and input_cost_per_second > 0: verbose_logger.debug( - f"For model={model} - input_cost_per_second: {model_info.get('input_cost_per_second')}; duration: {duration}" + f"For model={model} - input_cost_per_second: {input_cost_per_second}; duration: {duration}" ) ## COST PER SECOND ## - prompt_cost = model_info["input_cost_per_second"] * duration - completion_cost = 0.0 + prompt_cost = input_cost_per_second * duration return prompt_cost, completion_cost diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 4d6599fc1b5..42bcfca2311 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -192,6 +192,40 @@ def test_transcription_cost_falls_back_to_duration(): assert pytest.approx(cost, rel=1e-6) == expected_cost +def test_transcription_cost_prefers_input_when_output_zero(monkeypatch): + from litellm import completion_cost + + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + model_name = "custom-whisper-input-only" + custom_model_info = { + "input_cost_per_second": 0.00005, + "output_cost_per_second": 0.0, + "litellm_provider": "openai", + "mode": "audio_transcription", + "supported_endpoints": ["/v1/audio/transcriptions"], + } + monkeypatch.setattr( + litellm, + "model_cost", + {**litellm.model_cost, model_name: custom_model_info}, + ) + + response = TranscriptionResponse(text="demo text") + response.duration = 300.0 + + cost = completion_cost( + completion_response=response, + model=model_name, + custom_llm_provider="openai", + call_type="atranscription", + ) + + expected_cost = 300.0 * 0.00005 + assert pytest.approx(cost, rel=1e-6) == expected_cost + + def test_handle_realtime_stream_cost_calculation(): from litellm.cost_calculator import RealtimeAPITokenUsageProcessor From ce1a2c3209ef46e6f48565e561739cb2579dcbb1 Mon Sep 17 00:00:00 2001 From: Ostap Bodnar <171420608+obod-mpw@users.noreply.github.com> Date: Sat, 17 Jan 2026 01:18:46 +0200 Subject: [PATCH 03/77] add openai/dall-e base pricing entries (#19133) --- model_prices_and_context_window.json | 18 ++++++++++++++++++ 1 file changed, 18 insertions(+) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index a130aefa5de..ec55c251a70 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -7826,6 +7826,24 @@ "supports_tool_choice": true, "supports_vision": true }, + "dall-e-2": { + "input_cost_per_image": 0.02, + "litellm_provider": "openai", + "mode": "image_generation", + "supported_endpoints": [ + "/v1/images/generations", + "/v1/images/edits", + "/v1/images/variations" + ] + }, + "dall-e-3": { + "input_cost_per_image": 0.04, + "litellm_provider": "openai", + "mode": "image_generation", + "supported_endpoints": [ + "/v1/images/generations" + ] + }, "deepseek-chat": { "cache_read_input_token_cost": 2.8e-08, "input_cost_per_token": 2.8e-07, From 5a9f6e90cfa7063559bb99a68ed981b1a84cde0a Mon Sep 17 00:00:00 2001 From: Ryan Malloy Date: Fri, 16 Jan 2026 16:24:28 -0700 Subject: [PATCH 04/77] fix(tools): prevent OOM with nested $defs in tool schemas (#19098) (#19112) --- .../prompt_templates/factory.py | 7 +- litellm/llms/vertex_ai/common_utils.py | 7 +- ...llm_core_utils_prompt_templates_factory.py | 131 ++++++++++++++++++ 3 files changed, 139 insertions(+), 6 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 89a708077f3..6eb4ee74490 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -4468,9 +4468,10 @@ def _bedrock_tools_pt(tools: List) -> List[BedrockToolBlock]: defs = parameters.pop("$defs", {}) defs_copy = copy.deepcopy(defs) - # flatten the defs - for _, value in defs_copy.items(): - unpack_defs(value, defs_copy) + # Expand $ref references in parameters using the definitions + # Note: We don't pre-flatten defs as that causes exponential memory growth + # with circular references (see issue #19098). unpack_defs handles nested + # refs recursively and correctly detects/skips circular references. unpack_defs(parameters, defs_copy) tool_input_schema = BedrockToolInputSchemaBlock( json=BedrockToolJsonSchemaBlock( diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index 1864ef734c0..77c85742284 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -453,9 +453,10 @@ def _build_vertex_schema(parameters: dict, add_property_ordering: bool = False): valid_schema_fields = set(get_type_hints(Schema).keys()) defs = parameters.pop("$defs", {}) - # flatten the defs - for name, value in defs.items(): - unpack_defs(value, defs) + # Expand $ref references in parameters using the definitions + # Note: We don't pre-flatten defs as that causes exponential memory growth + # with circular references (see issue #19098). unpack_defs handles nested + # refs recursively and correctly detects/skips circular references. unpack_defs(parameters, defs) # 5. Nullable fields: diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 42a2b5d0971..a22fe13798f 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -1392,3 +1392,134 @@ def test_anthropic_messages_pt_server_tool_use_passthrough(): b for b in assistant_msg["content"] if b.get("type") == "text" ) assert text_block["text"] == "I found the time tool. How can I help you?" + + +def test_bedrock_tools_unpack_defs_no_oom_with_nested_refs(): + """ + Regression test for issue #19098: unpack_defs() causes OOM with nested tool schemas. + + The old implementation had a "flatten defs" loop that would pre-expand each def + using unpack_defs(), but since defs often reference each other, each subsequent + call would copy already-expanded content, causing exponential memory growth. + + This test creates a schema with multiple nested $defs that reference each other + to verify the fix prevents memory explosion while still correctly resolving refs. + """ + import sys + import copy + + from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_tools_pt + + # Schema with multiple nested $defs that reference each other + # This pattern would cause OOM with the old "flatten defs" loop + complex_nested_schema = { + "type": "object", + "properties": { + "query": {"$ref": "#/$defs/Expression"}, + }, + "$defs": { + "Expression": { + "type": "object", + "properties": { + "type": {"type": "string", "enum": ["and", "or", "not", "comparison"]}, + "left": {"$ref": "#/$defs/Operand"}, + "right": {"$ref": "#/$defs/Operand"}, + "operator": {"$ref": "#/$defs/Operator"}, + }, + }, + "Operand": { + "type": "object", + "anyOf": [ + {"$ref": "#/$defs/Literal"}, + {"$ref": "#/$defs/FieldRef"}, + {"$ref": "#/$defs/Expression"}, # Circular: Operand -> Expression -> Operand + ], + }, + "Literal": { + "type": "object", + "properties": { + "type": {"type": "string", "const": "literal"}, + "value": {"$ref": "#/$defs/LiteralValue"}, + }, + }, + "LiteralValue": { + "oneOf": [ + {"type": "string"}, + {"type": "number"}, + {"type": "boolean"}, + {"type": "null"}, + ], + }, + "FieldRef": { + "type": "object", + "properties": { + "type": {"type": "string", "const": "field"}, + "name": {"type": "string"}, + "table": {"$ref": "#/$defs/TableRef"}, + }, + }, + "TableRef": { + "type": "object", + "properties": { + "name": {"type": "string"}, + "alias": {"type": "string"}, + }, + }, + "Operator": { + "type": "string", + "enum": ["=", "!=", "<", ">", "<=", ">=", "LIKE", "IN"], + }, + }, + } + + tools = [ + { + "type": "function", + "function": { + "name": "execute_query", + "description": "Execute a query with complex expressions", + "parameters": complex_nested_schema, + }, + } + ] + + # Measure initial size + def get_size(obj, seen=None): + size = sys.getsizeof(obj) + if seen is None: + seen = set() + obj_id = id(obj) + if obj_id in seen: + return 0 + seen.add(obj_id) + if isinstance(obj, dict): + size += sum([get_size(v, seen) for v in obj.values()]) + size += sum([get_size(k, seen) for k in obj.keys()]) + elif hasattr(obj, "__iter__") and not isinstance(obj, (str, bytes, bytearray)): + size += sum([get_size(i, seen) for i in obj]) + return size + + initial_size = get_size(tools) + + # Process through _bedrock_tools_pt - this should complete without OOM + tools_copy = copy.deepcopy(tools) + result = _bedrock_tools_pt(tools=tools_copy) + + final_size = get_size(result) + + # The expansion factor should be reasonable (< 100x), not exponential (35000x as in #19098) + expansion_factor = final_size / initial_size + assert expansion_factor < 100, ( + f"Memory expansion factor {expansion_factor:.1f}x is too high. " + f"Initial: {initial_size} bytes, Final: {final_size} bytes" + ) + + # Verify the result is valid Bedrock tools format + assert isinstance(result, list) + assert len(result) == 1 + assert "toolSpec" in result[0] + assert result[0]["toolSpec"]["name"] == "execute_query" + + # Verify $defs have been removed (Bedrock doesn't support them) + tool_schema = result[0]["toolSpec"].get("inputSchema", {}).get("json", {}) + assert "$defs" not in tool_schema, "$defs should be removed after expansion" From 866bd4674837d64462c70e63fa13343c9ec6a1fb Mon Sep 17 00:00:00 2001 From: Harshit Jain <48647625+Harshit28j@users.noreply.github.com> Date: Sat, 17 Jan 2026 04:56:52 +0530 Subject: [PATCH 05/77] chore: resolve ModuleNotFoundError for Microsoft Foundry Agents (#18991) --- litellm/a2a_protocol/main.py | 44 ++-- .../proxy/agent_endpoints/a2a_endpoints.py | 88 ++++++-- poetry.lock | 200 ++++++++++-------- pyproject.toml | 16 +- requirements.txt | 13 +- .../agent_endpoints/test_a2a_endpoints.py | 53 +++-- 6 files changed, 257 insertions(+), 157 deletions(-) diff --git a/litellm/a2a_protocol/main.py b/litellm/a2a_protocol/main.py index 167aad7959a..2d36dbeacda 100644 --- a/litellm/a2a_protocol/main.py +++ b/litellm/a2a_protocol/main.py @@ -113,7 +113,9 @@ def _get_a2a_model_info(a2a_client: Any, kwargs: Dict[str, Any]) -> str: litellm_logging_obj.model = model litellm_logging_obj.custom_llm_provider = custom_llm_provider litellm_logging_obj.model_call_details["model"] = model - litellm_logging_obj.model_call_details["custom_llm_provider"] = custom_llm_provider + litellm_logging_obj.model_call_details[ + "custom_llm_provider" + ] = custom_llm_provider return agent_name @@ -197,7 +199,11 @@ async def asend_message( ) # Extract params from request - params = request.params.model_dump(mode="json") if hasattr(request.params, "model_dump") else dict(request.params) + params = ( + request.params.model_dump(mode="json") + if hasattr(request.params, "model_dump") + else dict(request.params) + ) response_dict = await A2ACompletionBridgeHandler.handle_non_streaming( request_id=str(request.id), @@ -216,7 +222,9 @@ async def asend_message( # Create A2A client if not provided but api_base is available if a2a_client is None: if api_base is None: - raise ValueError("Either a2a_client or api_base is required for standard A2A flow") + raise ValueError( + "Either a2a_client or api_base is required for standard A2A flow" + ) a2a_client = await create_a2a_client(base_url=api_base) # Type assertion: a2a_client is guaranteed to be non-None here @@ -235,7 +243,11 @@ async def asend_message( # Calculate token usage from request and response response_dict = a2a_response.model_dump(mode="json", exclude_none=True) - prompt_tokens, completion_tokens, _ = A2ARequestUtils.calculate_usage_from_request_response( + ( + prompt_tokens, + completion_tokens, + _, + ) = A2ARequestUtils.calculate_usage_from_request_response( request=request, response_dict=response_dict, ) @@ -280,7 +292,9 @@ def send_message( if loop is not None: return asend_message(a2a_client=a2a_client, request=request, **kwargs) else: - return asyncio.run(asend_message(a2a_client=a2a_client, request=request, **kwargs)) + return asyncio.run( + asend_message(a2a_client=a2a_client, request=request, **kwargs) + ) async def asend_message_streaming( @@ -347,7 +361,11 @@ async def asend_message_streaming( ) # Extract params from request - params = request.params.model_dump(mode="json") if hasattr(request.params, "model_dump") else dict(request.params) + params = ( + request.params.model_dump(mode="json") + if hasattr(request.params, "model_dump") + else dict(request.params) + ) async for chunk in A2ACompletionBridgeHandler.handle_streaming( request_id=str(request.id), @@ -365,7 +383,9 @@ async def asend_message_streaming( # Create A2A client if not provided but api_base is available if a2a_client is None: if api_base is None: - raise ValueError("Either a2a_client or api_base is required for standard A2A flow") + raise ValueError( + "Either a2a_client or api_base is required for standard A2A flow" + ) a2a_client = await create_a2a_client(base_url=api_base) # Type assertion: a2a_client is guaranteed to be non-None here @@ -378,7 +398,9 @@ async def asend_message_streaming( stream = a2a_client.send_message_streaming(request) # Build logging object for streaming completion callbacks - agent_card = getattr(a2a_client, "_litellm_agent_card", None) or getattr(a2a_client, "agent_card", None) + agent_card = getattr(a2a_client, "_litellm_agent_card", None) or getattr( + a2a_client, "agent_card", None + ) agent_name = getattr(agent_card, "name", "unknown") if agent_card else "unknown" model = f"a2a_agent/{agent_name}" @@ -456,7 +478,7 @@ async def create_a2a_client( if not A2A_SDK_AVAILABLE: raise ImportError( "The 'a2a' package is required for A2A agent invocation. " - "Install it with: pip install a2a" + "Install it with: pip install a2a-sdk" ) verbose_logger.info(f"Creating A2A client for {base_url}") @@ -512,7 +534,7 @@ async def aget_agent_card( if not A2A_SDK_AVAILABLE: raise ImportError( "The 'a2a' package is required for A2A agent invocation. " - "Install it with: pip install a2a" + "Install it with: pip install a2a-sdk" ) verbose_logger.info(f"Fetching agent card from {base_url}") @@ -534,5 +556,3 @@ async def aget_agent_card( f"Fetched agent card: {agent_card.name if hasattr(agent_card, 'name') else 'unknown'}" ) return agent_card - - diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index c2d53b40b7b..a21f3291d31 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -55,9 +55,30 @@ async def _handle_stream_message( proxy_server_request: Optional[dict] = None, ) -> StreamingResponse: """Handle message/stream method via SDK functions.""" - from a2a.types import MessageSendParams, SendStreamingMessageRequest - from litellm.a2a_protocol import asend_message_streaming + from litellm.a2a_protocol.main import A2A_SDK_AVAILABLE + + # Check is handled in invoke_agent_a2a, but if called directly: + if not A2A_SDK_AVAILABLE: + # Return a streaming response that yields an error + async def _error_stream(): + yield json.dumps( + { + "jsonrpc": "2.0", + "id": request_id, + "error": { + "code": -32603, + "message": "Server error: 'a2a' package not installed", + }, + } + ) + "\n" + + return StreamingResponse(_error_stream(), media_type="application/x-ndjson") + + from a2a.types import ( + MessageSendParams, + SendStreamingMessageRequest, + ) async def stream_response(): try: @@ -75,16 +96,20 @@ async def _handle_stream_message( ): # Chunk may be dict or object depending on bridge vs standard path if hasattr(chunk, "model_dump"): - yield json.dumps(chunk.model_dump(mode="json", exclude_none=True)) + "\n" + yield json.dumps( + chunk.model_dump(mode="json", exclude_none=True) + ) + "\n" else: yield json.dumps(chunk) + "\n" except Exception as e: verbose_proxy_logger.exception(f"Error streaming A2A response: {e}") - yield json.dumps({ - "jsonrpc": "2.0", - "id": request_id, - "error": {"code": -32603, "message": f"Streaming error: {str(e)}"}, - }) + "\n" + yield json.dumps( + { + "jsonrpc": "2.0", + "id": request_id, + "error": {"code": -32603, "message": f"Streaming error: {str(e)}"}, + } + ) + "\n" return StreamingResponse(stream_response(), media_type="application/x-ndjson") @@ -169,9 +194,8 @@ async def invoke_agent_a2a( - message/send: Send a message and get a response - message/stream: Send a message and stream the response """ - from a2a.types import MessageSendParams, SendMessageRequest - from litellm.a2a_protocol import asend_message + from litellm.a2a_protocol.main import A2A_SDK_AVAILABLE from litellm.proxy.agent_endpoints.auth.agent_permission_handler import ( AgentRequestHandler, ) @@ -189,16 +213,28 @@ async def invoke_agent_a2a( # Validate JSON-RPC format if body.get("jsonrpc") != "2.0": - return _jsonrpc_error(body.get("id"), -32600, "Invalid Request: jsonrpc must be '2.0'") + return _jsonrpc_error( + body.get("id"), -32600, "Invalid Request: jsonrpc must be '2.0'" + ) request_id = body.get("id") method = body.get("method") params = body.get("params", {}) + if not A2A_SDK_AVAILABLE: + return _jsonrpc_error( + request_id, + -32603, + "Server error: 'a2a' package not installed. Please install 'a2a-sdk'.", + 500, + ) + # Find the agent agent = _get_agent(agent_id) if agent is None: - return _jsonrpc_error(request_id, -32000, f"Agent '{agent_id}' not found", 404) + return _jsonrpc_error( + request_id, -32000, f"Agent '{agent_id}' not found", 404 + ) is_allowed = await AgentRequestHandler.is_agent_allowed( agent_id=agent.agent_id, @@ -213,23 +249,29 @@ async def invoke_agent_a2a( # Get backend URL and agent name agent_url = agent.agent_card_params.get("url") agent_name = agent.agent_card_params.get("name", agent_id) - + # Get litellm_params (may include custom_llm_provider for completion bridge) litellm_params = agent.litellm_params or {} custom_llm_provider = litellm_params.get("custom_llm_provider") - + # URL is required unless using completion bridge with a provider that derives endpoint from model # (e.g., bedrock/agentcore derives endpoint from ARN in model string) if not agent_url and not custom_llm_provider: - return _jsonrpc_error(request_id, -32000, f"Agent '{agent_id}' has no URL configured", 500) + return _jsonrpc_error( + request_id, -32000, f"Agent '{agent_id}' has no URL configured", 500 + ) - verbose_proxy_logger.info(f"Proxying A2A request to agent '{agent_id}' at {agent_url or 'completion-bridge'}") + verbose_proxy_logger.info( + f"Proxying A2A request to agent '{agent_id}' at {agent_url or 'completion-bridge'}" + ) # Set up data dict for litellm processing - body.update({ - "model": f"a2a_agent/{agent_name}", - "custom_llm_provider": "a2a_agent", - }) + body.update( + { + "model": f"a2a_agent/{agent_name}", + "custom_llm_provider": "a2a_agent", + } + ) # Add litellm data (user_api_key, user_id, team_id, etc.) data = await add_litellm_data_to_request( @@ -243,6 +285,8 @@ async def invoke_agent_a2a( # Route through SDK functions if method == "message/send": + from a2a.types import MessageSendParams, SendMessageRequest + a2a_request = SendMessageRequest( id=request_id, params=MessageSendParams(**params), @@ -255,7 +299,9 @@ async def invoke_agent_a2a( metadata=data.get("metadata", {}), proxy_server_request=data.get("proxy_server_request"), ) - return JSONResponse(content=response.model_dump(mode="json", exclude_none=True)) + return JSONResponse( + content=response.model_dump(mode="json", exclude_none=True) + ) elif method == "message/stream": return await _handle_stream_message( diff --git a/poetry.lock b/poetry.lock index 3bafdb157ca..fd6b28533e7 100644 --- a/poetry.lock +++ b/poetry.lock @@ -1,4 +1,36 @@ -# This file is automatically @generated by Poetry 2.2.0 and should not be changed by hand. +# This file is automatically @generated by Poetry 2.2.1 and should not be changed by hand. + +[[package]] +name = "a2a-sdk" +version = "0.3.22" +description = "A2A Python SDK" +optional = true +python-versions = ">=3.10" +groups = ["main"] +markers = "python_version >= \"3.10\" and extra == \"extra-proxy\"" +files = [ + {file = "a2a_sdk-0.3.22-py3-none-any.whl", hash = "sha256:b98701135bb90b0ff85d35f31533b6b7a299bf810658c1c65f3814a6c15ea385"}, + {file = "a2a_sdk-0.3.22.tar.gz", hash = "sha256:77a5694bfc4f26679c11b70c7f1062522206d430b34bc1215cfbb1eba67b7e7d"}, +] + +[package.dependencies] +google-api-core = ">=1.26.0" +httpx = ">=0.28.1" +httpx-sse = ">=0.4.0" +protobuf = ">=5.29.5" +pydantic = ">=2.11.3" + +[package.extras] +all = ["cryptography (>=43.0.0)", "fastapi (>=0.115.2)", "grpcio (>=1.60)", "grpcio-reflection (>=1.7.0)", "grpcio-tools (>=1.60)", "opentelemetry-api (>=1.33.0)", "opentelemetry-sdk (>=1.33.0)", "pyjwt (>=2.0.0)", "sqlalchemy[aiomysql,asyncio] (>=2.0.0)", "sqlalchemy[aiosqlite,asyncio] (>=2.0.0)", "sqlalchemy[asyncio,postgresql-asyncpg] (>=2.0.0)", "sse-starlette", "starlette"] +encryption = ["cryptography (>=43.0.0)"] +grpc = ["grpcio (>=1.60)", "grpcio-reflection (>=1.7.0)", "grpcio-tools (>=1.60)"] +http-server = ["fastapi (>=0.115.2)", "sse-starlette", "starlette"] +mysql = ["sqlalchemy[aiomysql,asyncio] (>=2.0.0)"] +postgresql = ["sqlalchemy[asyncio,postgresql-asyncpg] (>=2.0.0)"] +signing = ["pyjwt (>=2.0.0)"] +sql = ["sqlalchemy[aiomysql,asyncio] (>=2.0.0)", "sqlalchemy[aiosqlite,asyncio] (>=2.0.0)", "sqlalchemy[asyncio,postgresql-asyncpg] (>=2.0.0)"] +sqlite = ["sqlalchemy[aiosqlite,asyncio] (>=2.0.0)"] +telemetry = ["opentelemetry-api (>=1.33.0)", "opentelemetry-sdk (>=1.33.0)"] [[package]] name = "aiofiles" @@ -1268,25 +1300,6 @@ dev = ["autoflake", "black", "build", "databricks-connect", "httpx", "ipython", notebook = ["ipython (>=8,<10)", "ipywidgets (>=8,<9)"] openai = ["httpx", "langchain-openai ; python_version > \"3.7\"", "openai"] -[[package]] -name = "deprecated" -version = "1.3.1" -description = "Python @deprecated decorator to deprecate old python classes, functions or methods." -optional = false -python-versions = "!=3.0.*,!=3.1.*,!=3.2.*,!=3.3.*,>=2.7" -groups = ["main", "dev", "proxy-dev"] -files = [ - {file = "deprecated-1.3.1-py2.py3-none-any.whl", hash = "sha256:597bfef186b6f60181535a29fbe44865ce137a5079f295b479886c82729d5f3f"}, - {file = "deprecated-1.3.1.tar.gz", hash = "sha256:b1b50e0ff0c1fddaa5708a2c6b0a6588bb09b892825ab2b214ac9ea9d92a5223"}, -] -markers = {main = "python_version >= \"3.10\""} - -[package.dependencies] -wrapt = ">=1.10,<3" - -[package.extras] -dev = ["PyTest", "PyTest-Cov", "bump2version (<1)", "setuptools ; python_version >= \"3.12\"", "tox"] - [[package]] name = "diskcache" version = "5.6.3" @@ -2521,7 +2534,7 @@ description = "Consume Server-Sent Event (SSE) messages with HTTPX." optional = true python-versions = ">=3.9" groups = ["main"] -markers = "python_version >= \"3.10\" and extra == \"proxy\"" +markers = "python_version >= \"3.10\" and (extra == \"proxy\" or extra == \"extra-proxy\")" files = [ {file = "httpx_sse-0.4.3-py3-none-any.whl", hash = "sha256:0ac1c9fe3c0afad2e0ebb25a934a59f4c7823b60792691f779fad2c5568830fc"}, {file = "httpx_sse-0.4.3.tar.gz", hash = "sha256:9b1ed0127459a66014aec3c56bebd93da3c1bc8bb6618c8082039a44889a755d"}, @@ -4036,143 +4049,153 @@ voice-helpers = ["numpy (>=2.0.2)", "sounddevice (>=0.5.1)"] [[package]] name = "opentelemetry-api" -version = "1.25.0" +version = "1.39.1" description = "OpenTelemetry Python API" optional = false -python-versions = ">=3.8" +python-versions = ">=3.9" groups = ["main", "dev", "proxy-dev"] files = [ - {file = "opentelemetry_api-1.25.0-py3-none-any.whl", hash = "sha256:757fa1aa020a0f8fa139f8959e53dec2051cc26b832e76fa839a6d76ecefd737"}, - {file = "opentelemetry_api-1.25.0.tar.gz", hash = "sha256:77c4985f62f2614e42ce77ee4c9da5fa5f0bc1e1821085e9a47533a9323ae869"}, + {file = "opentelemetry_api-1.39.1-py3-none-any.whl", hash = "sha256:2edd8463432a7f8443edce90972169b195e7d6a05500cd29e6d13898187c9950"}, + {file = "opentelemetry_api-1.39.1.tar.gz", hash = "sha256:fbde8c80e1b937a2c61f20347e91c0c18a1940cecf012d62e65a7caf08967c9c"}, ] markers = {main = "python_version >= \"3.10\""} [package.dependencies] -deprecated = ">=1.2.6" -importlib-metadata = ">=6.0,<=7.1" +importlib-metadata = ">=6.0,<8.8.0" +typing-extensions = ">=4.5.0" [[package]] name = "opentelemetry-exporter-otlp" -version = "1.25.0" +version = "1.39.1" description = "OpenTelemetry Collector Exporters" optional = false -python-versions = ">=3.8" +python-versions = ">=3.9" groups = ["dev", "proxy-dev"] files = [ - {file = "opentelemetry_exporter_otlp-1.25.0-py3-none-any.whl", hash = "sha256:d67a831757014a3bc3174e4cd629ae1493b7ba8d189e8a007003cacb9f1a6b60"}, - {file = "opentelemetry_exporter_otlp-1.25.0.tar.gz", hash = "sha256:ce03199c1680a845f82e12c0a6a8f61036048c07ec7a0bd943142aca8fa6ced0"}, + {file = "opentelemetry_exporter_otlp-1.39.1-py3-none-any.whl", hash = "sha256:68ae69775291f04f000eb4b698ff16ff685fdebe5cb52871bc4e87938a7b00fe"}, + {file = "opentelemetry_exporter_otlp-1.39.1.tar.gz", hash = "sha256:7cf7470e9fd0060c8a38a23e4f695ac686c06a48ad97f8d4867bc9b420180b9c"}, ] [package.dependencies] -opentelemetry-exporter-otlp-proto-grpc = "1.25.0" -opentelemetry-exporter-otlp-proto-http = "1.25.0" +opentelemetry-exporter-otlp-proto-grpc = "1.39.1" +opentelemetry-exporter-otlp-proto-http = "1.39.1" [[package]] name = "opentelemetry-exporter-otlp-proto-common" -version = "1.25.0" +version = "1.39.1" description = "OpenTelemetry Protobuf encoding" optional = false -python-versions = ">=3.8" +python-versions = ">=3.9" groups = ["dev", "proxy-dev"] files = [ - {file = "opentelemetry_exporter_otlp_proto_common-1.25.0-py3-none-any.whl", hash = "sha256:15637b7d580c2675f70246563363775b4e6de947871e01d0f4e3881d1848d693"}, - {file = "opentelemetry_exporter_otlp_proto_common-1.25.0.tar.gz", hash = "sha256:c93f4e30da4eee02bacd1e004eb82ce4da143a2f8e15b987a9f603e0a85407d3"}, + {file = "opentelemetry_exporter_otlp_proto_common-1.39.1-py3-none-any.whl", hash = "sha256:08f8a5862d64cc3435105686d0216c1365dc5701f86844a8cd56597d0c764fde"}, + {file = "opentelemetry_exporter_otlp_proto_common-1.39.1.tar.gz", hash = "sha256:763370d4737a59741c89a67b50f9e39271639ee4afc999dadfe768541c027464"}, ] [package.dependencies] -opentelemetry-proto = "1.25.0" +opentelemetry-proto = "1.39.1" [[package]] name = "opentelemetry-exporter-otlp-proto-grpc" -version = "1.25.0" +version = "1.39.1" description = "OpenTelemetry Collector Protobuf over gRPC Exporter" optional = false -python-versions = ">=3.8" +python-versions = ">=3.9" groups = ["dev", "proxy-dev"] files = [ - {file = "opentelemetry_exporter_otlp_proto_grpc-1.25.0-py3-none-any.whl", hash = "sha256:3131028f0c0a155a64c430ca600fd658e8e37043cb13209f0109db5c1a3e4eb4"}, - {file = "opentelemetry_exporter_otlp_proto_grpc-1.25.0.tar.gz", hash = "sha256:c0b1661415acec5af87625587efa1ccab68b873745ca0ee96b69bb1042087eac"}, + {file = "opentelemetry_exporter_otlp_proto_grpc-1.39.1-py3-none-any.whl", hash = "sha256:fa1c136a05c7e9b4c09f739469cbdb927ea20b34088ab1d959a849b5cc589c18"}, + {file = "opentelemetry_exporter_otlp_proto_grpc-1.39.1.tar.gz", hash = "sha256:772eb1c9287485d625e4dbe9c879898e5253fea111d9181140f51291b5fec3ad"}, ] [package.dependencies] -deprecated = ">=1.2.6" -googleapis-common-protos = ">=1.52,<2.0" -grpcio = ">=1.0.0,<2.0.0" +googleapis-common-protos = ">=1.57,<2.0" +grpcio = [ + {version = ">=1.63.2,<2.0.0", markers = "python_version < \"3.13\""}, + {version = ">=1.66.2,<2.0.0", markers = "python_version >= \"3.13\""}, +] opentelemetry-api = ">=1.15,<2.0" -opentelemetry-exporter-otlp-proto-common = "1.25.0" -opentelemetry-proto = "1.25.0" -opentelemetry-sdk = ">=1.25.0,<1.26.0" +opentelemetry-exporter-otlp-proto-common = "1.39.1" +opentelemetry-proto = "1.39.1" +opentelemetry-sdk = ">=1.39.1,<1.40.0" +typing-extensions = ">=4.6.0" + +[package.extras] +gcp-auth = ["opentelemetry-exporter-credential-provider-gcp (>=0.59b0)"] [[package]] name = "opentelemetry-exporter-otlp-proto-http" -version = "1.25.0" +version = "1.39.1" description = "OpenTelemetry Collector Protobuf over HTTP Exporter" optional = false -python-versions = ">=3.8" +python-versions = ">=3.9" groups = ["dev", "proxy-dev"] files = [ - {file = "opentelemetry_exporter_otlp_proto_http-1.25.0-py3-none-any.whl", hash = "sha256:2eca686ee11b27acd28198b3ea5e5863a53d1266b91cda47c839d95d5e0541a6"}, - {file = "opentelemetry_exporter_otlp_proto_http-1.25.0.tar.gz", hash = "sha256:9f8723859e37c75183ea7afa73a3542f01d0fd274a5b97487ea24cb683d7d684"}, + {file = "opentelemetry_exporter_otlp_proto_http-1.39.1-py3-none-any.whl", hash = "sha256:d9f5207183dd752a412c4cd564ca8875ececba13be6e9c6c370ffb752fd59985"}, + {file = "opentelemetry_exporter_otlp_proto_http-1.39.1.tar.gz", hash = "sha256:31bdab9745c709ce90a49a0624c2bd445d31a28ba34275951a6a362d16a0b9cb"}, ] [package.dependencies] -deprecated = ">=1.2.6" googleapis-common-protos = ">=1.52,<2.0" opentelemetry-api = ">=1.15,<2.0" -opentelemetry-exporter-otlp-proto-common = "1.25.0" -opentelemetry-proto = "1.25.0" -opentelemetry-sdk = ">=1.25.0,<1.26.0" +opentelemetry-exporter-otlp-proto-common = "1.39.1" +opentelemetry-proto = "1.39.1" +opentelemetry-sdk = ">=1.39.1,<1.40.0" requests = ">=2.7,<3.0" +typing-extensions = ">=4.5.0" + +[package.extras] +gcp-auth = ["opentelemetry-exporter-credential-provider-gcp (>=0.59b0)"] [[package]] name = "opentelemetry-proto" -version = "1.25.0" +version = "1.39.1" description = "OpenTelemetry Python Proto" optional = false -python-versions = ">=3.8" +python-versions = ">=3.9" groups = ["main", "dev", "proxy-dev"] files = [ - {file = "opentelemetry_proto-1.25.0-py3-none-any.whl", hash = "sha256:f07e3341c78d835d9b86665903b199893befa5e98866f63d22b00d0b7ca4972f"}, - {file = "opentelemetry_proto-1.25.0.tar.gz", hash = "sha256:35b6ef9dc4a9f7853ecc5006738ad40443701e52c26099e197895cbda8b815a3"}, + {file = "opentelemetry_proto-1.39.1-py3-none-any.whl", hash = "sha256:22cdc78efd3b3765d09e68bfbd010d4fc254c9818afd0b6b423387d9dee46007"}, + {file = "opentelemetry_proto-1.39.1.tar.gz", hash = "sha256:6c8e05144fc0d3ed4d22c2289c6b126e03bcd0e6a7da0f16cedd2e1c2772e2c8"}, ] markers = {main = "python_version >= \"3.10\" and extra == \"mlflow\""} [package.dependencies] -protobuf = ">=3.19,<5.0" +protobuf = ">=5.0,<7.0" [[package]] name = "opentelemetry-sdk" -version = "1.25.0" +version = "1.39.1" description = "OpenTelemetry Python SDK" optional = false -python-versions = ">=3.8" +python-versions = ">=3.9" groups = ["main", "dev", "proxy-dev"] files = [ - {file = "opentelemetry_sdk-1.25.0-py3-none-any.whl", hash = "sha256:d97ff7ec4b351692e9d5a15af570c693b8715ad78b8aafbec5c7100fe966b4c9"}, - {file = "opentelemetry_sdk-1.25.0.tar.gz", hash = "sha256:ce7fc319c57707ef5bf8b74fb9f8ebdb8bfafbe11898410e0d2a761d08a98ec7"}, + {file = "opentelemetry_sdk-1.39.1-py3-none-any.whl", hash = "sha256:4d5482c478513ecb0a5d938dcc61394e647066e0cc2676bee9f3af3f3f45f01c"}, + {file = "opentelemetry_sdk-1.39.1.tar.gz", hash = "sha256:cf4d4563caf7bff906c9f7967e2be22d0d6b349b908be0d90fb21c8e9c995cc6"}, ] markers = {main = "python_version >= \"3.10\""} [package.dependencies] -opentelemetry-api = "1.25.0" -opentelemetry-semantic-conventions = "0.46b0" -typing-extensions = ">=3.7.4" +opentelemetry-api = "1.39.1" +opentelemetry-semantic-conventions = "0.60b1" +typing-extensions = ">=4.5.0" [[package]] name = "opentelemetry-semantic-conventions" -version = "0.46b0" +version = "0.60b1" description = "OpenTelemetry Semantic Conventions" optional = false -python-versions = ">=3.8" +python-versions = ">=3.9" groups = ["main", "dev", "proxy-dev"] files = [ - {file = "opentelemetry_semantic_conventions-0.46b0-py3-none-any.whl", hash = "sha256:6daef4ef9fa51d51855d9f8e0ccd3a1bd59e0e545abe99ac6203804e36ab3e07"}, - {file = "opentelemetry_semantic_conventions-0.46b0.tar.gz", hash = "sha256:fbc982ecbb6a6e90869b15c1673be90bd18c8a56ff1cffc0864e38e2edffaefa"}, + {file = "opentelemetry_semantic_conventions-0.60b1-py3-none-any.whl", hash = "sha256:9fa8c8b0c110da289809292b0591220d3a7b53c1526a23021e977d68597893fb"}, + {file = "opentelemetry_semantic_conventions-0.60b1.tar.gz", hash = "sha256:87c228b5a0669b748c76d76df6c364c369c28f1c465e50f661e39737e84bc953"}, ] markers = {main = "python_version >= \"3.10\""} [package.dependencies] -opentelemetry-api = "1.25.0" +opentelemetry-api = "1.39.1" +typing-extensions = ">=4.5.0" [[package]] name = "orjson" @@ -4828,23 +4851,23 @@ testing = ["google-api-core (>=1.31.5)"] [[package]] name = "protobuf" -version = "4.25.8" +version = "5.29.5" description = "" optional = false python-versions = ">=3.8" groups = ["main", "dev", "proxy-dev"] files = [ - {file = "protobuf-4.25.8-cp310-abi3-win32.whl", hash = "sha256:504435d831565f7cfac9f0714440028907f1975e4bed228e58e72ecfff58a1e0"}, - {file = "protobuf-4.25.8-cp310-abi3-win_amd64.whl", hash = "sha256:bd551eb1fe1d7e92c1af1d75bdfa572eff1ab0e5bf1736716814cdccdb2360f9"}, - {file = "protobuf-4.25.8-cp37-abi3-macosx_10_9_universal2.whl", hash = "sha256:ca809b42f4444f144f2115c4c1a747b9a404d590f18f37e9402422033e464e0f"}, - {file = "protobuf-4.25.8-cp37-abi3-manylinux2014_aarch64.whl", hash = "sha256:9ad7ef62d92baf5a8654fbb88dac7fa5594cfa70fd3440488a5ca3bfc6d795a7"}, - {file = "protobuf-4.25.8-cp37-abi3-manylinux2014_x86_64.whl", hash = "sha256:83e6e54e93d2b696a92cad6e6efc924f3850f82b52e1563778dfab8b355101b0"}, - {file = "protobuf-4.25.8-cp38-cp38-win32.whl", hash = "sha256:27d498ffd1f21fb81d987a041c32d07857d1d107909f5134ba3350e1ce80a4af"}, - {file = "protobuf-4.25.8-cp38-cp38-win_amd64.whl", hash = "sha256:d552c53d0415449c8d17ced5c341caba0d89dbf433698e1436c8fa0aae7808a3"}, - {file = "protobuf-4.25.8-cp39-cp39-win32.whl", hash = "sha256:077ff8badf2acf8bc474406706ad890466274191a48d0abd3bd6987107c9cde5"}, - {file = "protobuf-4.25.8-cp39-cp39-win_amd64.whl", hash = "sha256:f4510b93a3bec6eba8fd8f1093e9d7fb0d4a24d1a81377c10c0e5bbfe9e4ed24"}, - {file = "protobuf-4.25.8-py3-none-any.whl", hash = "sha256:15a0af558aa3b13efef102ae6e4f3efac06f1eea11afb3a57db2901447d9fb59"}, - {file = "protobuf-4.25.8.tar.gz", hash = "sha256:6135cf8affe1fc6f76cced2641e4ea8d3e59518d1f24ae41ba97bcad82d397cd"}, + {file = "protobuf-5.29.5-cp310-abi3-win32.whl", hash = "sha256:3f1c6468a2cfd102ff4703976138844f78ebd1fb45f49011afc5139e9e283079"}, + {file = "protobuf-5.29.5-cp310-abi3-win_amd64.whl", hash = "sha256:3f76e3a3675b4a4d867b52e4a5f5b78a2ef9565549d4037e06cf7b0942b1d3fc"}, + {file = "protobuf-5.29.5-cp38-abi3-macosx_10_9_universal2.whl", hash = "sha256:e38c5add5a311f2a6eb0340716ef9b039c1dfa428b28f25a7838ac329204a671"}, + {file = "protobuf-5.29.5-cp38-abi3-manylinux2014_aarch64.whl", hash = "sha256:fa18533a299d7ab6c55a238bf8629311439995f2e7eca5caaff08663606e9015"}, + {file = "protobuf-5.29.5-cp38-abi3-manylinux2014_x86_64.whl", hash = "sha256:63848923da3325e1bf7e9003d680ce6e14b07e55d0473253a690c3a8b8fd6e61"}, + {file = "protobuf-5.29.5-cp38-cp38-win32.whl", hash = "sha256:ef91363ad4faba7b25d844ef1ada59ff1604184c0bcd8b39b8a6bef15e1af238"}, + {file = "protobuf-5.29.5-cp38-cp38-win_amd64.whl", hash = "sha256:7318608d56b6402d2ea7704ff1e1e4597bee46d760e7e4dd42a3d45e24b87f2e"}, + {file = "protobuf-5.29.5-cp39-cp39-win32.whl", hash = "sha256:6f642dc9a61782fa72b90878af134c5afe1917c89a568cd3476d758d3c3a0736"}, + {file = "protobuf-5.29.5-cp39-cp39-win_amd64.whl", hash = "sha256:470f3af547ef17847a28e1f47200a1cbf0ba3ff57b7de50d22776607cd2ea353"}, + {file = "protobuf-5.29.5-py3-none-any.whl", hash = "sha256:6cf42630262c59b2d8de33954443d94b746c952b01434fc58a417fdbd2e84bd5"}, + {file = "protobuf-5.29.5.tar.gz", hash = "sha256:bc1463bafd4b0929216c35f437a8e28731a2b7fe3d98bb77a600efced5a15c84"}, ] markers = {main = "python_version >= \"3.10\" and (extra == \"mlflow\" or extra == \"extra-proxy\") or extra == \"extra-proxy\""} @@ -7687,7 +7710,7 @@ version = "1.17.3" description = "Module for decorators, wrappers and monkey patching." optional = false python-versions = ">=3.8" -groups = ["main", "dev", "proxy-dev"] +groups = ["dev"] files = [ {file = "wrapt-1.17.3-cp310-cp310-macosx_10_9_universal2.whl", hash = "sha256:88bbae4d40d5a46142e70d58bf664a89b6b4befaea7b2ecc14e03cedb8e06c04"}, {file = "wrapt-1.17.3-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:e6b13af258d6a9ad602d57d889f83b9d5543acd471eee12eb51f5b01f8eb1bc2"}, @@ -7771,7 +7794,6 @@ files = [ {file = "wrapt-1.17.3-py3-none-any.whl", hash = "sha256:7171ae35d2c33d326ac19dd8facb1e82e5fd04ef8c6c0e394d7af55a55051c22"}, {file = "wrapt-1.17.3.tar.gz", hash = "sha256:f66eb08feaa410fe4eebd17f2a2c8e2e46d3476e9f8c783daa8e09e0faa666d0"}, ] -markers = {main = "python_version >= \"3.10\""} [[package]] name = "wsproto" @@ -7972,7 +7994,7 @@ type = ["pytest-mypy"] [extras] caching = ["diskcache"] -extra-proxy = ["azure-identity", "azure-keyvault-secrets", "google-cloud-iam", "google-cloud-kms", "prisma", "redisvl", "resend"] +extra-proxy = ["a2a-sdk", "azure-identity", "azure-keyvault-secrets", "google-cloud-iam", "google-cloud-kms", "prisma", "redisvl", "resend"] mlflow = ["mlflow"] proxy = ["PyJWT", "apscheduler", "azure-identity", "azure-storage-blob", "backoff", "boto3", "cryptography", "fastapi", "fastapi-sso", "gunicorn", "litellm-enterprise", "litellm-proxy-extras", "mcp", "orjson", "polars", "pynacl", "python-multipart", "pyyaml", "rich", "rq", "soundfile", "uvicorn", "uvloop", "websockets"] semantic-router = ["semantic-router"] @@ -7981,4 +8003,4 @@ utils = ["numpydoc"] [metadata] lock-version = "2.1" python-versions = ">=3.9,<4.0" -content-hash = "ea62b77c662ab9fc486e421c576f0868bcde16d62a24703ee1f4916a0465ffb2" +content-hash = "82a9400ee6e7a550758024ca3c076c4063c6a22ca9535b125e585b03aca7b2b2" diff --git a/pyproject.toml b/pyproject.toml index aa8e6fd97be..c33c8c995f0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -59,6 +59,7 @@ websockets = {version = "^15.0.1", optional = true} boto3 = {version = "1.36.0", optional = true} redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.9' and python_version < '3.14'"} mcp = {version = "^1.21.2", optional = true, python = ">=3.10"} +a2a-sdk = {version = "^0.3.22", optional = true, python = ">=3.10"} litellm-proxy-extras = {version = "0.4.21", optional = true} rich = {version = "13.7.1", optional = true} litellm-enterprise = {version = "0.1.27", optional = true} @@ -111,7 +112,8 @@ extra_proxy = [ "google-cloud-kms", "google-cloud-iam", "resend", - "redisvl" + "redisvl", + "a2a-sdk" ] utils = [ @@ -147,9 +149,9 @@ types-requests = "*" types-setuptools = "*" types-redis = "*" types-PyYAML = "*" -opentelemetry-api = "1.25.0" -opentelemetry-sdk = "1.25.0" -opentelemetry-exporter-otlp = "1.25.0" +opentelemetry-api = "^1.28.0" +opentelemetry-sdk = "^1.28.0" +opentelemetry-exporter-otlp = "^1.28.0" langfuse = "^2.45.0" fastapi-offline = "^1.7.3" @@ -157,9 +159,9 @@ fastapi-offline = "^1.7.3" prisma = "0.11.0" hypercorn = "^0.15.0" prometheus-client = "0.20.0" -opentelemetry-api = "1.25.0" -opentelemetry-sdk = "1.25.0" -opentelemetry-exporter-otlp = "1.25.0" +opentelemetry-api = "^1.28.0" +opentelemetry-sdk = "^1.28.0" +opentelemetry-exporter-otlp = "^1.28.0" azure-identity = {version = "^1.15.0", python = ">=3.9"} [build-system] diff --git a/requirements.txt b/requirements.txt index 7ec71fb7dc9..4fae2591dd1 100644 --- a/requirements.txt +++ b/requirements.txt @@ -16,12 +16,12 @@ prisma==0.11.0 # for db nodejs-wheel-binaries==24.12.0 ## required by prisma for migrations, prevents runtime download (updated from nodejs-bin for security fixes) mangum==0.17.0 # for aws lambda functions pynacl==1.6.2 # for encrypting keys -google-cloud-aiplatform==1.47.0 # for vertex ai calls +google-cloud-aiplatform==1.133.0 # for vertex ai calls google-cloud-iam==2.19.1 # for GCP IAM Redis authentication -google-genai==1.22.0 +google-genai==1.37.0 anthropic[vertex]==0.54.0 mcp==1.25.0 ; python_version >= "3.10" # for MCP server -google-generativeai==0.5.0 # for vertex ai calls +# google-generativeai removed - deprecated, replaced by google-genai (line 21) async_generator==1.10.0 # for async ollama calls langfuse==2.59.7 # for langfuse self-hosted logging prometheus_client==0.20.0 # for /metrics endpoint on proxy @@ -37,9 +37,10 @@ azure-ai-contentsafety==1.0.0 # for azure content safety azure-identity==1.16.1 ; python_version >= "3.9" # for azure content safety azure-keyvault==4.2.0 # for azure KMS integration azure-storage-file-datalake==12.20.0 # for azure buck storage logging -opentelemetry-api==1.25.0 -opentelemetry-sdk==1.25.0 -opentelemetry-exporter-otlp==1.25.0 +opentelemetry-api==1.28.0 +opentelemetry-sdk==1.28.0 +opentelemetry-exporter-otlp==1.28.0 +a2a-sdk>=0.3.22 ; python_version >= "3.10" # grpcio: 1.68.0-1.68.1 has reconnect bug (#38290), 1.75+ has Python 3.14 wheels + fix grpcio>=1.62.3,!=1.68.*,!=1.69.*,!=1.70.*,!=1.71.0,!=1.71.1,!=1.72.0,!=1.72.1,!=1.73.0; python_version < "3.14" grpcio>=1.75.0; python_version >= "3.14" diff --git a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py index 061e27da919..9588c3b55c3 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py @@ -49,18 +49,20 @@ async def test_invoke_agent_a2a_adds_litellm_data(): # Mock request mock_request = MagicMock() - mock_request.json = AsyncMock(return_value={ - "jsonrpc": "2.0", - "id": "test-id", - "method": "message/send", - "params": { - "message": { - "role": "user", - "parts": [{"kind": "text", "text": "Hello"}], - "messageId": "msg-123", - } - }, - }) + mock_request.json = AsyncMock( + return_value={ + "jsonrpc": "2.0", + "id": "test-id", + "method": "message/send", + "params": { + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "Hello"}], + "messageId": "msg-123", + } + }, + } + ) mock_user_api_key_dict = UserAPIKeyAuth( api_key="sk-test-key", @@ -77,40 +79,44 @@ async def test_invoke_agent_a2a_adds_litellm_data(): SendMessageRequest, SendStreamingMessageRequest, ) + # Real types available - use them - use_real_types = True + pass except ImportError: # Real types not available - create realistic mocks - use_real_types = False - + pass + def make_mock_pydantic_class(name): """Create a mock class that behaves like a Pydantic model.""" + class MockPydanticClass: def __init__(self, **kwargs): self.__dict__.update(kwargs) # Store kwargs for model_dump() if needed self._kwargs = kwargs - + def model_dump(self, mode="json", exclude_none=False): """Mock model_dump method.""" result = dict(self._kwargs) if exclude_none: result = {k: v for k, v in result.items() if v is not None} return result - + MockPydanticClass.__name__ = name return MockPydanticClass - + MessageSendParams = make_mock_pydantic_class("MessageSendParams") SendMessageRequest = make_mock_pydantic_class("SendMessageRequest") - SendStreamingMessageRequest = make_mock_pydantic_class("SendStreamingMessageRequest") - + SendStreamingMessageRequest = make_mock_pydantic_class( + "SendStreamingMessageRequest" + ) + # Create a mock module for a2a.types mock_a2a_types = MagicMock() mock_a2a_types.MessageSendParams = MessageSendParams mock_a2a_types.SendMessageRequest = SendMessageRequest mock_a2a_types.SendStreamingMessageRequest = SendStreamingMessageRequest - + # Patch at the source modules with patch( "litellm.proxy.agent_endpoints.a2a_endpoints._get_agent", @@ -137,12 +143,15 @@ async def test_invoke_agent_a2a_adds_litellm_data(): ), patch.dict( sys.modules, {"a2a": MagicMock(), "a2a.types": mock_a2a_types}, + ), patch( + "litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", + True, ): from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a mock_fastapi_response = MagicMock() - result = await invoke_agent_a2a( + await invoke_agent_a2a( agent_id="test-agent", request=mock_request, fastapi_response=mock_fastapi_response, From 0683f29671e7989ea390444080daf1658a2c8e3a Mon Sep 17 00:00:00 2001 From: Harshit Jain Date: Sat, 17 Jan 2026 17:48:54 +0530 Subject: [PATCH 06/77] feat(panw_prisma_airs): add custom violation message support --- .../docs/proxy/guardrails/panw_prisma_airs.md | 28 +++++++++++++++++++ .../panw_prisma_airs/panw_prisma_airs.py | 15 +++++++++- .../guardrails/guardrail_initializers.py | 1 + 3 files changed, 43 insertions(+), 1 deletion(-) diff --git a/docs/my-website/docs/proxy/guardrails/panw_prisma_airs.md b/docs/my-website/docs/proxy/guardrails/panw_prisma_airs.md index 53f8a03f5bb..e3273a01c17 100644 --- a/docs/my-website/docs/proxy/guardrails/panw_prisma_airs.md +++ b/docs/my-website/docs/proxy/guardrails/panw_prisma_airs.md @@ -206,6 +206,7 @@ Expected successful response: | `mode` | No | When to run the guardrail | `pre_call` | | `fallback_on_error` | No | Action when PANW API is unavailable: `"block"` (fail-closed, default) or `"allow"` (fail-open). Config errors always block. | `block` | | `timeout` | No | PANW API call timeout in seconds (1-60) | `10.0` | +| `violation_message_template` | No | Custom template for error message when request is blocked. Supports `{guardrail_name}`, `{category}`, `{action_type}`, `{default_message}` placeholders. | - | ### Regional Endpoints @@ -449,6 +450,33 @@ LiteLLM does not alter or configure your PANW security profile. To change what c The guardrail is **fail-closed** by default - if the PANW API is unavailable, requests are blocked to ensure no unscanned content reaches your LLM. This provides maximum security. ::: +### Custom Violation Messages + +You can customize the error message returned to the user when a request is blocked by configuring the `violation_message_template` parameter. This is useful for providing user-friendly feedback instead of technical details. + +```yaml +guardrails: + - guardrail_name: "panw-custom-message" + litellm_params: + guardrail: panw_prisma_airs + api_key: os.environ/PANW_PRISMA_AIRS_API_KEY + # Simple message + violation_message_template: "Your request was blocked by our AI Security Policy." + + - guardrail_name: "panw-detailed-message" + litellm_params: + guardrail: panw_prisma_airs + api_key: os.environ/PANW_PRISMA_AIRS_API_KEY + # Message with placeholders + violation_message_template: "{action_type} blocked due to {category} violation. Please contact support." +``` + +**Supported Placeholders:** +- `{guardrail_name}`: Name of the guardrail (e.g. "panw-custom-message") +- `{category}`: Violation category (e.g. "malicious", "injection", "dlp") +- `{action_type}`: "Prompt" or "Response" +- `{default_message}`: The original technical error message + ### Fail-Open Configuration By default, the PANW guardrail operates in **fail-closed** mode for maximum security. If the PANW API is unavailable (timeout, rate limit, network error), requests are blocked. You can configure **fail-open** mode for high-availability scenarios where service continuity is critical. diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py index 02e481acddd..b98eeff99d6 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py @@ -62,6 +62,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): app_name: Optional[str] = None, fallback_on_error: Literal["block", "allow"] = "block", timeout: float = 10.0, + violation_message_template: Optional[str] = None, **kwargs, ): """Initialize PANW Prisma AIRS guardrail handler.""" @@ -77,6 +78,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): default_on=default_on, mask_request_content=_mask_request_content, mask_response_content=_mask_response_content, + violation_message_template=violation_message_template, **kwargs, ) @@ -489,7 +491,18 @@ class PanwPrismaAirsHandler(CustomGuardrail): detection_key = "response_detected" if is_response else "prompt_detected" category = scan_result.get("category", "unknown") - error_msg = f"{action_type} blocked by PANW Prisma AI Security policy (Category: {category})" + default_msg = f"{action_type} blocked by PANW Prisma AI Security policy (Category: {category})" + + # Use custom violation message template if configured + error_msg = self.render_violation_message( + default=default_msg, + context={ + "guardrail_name": self.guardrail_name, + "category": category, + "action_type": action_type, + "default_message": default_msg, + }, + ) error_detail = { "error": { diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index 66b41005c4e..639aebf45c9 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -217,6 +217,7 @@ def initialize_panw_prisma_airs(litellm_params, guardrail): app_name=getattr(litellm_params, "app_name", None), fallback_on_error=getattr(litellm_params, "fallback_on_error", "block"), timeout=float(getattr(litellm_params, "timeout", 10.0)), + violation_message_template=litellm_params.violation_message_template, ) litellm.logging_callback_manager.add_litellm_callback(_panw_callback) From 1301896e03f3267bda8bf56cf40fae5c44342b23 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Sat, 17 Jan 2026 09:08:41 -0800 Subject: [PATCH 07/77] Adjusting new badges --- .../components/EntityUsage/EntityUsage.tsx | 15 ++++---- .../UsagePage/components/UsagePageView.tsx | 35 ++++++++----------- .../src/components/leftnav.tsx | 12 +++---- .../src/components/view_logs/index.tsx | 5 +-- 4 files changed, 30 insertions(+), 37 deletions(-) diff --git a/ui/litellm-dashboard/src/components/UsagePage/components/EntityUsage/EntityUsage.tsx b/ui/litellm-dashboard/src/components/UsagePage/components/EntityUsage/EntityUsage.tsx index 7b8fcc4896f..32882341921 100644 --- a/ui/litellm-dashboard/src/components/UsagePage/components/EntityUsage/EntityUsage.tsx +++ b/ui/litellm-dashboard/src/components/UsagePage/components/EntityUsage/EntityUsage.tsx @@ -24,7 +24,6 @@ import { } from "@tremor/react"; import React, { useEffect, useState } from "react"; import { ActivityMetrics, processActivityData } from "../../../activity_metrics"; -import NewBadge from "../../../common_components/NewBadge"; import { UsageExportHeader } from "../../../EntityUsageExport"; import type { EntityType } from "../../../EntityUsageExport/types"; import { @@ -395,14 +394,12 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti teams={teams || []} /> - - - Cost - {entityType === "agent" ? "Request / Token Consumption" : "Model Activity"} - Key Activity - Endpoint Activity - - + + Cost + {entityType === "agent" ? "Request / Token Consumption" : "Model Activity"} + Key Activity + Endpoint Activity + diff --git a/ui/litellm-dashboard/src/components/UsagePage/components/UsagePageView.tsx b/ui/litellm-dashboard/src/components/UsagePage/components/UsagePageView.tsx index aaecbd063b5..88385248b6e 100644 --- a/ui/litellm-dashboard/src/components/UsagePage/components/UsagePageView.tsx +++ b/ui/litellm-dashboard/src/components/UsagePage/components/UsagePageView.tsx @@ -39,7 +39,6 @@ import { Button } from "@tremor/react"; import { all_admin_roles } from "../../../utils/roles"; import { ActivityMetrics, processActivityData } from "../../activity_metrics"; import CloudZeroExportModal from "../../cloudzero_export_modal"; -import NewBadge from "../../common_components/NewBadge"; import EntityUsageExportModal from "../../EntityUsageExport"; import { Team } from "../../key_team_helpers/key_list"; import { Organization, tagListCall, userDailyActivityAggregatedCall, userDailyActivityCall } from "../../networking"; @@ -438,15 +437,13 @@ const UsagePage: React.FC = ({ teams, organizations }) => { {usageView === "global" && (
- - - Cost - Model Activity - Key Activity - MCP Server Activity - Endpoint Activity - - + + Cost + Model Activity + Key Activity + MCP Server Activity + Endpoint Activity +
)} diff --git a/ui/litellm-dashboard/src/components/leftnav.tsx b/ui/litellm-dashboard/src/components/leftnav.tsx index 9079fd6a7ed..5bd756f1d1d 100644 --- a/ui/litellm-dashboard/src/components/leftnav.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.tsx @@ -7,6 +7,7 @@ import { BarChartOutlined, BgColorsOutlined, BlockOutlined, + BookOutlined, CreditCardOutlined, DatabaseOutlined, ExperimentOutlined, @@ -47,6 +48,7 @@ interface MenuItem { roles?: string[]; children?: MenuItem[]; icon?: React.ReactNode; + external_url?: string; } // Group configuration @@ -213,6 +215,13 @@ const Sidebar: React.FC = ({ setPage, defaultSelectedKey, collapse label: "AI Hub", icon: , }, + { + key: "learning-resources", + page: "learning-resources", + label: "Learning Resources", + icon: , + external_url: "https://models.litellm.ai/cookbook", + }, { key: "experimental", page: "experimental", @@ -252,7 +261,7 @@ const Sidebar: React.FC = ({ setPage, defaultSelectedKey, collapse page: "usage", label: "Old Usage", icon: , - }, + } ], }, ], @@ -364,9 +373,23 @@ const Sidebar: React.FC = ({ setPage, defaultSelectedKey, collapse key: child.key, icon: child.icon, label: child.label, - onClick: () => navigateToPage(child.page), + onClick: () => { + if (child.external_url) { + window.open(child.external_url, "_blank"); + } else { + navigateToPage(child.page); + } + }, })), - onClick: !item.children ? () => navigateToPage(item.page) : undefined, + onClick: !item.children + ? () => { + if (item.external_url) { + window.open(item.external_url, "_blank"); + } else { + navigateToPage(item.page); + } + } + : undefined, })), }); }); diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 82894dfb0e2..c975877d06c 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -35,6 +35,36 @@ export const getCallbackConfigsCall = async (accessToken: string) => { throw error; } }; + +export const getInProductNudgesCall = async (accessToken: string) => { + /** + * Get in-product nudges configuration. + */ + try { + let url = proxyBaseUrl ? `${proxyBaseUrl}/in_product_nudges` : `/in_product_nudges`; + + const response = await fetch(url, { + method: "GET", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + + const data = await response.json(); + return data; + } catch (error) { + console.error("Failed to get in-product nudges:", error); + throw error; + } +}; /** * Helper file for calls being made to proxy */ diff --git a/ui/litellm-dashboard/src/components/survey/ClaudeCodeModal.tsx b/ui/litellm-dashboard/src/components/survey/ClaudeCodeModal.tsx new file mode 100644 index 00000000000..94527c160c1 --- /dev/null +++ b/ui/litellm-dashboard/src/components/survey/ClaudeCodeModal.tsx @@ -0,0 +1,69 @@ +import React from "react"; +import { X, Code, ExternalLink } from "lucide-react"; +import { Button } from "antd"; + +interface ClaudeCodeModalProps { + isOpen: boolean; + onClose: () => void; + onComplete: () => void; +} + +const GOOGLE_FORM_URL = "https://forms.gle/LZeJQ3XytBakckYa9"; + +export function ClaudeCodeModal({ isOpen, onClose, onComplete }: ClaudeCodeModalProps) { + if (!isOpen) return null; + + const handleOpenForm = () => { + window.open(GOOGLE_FORM_URL, "_blank", "noopener,noreferrer"); + onComplete(); + }; + + return ( +
+ {/* Backdrop */} +
+ + {/* Modal */} +
+ {/* Header */} +
+
+ + Claude Code Feedback +
+ +
+ + {/* Content */} +
+

+ Help us improve your experience +

+

+ We'd love to hear about your experience using LiteLLM with Claude Code. Your feedback helps us improve the product for everyone. +

+

+ This brief survey takes about 2-3 minutes to complete. +

+ + +
+
+
+ ); +} + diff --git a/ui/litellm-dashboard/src/components/survey/ClaudeCodePrompt.tsx b/ui/litellm-dashboard/src/components/survey/ClaudeCodePrompt.tsx new file mode 100644 index 00000000000..9006575ce0a --- /dev/null +++ b/ui/litellm-dashboard/src/components/survey/ClaudeCodePrompt.tsx @@ -0,0 +1,26 @@ +import React from "react"; +import { Code } from "lucide-react"; +import { NudgePrompt } from "./NudgePrompt"; + +interface ClaudeCodePromptProps { + onOpen: () => void; + onDismiss: () => void; + isVisible: boolean; +} + +export function ClaudeCodePrompt({ onOpen, onDismiss, isVisible }: ClaudeCodePromptProps) { + return ( + + ); +} + diff --git a/ui/litellm-dashboard/src/components/survey/NudgePrompt.tsx b/ui/litellm-dashboard/src/components/survey/NudgePrompt.tsx new file mode 100644 index 00000000000..9095c6c21c5 --- /dev/null +++ b/ui/litellm-dashboard/src/components/survey/NudgePrompt.tsx @@ -0,0 +1,91 @@ +import React, { useEffect, useState } from "react"; +import { X, LucideIcon } from "lucide-react"; +import { Button } from "antd"; + +interface NudgePromptProps { + onOpen: () => void; + onDismiss: () => void; + isVisible: boolean; + title: string; + description: string; + buttonText: string; + icon: LucideIcon; + accentColor: string; + buttonStyle?: React.CSSProperties; +} + +const DISMISS_DURATION = 15000; // 15 seconds + +export function NudgePrompt({ + onOpen, + onDismiss, + isVisible, + title, + description, + buttonText, + icon: Icon, + accentColor, + buttonStyle, +}: NudgePromptProps) { + const [progress, setProgress] = useState(100); + + useEffect(() => { + if (!isVisible) { + setProgress(100); + return; + } + + const startTime = Date.now(); + const interval = setInterval(() => { + const elapsed = Date.now() - startTime; + const remaining = Math.max(0, 100 - (elapsed / DISMISS_DURATION) * 100); + setProgress(remaining); + + if (remaining <= 0) { + clearInterval(interval); + } + }, 50); + + return () => clearInterval(interval); + }, [isVisible]); + + if (!isVisible) return null; + + return ( +
+ {/* Progress bar at top showing time remaining */} +
+
+
+ +
+
+
+ + {title} +
+ +
+ +

{description}

+ + +
+
+ ); +} + diff --git a/ui/litellm-dashboard/src/components/survey/SurveyPrompt.tsx b/ui/litellm-dashboard/src/components/survey/SurveyPrompt.tsx index 69a886bddeb..a41acb265a7 100644 --- a/ui/litellm-dashboard/src/components/survey/SurveyPrompt.tsx +++ b/ui/litellm-dashboard/src/components/survey/SurveyPrompt.tsx @@ -1,6 +1,6 @@ -import React, { useEffect, useState } from "react"; -import { MessageSquare, X } from "lucide-react"; -import { Button } from "antd"; +import React from "react"; +import { MessageSquare } from "lucide-react"; +import { NudgePrompt } from "./NudgePrompt"; interface SurveyPromptProps { onOpen: () => void; @@ -8,70 +8,18 @@ interface SurveyPromptProps { isVisible: boolean; } -const DISMISS_DURATION = 15000; // 15 seconds - export function SurveyPrompt({ onOpen, onDismiss, isVisible }: SurveyPromptProps) { - const [progress, setProgress] = useState(100); - - useEffect(() => { - if (!isVisible) { - setProgress(100); - return; - } - - const startTime = Date.now(); - const interval = setInterval(() => { - const elapsed = Date.now() - startTime; - const remaining = Math.max(0, 100 - (elapsed / DISMISS_DURATION) * 100); - setProgress(remaining); - - if (remaining <= 0) { - clearInterval(interval); - } - }, 50); - - return () => clearInterval(interval); - }, [isVisible]); - - if (!isVisible) return null; - return ( -
- {/* Progress bar at top showing time remaining */} -
-
-
- -
-
-
- - Quick feedback -
- -
- -

- Help us improve LiteLLM! Share your experience in 5 quick questions. -

- - -
-
+ ); } diff --git a/ui/litellm-dashboard/src/components/survey/index.tsx b/ui/litellm-dashboard/src/components/survey/index.tsx index 7c36419980c..fde05084b7e 100644 --- a/ui/litellm-dashboard/src/components/survey/index.tsx +++ b/ui/litellm-dashboard/src/components/survey/index.tsx @@ -1,3 +1,6 @@ export { SurveyPrompt } from "./SurveyPrompt"; export { SurveyModal } from "./SurveyModal"; +export { ClaudeCodePrompt } from "./ClaudeCodePrompt"; +export { ClaudeCodeModal } from "./ClaudeCodeModal"; +export { NudgePrompt } from "./NudgePrompt"; From bc4fbf282a67e77f26fe9e0ccbd53c071c093d08 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Sun, 18 Jan 2026 14:45:49 +0900 Subject: [PATCH 46/77] fix: Avoid attaching tool calls when a call_id already exists --- .../transformation.py | 91 ++++++++++++++++++- 1 file changed, 87 insertions(+), 4 deletions(-) diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index ad910c9cd97..a69f2c48bcf 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -2,7 +2,7 @@ Handles transforming from Responses API -> LiteLLM completion (Chat Completion API) """ -from typing import Any, Dict, List, Literal, Optional, Tuple, Union, cast +from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast from openai.types.responses import ResponseFunctionToolCall from openai.types.responses.tool_param import FunctionToolParam @@ -378,6 +378,7 @@ class LiteLLMCompletionResponsesConfig: if isinstance(input, str): messages.append(ChatCompletionUserMessage(role="user", content=input)) elif isinstance(input, list): + existing_tool_call_ids: Set[str] = set() for _input in input: chat_completion_messages = LiteLLMCompletionResponsesConfig._transform_responses_api_input_item_to_chat_completion_message( input_item=_input @@ -390,12 +391,94 @@ class LiteLLMCompletionResponsesConfig: input_item=_input ): tool_call_output_messages.extend(chat_completion_messages) - else: - messages.extend(chat_completion_messages) + continue - messages.extend(tool_call_output_messages) + if LiteLLMCompletionResponsesConfig._is_input_item_function_call( + input_item=_input + ): + call_id_raw = _input.get("call_id") or _input.get("id") or "" + if call_id_raw: + existing_tool_call_ids.add(str(call_id_raw)) + + messages.extend(chat_completion_messages) + + deduped_tool_call_messages = ( + LiteLLMCompletionResponsesConfig._deduplicate_tool_call_output_messages( + tool_call_output_messages=tool_call_output_messages, + existing_tool_call_ids=existing_tool_call_ids, + ) + ) + messages.extend(deduped_tool_call_messages) return messages + @staticmethod + def _deduplicate_tool_call_output_messages( + tool_call_output_messages: List[ + Union[ + AllMessageValues, + GenericChatCompletionMessage, + ChatCompletionMessageToolCall, + ChatCompletionResponseMessage, + ] + ], + existing_tool_call_ids: Set[str], + ) -> List[ + Union[ + AllMessageValues, + GenericChatCompletionMessage, + ChatCompletionMessageToolCall, + ChatCompletionResponseMessage, + ] + ]: + """Return tool call outputs after dropping assistant entries with duplicate call_ids.""" + if not tool_call_output_messages: + return [] + + filtered_messages: List[ + Union[ + AllMessageValues, + GenericChatCompletionMessage, + ChatCompletionMessageToolCall, + ChatCompletionResponseMessage, + ] + ] = [] + seen_tool_call_ids: Set[str] = set(existing_tool_call_ids) + + for tool_call_message in tool_call_output_messages: + if isinstance(tool_call_message, dict): + role = tool_call_message.get("role", "") + else: + role = getattr(tool_call_message, "role", "") + call_id = "" + + if role == "assistant": + tool_calls = None + if isinstance(tool_call_message, dict): + tool_calls = tool_call_message.get("tool_calls") + else: + tool_calls = getattr(tool_call_message, "tool_calls", None) + + if tool_calls and len(tool_calls) > 0: + first_call = tool_calls[0] + call_id_raw = None + if isinstance(first_call, dict): + call_id_raw = first_call.get("id") + else: + call_id_raw = getattr(first_call, "id", None) + + if call_id_raw: + call_id = str(call_id_raw) + + if call_id and call_id in seen_tool_call_ids and role == "assistant": + continue + + if call_id and role == "assistant": + seen_tool_call_ids.add(call_id) + + filtered_messages.append(tool_call_message) + + return filtered_messages + @staticmethod def _ensure_tool_call_output_has_corresponding_tool_call( messages: List[Union[AllMessageValues, GenericChatCompletionMessage]], From 38afdb71ff29160543eb8988ca3123df66d07d4a Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Mon, 19 Jan 2026 06:03:10 +0900 Subject: [PATCH 47/77] fix: Prevent MCP responses from reviving past tool calls via previous_response_id --- litellm/responses/mcp/mcp_streaming_iterator.py | 1 - 1 file changed, 1 deletion(-) diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index c00c2a2f3b2..ac040d3d6ec 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -655,7 +655,6 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): follow_up_params.update( { "input": follow_up_input, - "previous_response_id": self.collected_response.id, # type: ignore[attr-defined] "stream": True, } ) From cd19039e39b3dc6ef1e4bfa17f853e2bbbb0d3d3 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Mon, 19 Jan 2026 06:22:05 +0900 Subject: [PATCH 48/77] test: Parametrize MCP streaming test to cover OpenAI and Anthropic models --- tests/mcp_tests/test_aresponses_api_with_mcp.py | 15 +++++++++++---- 1 file changed, 11 insertions(+), 4 deletions(-) diff --git a/tests/mcp_tests/test_aresponses_api_with_mcp.py b/tests/mcp_tests/test_aresponses_api_with_mcp.py index 865a580f0ca..2ec90ab5858 100644 --- a/tests/mcp_tests/test_aresponses_api_with_mcp.py +++ b/tests/mcp_tests/test_aresponses_api_with_mcp.py @@ -660,16 +660,23 @@ async def test_streaming_mcp_events_validation(): @pytest.mark.asyncio -async def test_streaming_responses_api_with_mcp_tools(): +@pytest.mark.parametrize( + "model", + [ + pytest.param("gpt-4o-mini", id="openai"), + pytest.param("anthropic/claude-4-5-haiku", id="anthropic"), + ], +) +async def test_streaming_responses_api_with_mcp_tools(model: str): """ Test the streaming responses API with MCP tools when using server_url="litellm_proxy" Under the hood the follow occurs - MCP: responses called litellm MCP manager.list_tools (MOCKED) - - Request 1: Made to gpt-4o with fetched tools (REAL LLM CALL) + - Request 1: Made to model under test with fetched tools (REAL LLM CALL) - MCP: Execute tool call from request 1 and returns result (MOCKED) - - Request 2: Made to gpt-4o with fetched tools and tool results (REAL LLM CALL) + - Request 2: Made to model under test with fetched tools and tool results (REAL LLM CALL) Return the user the result of request 2 """ @@ -729,7 +736,7 @@ async def test_streaming_responses_api_with_mcp_tools(): "require_approval": "never" }) response = await litellm.aresponses( - model="gpt-4o-mini", + model=model, tools=[mcp_tool_config], tool_choice="required", input=[ From 4ad78236ab5a84e64f6104d3c1d5366db2ca2c6a Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Mon, 19 Jan 2026 06:46:22 +0900 Subject: [PATCH 49/77] test: Fail MCP streaming test when LiteLLM logs errors during follow-up calls --- .../mcp_tests/test_aresponses_api_with_mcp.py | 165 ++++++++++-------- 1 file changed, 96 insertions(+), 69 deletions(-) diff --git a/tests/mcp_tests/test_aresponses_api_with_mcp.py b/tests/mcp_tests/test_aresponses_api_with_mcp.py index 2ec90ab5858..5871cece869 100644 --- a/tests/mcp_tests/test_aresponses_api_with_mcp.py +++ b/tests/mcp_tests/test_aresponses_api_with_mcp.py @@ -1,3 +1,4 @@ +import logging import os import sys import pytest @@ -667,7 +668,9 @@ async def test_streaming_mcp_events_validation(): pytest.param("anthropic/claude-4-5-haiku", id="anthropic"), ], ) -async def test_streaming_responses_api_with_mcp_tools(model: str): +async def test_streaming_responses_api_with_mcp_tools( + model: str, caplog: pytest.LogCaptureFixture +): """ Test the streaming responses API with MCP tools when using server_url="litellm_proxy" @@ -700,75 +703,99 @@ async def test_streaming_responses_api_with_mcp_tools(model: str): ] # Only mock the MCP-specific operations, let LLM responses be real - with patch.object(LiteLLM_Proxy_MCP_Handler, '_get_mcp_tools_from_manager', new_callable=AsyncMock) as mock_get_tools, \ - patch.object(LiteLLM_Proxy_MCP_Handler, '_execute_tool_calls', new_callable=AsyncMock) as mock_execute_tools: - - # Setup MCP mocks only - mock_get_tools.return_value = (mock_mcp_tools, ["litellm_proxy"]) - - # Create a dynamic mock that will match the actual tool call ID from the LLM response - def mock_execute_tool_calls_side_effect(tool_calls, user_api_key_auth): - """Mock function that returns results matching the actual tool call IDs from the LLM""" - results = [] - for tool_call in tool_calls: - # Extract call_id from the tool call - call_id = None - if isinstance(tool_call, dict): - call_id = tool_call.get("call_id") or tool_call.get("id") - elif hasattr(tool_call, 'call_id'): - call_id = tool_call.call_id - elif hasattr(tool_call, 'id'): - call_id = tool_call.id - - if call_id: - results.append({ - "tool_call_id": call_id, - "result": "LiteLLM is a unified interface for 100+ LLMs that translates inputs to provider-specific completion endpoints and provides consistent OpenAI-format output." - }) - return results - - mock_execute_tools.side_effect = mock_execute_tool_calls_side_effect - - # Make the actual call - LLM responses will be real - mcp_tool_config = cast(Any, { - "type": "mcp", - "server_url": "litellm_proxy", - "require_approval": "never" - }) - response = await litellm.aresponses( - model=model, - tools=[mcp_tool_config], - tool_choice="required", - input=[ + with caplog.at_level(logging.ERROR): + with patch.object( + LiteLLM_Proxy_MCP_Handler, + '_get_mcp_tools_from_manager', + new_callable=AsyncMock, + ) as mock_get_tools, patch.object( + LiteLLM_Proxy_MCP_Handler, + '_execute_tool_calls', + new_callable=AsyncMock, + ) as mock_execute_tools: + # Setup MCP mocks only + mock_get_tools.return_value = (mock_mcp_tools, ["litellm_proxy"]) + + # Create a dynamic mock that will match the actual tool call ID from the LLM response + def mock_execute_tool_calls_side_effect(tool_calls, user_api_key_auth): + """Mock function that returns results matching the actual tool call IDs from the LLM""" + results = [] + for tool_call in tool_calls: + # Extract call_id from the tool call + call_id = None + if isinstance(tool_call, dict): + call_id = tool_call.get("call_id") or tool_call.get("id") + elif hasattr(tool_call, 'call_id'): + call_id = tool_call.call_id + elif hasattr(tool_call, 'id'): + call_id = tool_call.id + + if call_id: + results.append( + { + "tool_call_id": call_id, + "result": "LiteLLM is a unified interface for 100+ LLMs that translates inputs to provider-specific completion endpoints and provides consistent OpenAI-format output.", + } + ) + return results + + mock_execute_tools.side_effect = mock_execute_tool_calls_side_effect + + # Make the actual call - LLM responses will be real + mcp_tool_config = cast( + Any, { - "role": "user", - "type": "message", - "content": "give me a TLDR of what BerriAI/litellm is about" - } - ], - stream=True - ) - - print(f"📋 Response type: {type(response)}") - assert hasattr(response, '__aiter__'), "Response should be an async streaming response" - - # Collect streaming chunks - chunks = [] - async for chunk in response: - chunks.append(chunk) - print(f"📦 Chunk type: {getattr(chunk, 'type', 'unknown')}") - - print(f"📊 Total chunks received: {len(chunks)}") - - # Verify MCP mocks were called (may be called multiple times in streaming) - assert mock_get_tools.call_count >= 1, f"Expected MCP tools to be fetched at least once, got {mock_get_tools.call_count}" - print(f"MCP tools fetched: {len(mock_mcp_tools)}") - - # Verify we got a response - assert response is not None - assert len(chunks) > 0, "Should have received streaming chunks" - - print("Basic streaming responses API with MCP tools test passed!") + "type": "mcp", + "server_url": "litellm_proxy", + "require_approval": "never", + }, + ) + response = await litellm.aresponses( + model=model, + tools=[mcp_tool_config], + tool_choice="required", + input=[ + { + "role": "user", + "type": "message", + "content": "give me a TLDR of what BerriAI/litellm is about", + } + ], + stream=True, + ) + + print(f"📋 Response type: {type(response)}") + assert hasattr(response, '__aiter__'), "Response should be an async streaming response" + + # Collect streaming chunks + chunks = [] + async for chunk in response: + chunks.append(chunk) + print(f"📦 Chunk type: {getattr(chunk, 'type', 'unknown')}") + + print(f"📊 Total chunks received: {len(chunks)}") + + # Verify MCP mocks were called (may be called multiple times in streaming) + assert ( + mock_get_tools.call_count >= 1 + ), f"Expected MCP tools to be fetched at least once, got {mock_get_tools.call_count}" + print(f"MCP tools fetched: {len(mock_mcp_tools)}") + + # Verify we got a response + assert response is not None + assert len(chunks) > 0, "Should have received streaming chunks" + + print("Basic streaming responses API with MCP tools test passed!") + + lite_errors = [ + record + for record in caplog.records + if record.levelno >= logging.ERROR + and ("LiteLLM" in record.name or "LiteLLM" in record.getMessage()) + ] + assert not lite_errors, "Unexpected LiteLLM errors: " + ", ".join( + record.getMessage() for record in lite_errors + ) @pytest.mark.asyncio From d31c609600f16695fcd0a15b4c9b52d73e57d533 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Mon, 19 Jan 2026 07:00:14 +0900 Subject: [PATCH 50/77] test: Let MCP tool-execution mock accept new kwargs for streaming tests --- tests/mcp_tests/test_aresponses_api_with_mcp.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/tests/mcp_tests/test_aresponses_api_with_mcp.py b/tests/mcp_tests/test_aresponses_api_with_mcp.py index 5871cece869..57c79039ee0 100644 --- a/tests/mcp_tests/test_aresponses_api_with_mcp.py +++ b/tests/mcp_tests/test_aresponses_api_with_mcp.py @@ -665,7 +665,7 @@ async def test_streaming_mcp_events_validation(): "model", [ pytest.param("gpt-4o-mini", id="openai"), - pytest.param("anthropic/claude-4-5-haiku", id="anthropic"), + pytest.param("claude-haiku-4-5", id="anthropic"), ], ) async def test_streaming_responses_api_with_mcp_tools( @@ -717,7 +717,9 @@ async def test_streaming_responses_api_with_mcp_tools( mock_get_tools.return_value = (mock_mcp_tools, ["litellm_proxy"]) # Create a dynamic mock that will match the actual tool call ID from the LLM response - def mock_execute_tool_calls_side_effect(tool_calls, user_api_key_auth): + def mock_execute_tool_calls_side_effect( + tool_calls, user_api_key_auth, **kwargs + ): """Mock function that returns results matching the actual tool call IDs from the LLM""" results = [] for tool_call in tool_calls: From 60f40856e12e7bf6ea323a3f453be2d2063706fc Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Mon, 19 Jan 2026 07:10:32 +0900 Subject: [PATCH 51/77] chore: fix lint error --- .../litellm_completion_transformation/transformation.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index a69f2c48bcf..eaa80c6cfe4 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -2,6 +2,7 @@ Handles transforming from Responses API -> LiteLLM completion (Chat Completion API) """ +from collections.abc import Sequence from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast from openai.types.responses import ResponseFunctionToolCall @@ -452,13 +453,17 @@ class LiteLLMCompletionResponsesConfig: call_id = "" if role == "assistant": - tool_calls = None + tool_calls: Any = None if isinstance(tool_call_message, dict): tool_calls = tool_call_message.get("tool_calls") else: tool_calls = getattr(tool_call_message, "tool_calls", None) - if tool_calls and len(tool_calls) > 0: + if ( + isinstance(tool_calls, Sequence) + and not isinstance(tool_calls, (str, bytes)) + and len(tool_calls) > 0 + ): first_call = tool_calls[0] call_id_raw = None if isinstance(first_call, dict): From 737fec600f6afff8c976a03fa937285e874559bd Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Mon, 19 Jan 2026 10:49:39 +0900 Subject: [PATCH 52/77] test: add mcp e2e test --- .../test_configs/test_config_mcp_e2e.yaml | 20 ++++ tests/mcp_tests/test_proxy_mcp_e2e.py | 94 +++++++++++++++++++ 2 files changed, 114 insertions(+) create mode 100644 tests/mcp_tests/test_configs/test_config_mcp_e2e.yaml create mode 100644 tests/mcp_tests/test_proxy_mcp_e2e.py diff --git a/tests/mcp_tests/test_configs/test_config_mcp_e2e.yaml b/tests/mcp_tests/test_configs/test_config_mcp_e2e.yaml new file mode 100644 index 00000000000..06d0878f788 --- /dev/null +++ b/tests/mcp_tests/test_configs/test_config_mcp_e2e.yaml @@ -0,0 +1,20 @@ +general_settings: + master_key: sk-1234 + +litellm_settings: + drop_params: true + +model_list: + - model_name: openai-gpt-4o-mini + litellm_params: + model: gpt-4o-mini + - model_name: anthropic-claude-haiku-4-5 + litellm_params: + model: anthropic/claude-haiku-4-5 + +mcp_servers: + math_stdio: + transport: stdio + command: python3 + args: + - tests/mcp_tests/mcp_server.py diff --git a/tests/mcp_tests/test_proxy_mcp_e2e.py b/tests/mcp_tests/test_proxy_mcp_e2e.py new file mode 100644 index 00000000000..44d340dbcb7 --- /dev/null +++ b/tests/mcp_tests/test_proxy_mcp_e2e.py @@ -0,0 +1,94 @@ +import asyncio +import socket +import threading +import time +from pathlib import Path + +import pytest +import uvicorn +from mcp import ClientSession +from mcp.client.streamable_http import streamablehttp_client + +from litellm.proxy.proxy_server import ( + app as proxy_app, + cleanup_router_config_variables, + initialize, +) + + +CONFIG_TEMPLATE_PATH = Path("tests/mcp_tests/test_configs/test_config_mcp_e2e.yaml") +PROXY_START_TIMEOUT = 30 +MCP_HEADERS = { + "Authorization": "Bearer sk-1234", + "x-mcp-servers": "math_stdio", +} + + +def _initialize_proxy(config_path: str) -> None: + cleanup_router_config_variables() + asyncio.run(initialize(config=config_path, debug=True)) + + +def _start_proxy_server(config_path: str) -> tuple[str, uvicorn.Server, threading.Thread, socket.socket]: + _initialize_proxy(config_path) + + sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + sock.bind(("127.0.0.1", 0)) + host, port = sock.getsockname() + + config = uvicorn.Config(proxy_app, host=host, port=port, log_level="warning") + server = uvicorn.Server(config) + + def _run() -> None: + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + loop.run_until_complete(server.serve(sockets=[sock])) + + thread = threading.Thread(target=_run, daemon=True) + thread.start() + + start_time = time.time() + while not server.started: + if not thread.is_alive(): + raise RuntimeError("Proxy server failed to start") + if time.time() - start_time > PROXY_START_TIMEOUT: + raise TimeoutError("Proxy server did not start in time") + time.sleep(0.05) + + return f"http://{host}:{port}", server, thread, sock + + +@pytest.fixture(scope="session") +def proxy_server_url(tmp_path_factory: pytest.TempPathFactory): + config_dir = tmp_path_factory.mktemp("mcp_e2e") + config_path = config_dir / "config.yaml" + config_path.write_text(CONFIG_TEMPLATE_PATH.read_text()) + + server_url, server, thread, sock = _start_proxy_server(str(config_path)) + + yield server_url + + server.should_exit = True + thread.join(timeout=10) + sock.close() + + +@pytest.mark.asyncio +async def test_proxy_mcp_stdio_roundtrip(proxy_server_url: str) -> None: + async with asyncio.timeout(20): + async with streamablehttp_client( + url=f"{proxy_server_url}/mcp", headers=MCP_HEADERS + ) as (read, write, _get_session_id): + async with ClientSession(read, write) as session: + await session.initialize() + tools_result = await session.list_tools() + assert any(tool.name.endswith("add") for tool in tools_result.tools) + + result = await session.call_tool( + "add", arguments={"a": 3, "b": 4} + ) + assert result.content + first_content = result.content[0] + text = getattr(first_content, "text", None) + assert text == "7" From c2b5e9c6697e433e5b0af7363ee0c9d7ecd5031a Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Mon, 19 Jan 2026 11:12:36 +0900 Subject: [PATCH 53/77] test: MCP E2E streamable_http --- tests/mcp_tests/mcp_server.py | 47 ++++- .../test_configs/test_config_mcp_e2e.yaml | 3 + tests/mcp_tests/test_proxy_mcp_e2e.py | 177 ++++++++++++++++-- 3 files changed, 207 insertions(+), 20 deletions(-) diff --git a/tests/mcp_tests/mcp_server.py b/tests/mcp_tests/mcp_server.py index 99a67edd021..c4daff4ab1c 100644 --- a/tests/mcp_tests/mcp_server.py +++ b/tests/mcp_tests/mcp_server.py @@ -1,9 +1,35 @@ # math_server.py +from __future__ import annotations + +import argparse +import os + from mcp.server.fastmcp import FastMCP mcp = FastMCP("Math") +def _parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="MCP math test server") + parser.add_argument( + "--transport", + default=os.getenv("MCP_TRANSPORT", "stdio"), + help="Transport to use (stdio or http)", + ) + parser.add_argument( + "--host", + default=os.getenv("MCP_HOST", "127.0.0.1"), + help="Host to bind when serving over HTTP", + ) + parser.add_argument( + "--port", + type=int, + default=int(os.getenv("MCP_PORT", "0")), + help="Port to bind when serving over HTTP", + ) + return parser.parse_args() + + @mcp.tool() def add(a: int, b: int) -> int: """Add two numbers""" @@ -16,5 +42,24 @@ def multiply(a: int, b: int) -> int: return a * b +def main() -> None: + args = _parse_args() + transport = (args.transport or "stdio").lower() + + if transport == "stdio": + mcp.run(transport="stdio") + return + + if transport in {"http", "streamable_http", "streamable-http"}: + if args.port <= 0: + raise ValueError("HTTP transport requires a valid --port value") + mcp.settings.host = args.host + mcp.settings.port = args.port + mcp.run(transport="streamable-http") + return + + raise ValueError(f"Unsupported transport: {transport}") + + if __name__ == "__main__": - mcp.run(transport="stdio") + main() diff --git a/tests/mcp_tests/test_configs/test_config_mcp_e2e.yaml b/tests/mcp_tests/test_configs/test_config_mcp_e2e.yaml index 06d0878f788..37d5359e3ce 100644 --- a/tests/mcp_tests/test_configs/test_config_mcp_e2e.yaml +++ b/tests/mcp_tests/test_configs/test_config_mcp_e2e.yaml @@ -18,3 +18,6 @@ mcp_servers: command: python3 args: - tests/mcp_tests/mcp_server.py + math_streamable_http: + transport: http + url: http://127.0.0.1:0/mcp diff --git a/tests/mcp_tests/test_proxy_mcp_e2e.py b/tests/mcp_tests/test_proxy_mcp_e2e.py index 44d340dbcb7..8fbd80b8d62 100644 --- a/tests/mcp_tests/test_proxy_mcp_e2e.py +++ b/tests/mcp_tests/test_proxy_mcp_e2e.py @@ -1,11 +1,16 @@ import asyncio +import os import socket +import subprocess +import sys import threading import time +import typing from pathlib import Path import pytest import uvicorn +import yaml from mcp import ClientSession from mcp.client.streamable_http import streamablehttp_client @@ -17,11 +22,28 @@ from litellm.proxy.proxy_server import ( CONFIG_TEMPLATE_PATH = Path("tests/mcp_tests/test_configs/test_config_mcp_e2e.yaml") +MCP_SERVER_SCRIPT = Path("tests/mcp_tests/mcp_server.py") +PROJECT_ROOT = Path(__file__).resolve().parents[2] PROXY_START_TIMEOUT = 30 MCP_HEADERS = { "Authorization": "Bearer sk-1234", "x-mcp-servers": "math_stdio", } +STREAMABLE_HTTP_HEADERS = { + "Authorization": "Bearer sk-1234", + "x-mcp-servers": "math_streamable_http", +} + + +@pytest.fixture(scope="session", autouse=True) +def _clear_proxy_database_env() -> typing.Iterator[None]: + """Ensure local proxy DB settings don't leak into tests.""" + mp = pytest.MonkeyPatch() + mp.delenv("DATABASE_URL", raising=False) + try: + yield + finally: + mp.undo() def _initialize_proxy(config_path: str) -> None: @@ -60,10 +82,67 @@ def _start_proxy_server(config_path: str) -> tuple[str, uvicorn.Server, threadin @pytest.fixture(scope="session") -def proxy_server_url(tmp_path_factory: pytest.TempPathFactory): +def math_streamable_http_server() -> str: + host = "127.0.0.1" + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock: + sock.bind((host, 0)) + _, port = sock.getsockname() + + cmd = [ + sys.executable, + str(MCP_SERVER_SCRIPT), + "--transport", + "http", + "--host", + host, + "--port", + str(port), + ] + + env = os.environ.copy() + server_process = subprocess.Popen( + cmd, + cwd=str(PROJECT_ROOT), + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + ) + + start_time = time.time() + while True: + if server_process.poll() is not None: + stdout, stderr = server_process.communicate() + raise RuntimeError( + f"Streamable HTTP MCP server exited early.\nSTDOUT: {stdout.decode()}\nSTDERR: {stderr.decode()}" + ) + try: + with socket.create_connection((host, port), timeout=0.1): + break + except OSError: + if time.time() - start_time > PROXY_START_TIMEOUT: + server_process.terminate() + raise TimeoutError("Streamable HTTP MCP server did not start in time") + time.sleep(0.05) + + yield f"http://{host}:{port}" + + server_process.terminate() + try: + server_process.wait(timeout=5) + except subprocess.TimeoutExpired: + server_process.kill() + + +@pytest.fixture(scope="session") +def proxy_server_url( + tmp_path_factory: pytest.TempPathFactory, math_streamable_http_server: str +): config_dir = tmp_path_factory.mktemp("mcp_e2e") config_path = config_dir / "config.yaml" - config_path.write_text(CONFIG_TEMPLATE_PATH.read_text()) + config = yaml.safe_load(CONFIG_TEMPLATE_PATH.read_text()) + config["mcp_servers"]["math_streamable_http"][ + "url" + ] = f"{math_streamable_http_server}/mcp" + config_path.write_text(yaml.safe_dump(config)) server_url, server, thread, sock = _start_proxy_server(str(config_path)) @@ -74,21 +153,81 @@ def proxy_server_url(tmp_path_factory: pytest.TempPathFactory): sock.close() -@pytest.mark.asyncio -async def test_proxy_mcp_stdio_roundtrip(proxy_server_url: str) -> None: - async with asyncio.timeout(20): - async with streamablehttp_client( - url=f"{proxy_server_url}/mcp", headers=MCP_HEADERS - ) as (read, write, _get_session_id): - async with ClientSession(read, write) as session: - await session.initialize() - tools_result = await session.list_tools() - assert any(tool.name.endswith("add") for tool in tools_result.tools) +class TestProxyMcpSimpleConnections: + @pytest.mark.asyncio + async def test_proxy_mcp_stdio_roundtrip(self, proxy_server_url: str) -> None: + async with asyncio.timeout(20): + async with streamablehttp_client( + url=f"{proxy_server_url}/mcp", headers=MCP_HEADERS + ) as (read, write, _get_session_id): + async with ClientSession(read, write) as session: + await session.initialize() + tools_result = await session.list_tools() + assert any(tool.name.endswith("add") for tool in tools_result.tools) - result = await session.call_tool( - "add", arguments={"a": 3, "b": 4} - ) - assert result.content - first_content = result.content[0] - text = getattr(first_content, "text", None) - assert text == "7" + result = await session.call_tool( + "add", arguments={"a": 3, "b": 4} + ) + assert result.content + first_content = result.content[0] + text = getattr(first_content, "text", None) + assert text == "7" + + @pytest.mark.asyncio + async def test_proxy_mcp_streamable_http_roundtrip( + self, proxy_server_url: str + ) -> None: + async with asyncio.timeout(20): + async with streamablehttp_client( + url=f"{proxy_server_url}/mcp", headers=STREAMABLE_HTTP_HEADERS + ) as (read, write, _get_session_id): + async with ClientSession(read, write) as session: + await session.initialize() + tools_result = await session.list_tools() + assert any(tool.name.endswith("add") for tool in tools_result.tools) + + result = await session.call_tool( + "add", arguments={"a": 5, "b": 6} + ) + assert result.content + first_content = result.content[0] + text = getattr(first_content, "text", None) + assert text == "11" + + @pytest.mark.asyncio + async def test_proxy_mcp_lists_all_servers_without_header( + self, proxy_server_url: str + ) -> None: + async with asyncio.timeout(20): + async with streamablehttp_client( + url=f"{proxy_server_url}/mcp", + headers={"Authorization": "Bearer sk-1234"}, + ) as (read, write, _get_session_id): + async with ClientSession(read, write) as session: + await session.initialize() + tools_result = await session.list_tools() + tool_names = {tool.name for tool in tools_result.tools} + expected_tool_names = { + "math_stdio-add", + "math_stdio-multiply", + "math_streamable_http-add", + "math_streamable_http-multiply", + } + assert expected_tool_names <= tool_names + + async def _call_and_get_text( + tool_name: str, *, a: int, b: int + ) -> str | None: + result = await session.call_tool(tool_name, arguments={"a": a, "b": b}) + assert result.content + first_content = result.content[0] + return getattr(first_content, "text", None) + + stdio_result = await _call_and_get_text( + "math_stdio-add", a=2, b=3 + ) + streamable_result = await _call_and_get_text( + "math_streamable_http-add", a=4, b=5 + ) + assert stdio_result == "5" + assert streamable_result == "9" From 20b6468222414c6e4641d2d56c5745aed23c0cd6 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Mon, 19 Jan 2026 11:17:44 +0900 Subject: [PATCH 54/77] test: refactor --- tests/mcp_tests/test_proxy_mcp_e2e.py | 36 ++++++++++++++++----------- 1 file changed, 22 insertions(+), 14 deletions(-) diff --git a/tests/mcp_tests/test_proxy_mcp_e2e.py b/tests/mcp_tests/test_proxy_mcp_e2e.py index 8fbd80b8d62..b124fabe8b8 100644 --- a/tests/mcp_tests/test_proxy_mcp_e2e.py +++ b/tests/mcp_tests/test_proxy_mcp_e2e.py @@ -25,14 +25,12 @@ CONFIG_TEMPLATE_PATH = Path("tests/mcp_tests/test_configs/test_config_mcp_e2e.ya MCP_SERVER_SCRIPT = Path("tests/mcp_tests/mcp_server.py") PROJECT_ROOT = Path(__file__).resolve().parents[2] PROXY_START_TIMEOUT = 30 -MCP_HEADERS = { - "Authorization": "Bearer sk-1234", - "x-mcp-servers": "math_stdio", -} -STREAMABLE_HTTP_HEADERS = { - "Authorization": "Bearer sk-1234", - "x-mcp-servers": "math_streamable_http", -} + + +@pytest.fixture(scope="session") +def proxy_authorization_header() -> str: + """Shared Authorization header value for proxy calls.""" + return "Bearer sk-1234" @pytest.fixture(scope="session", autouse=True) @@ -155,10 +153,16 @@ def proxy_server_url( class TestProxyMcpSimpleConnections: @pytest.mark.asyncio - async def test_proxy_mcp_stdio_roundtrip(self, proxy_server_url: str) -> None: + async def test_proxy_mcp_stdio_roundtrip( + self, proxy_server_url: str, proxy_authorization_header: str + ) -> None: async with asyncio.timeout(20): async with streamablehttp_client( - url=f"{proxy_server_url}/mcp", headers=MCP_HEADERS + url=f"{proxy_server_url}/mcp", + headers={ + "Authorization": proxy_authorization_header, + "x-mcp-servers": "math_stdio", + }, ) as (read, write, _get_session_id): async with ClientSession(read, write) as session: await session.initialize() @@ -175,11 +179,15 @@ class TestProxyMcpSimpleConnections: @pytest.mark.asyncio async def test_proxy_mcp_streamable_http_roundtrip( - self, proxy_server_url: str + self, proxy_server_url: str, proxy_authorization_header: str ) -> None: async with asyncio.timeout(20): async with streamablehttp_client( - url=f"{proxy_server_url}/mcp", headers=STREAMABLE_HTTP_HEADERS + url=f"{proxy_server_url}/mcp", + headers={ + "Authorization": proxy_authorization_header, + "x-mcp-servers": "math_streamable_http", + }, ) as (read, write, _get_session_id): async with ClientSession(read, write) as session: await session.initialize() @@ -196,12 +204,12 @@ class TestProxyMcpSimpleConnections: @pytest.mark.asyncio async def test_proxy_mcp_lists_all_servers_without_header( - self, proxy_server_url: str + self, proxy_server_url: str, proxy_authorization_header: str ) -> None: async with asyncio.timeout(20): async with streamablehttp_client( url=f"{proxy_server_url}/mcp", - headers={"Authorization": "Bearer sk-1234"}, + headers={"Authorization": proxy_authorization_header}, ) as (read, write, _get_session_id): async with ClientSession(read, write) as session: await session.initialize() From e326b397c5cf55625e22b9c09692ae17313efd73 Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Mon, 19 Jan 2026 07:48:01 +0530 Subject: [PATCH 55/77] docs: Add Google Workload Identity Federation (WIF) documentation to Vertex AI (#19320) - Added new section documenting WIF support for Vertex AI authentication - Included SDK and Proxy configuration examples - Added sample WIF credentials file format for AWS federation - Mentioned LLM Credentials UI as an alternative for credential management - Added link to Google Cloud WIF documentation Co-authored-by: Cursor Agent --- docs/my-website/docs/providers/vertex.md | 71 ++++++++++++++++++++++++ 1 file changed, 71 insertions(+) diff --git a/docs/my-website/docs/providers/vertex.md b/docs/my-website/docs/providers/vertex.md index 33ebf535d29..be2bf86ab10 100644 --- a/docs/my-website/docs/providers/vertex.md +++ b/docs/my-website/docs/providers/vertex.md @@ -1390,6 +1390,77 @@ model_list: +### **Workload Identity Federation** + +LiteLLM supports [Google Cloud Workload Identity Federation (WIF)](https://cloud.google.com/iam/docs/workload-identity-federation), which allows you to grant on-premises or multi-cloud workloads access to Google Cloud resources without using a service account key. This is the recommended approach for workloads running in other cloud environments (AWS, Azure, etc.) or on-premises. + +To use Workload Identity Federation, pass the path to your WIF credentials configuration file via `vertex_credentials`: + + + + +```python +from litellm import completion + +response = completion( + model="vertex_ai/gemini-1.5-pro", + messages=[{"role": "user", "content": "Hello!"}], + vertex_credentials="/path/to/wif-credentials.json", # 👈 WIF credentials file + vertex_project="your-gcp-project-id", + vertex_location="us-central1" +) +``` + + + + +```yaml +model_list: + - model_name: gemini-model + litellm_params: + model: vertex_ai/gemini-1.5-pro + vertex_project: your-gcp-project-id + vertex_location: us-central1 + vertex_credentials: /path/to/wif-credentials.json # 👈 WIF credentials file +``` + +Alternatively, you can create credentials in **LLM Credentials** in the LiteLLM UI and use those to authenticate your models: + +```yaml +model_list: + - model_name: gemini-model + litellm_params: + model: vertex_ai/gemini-1.5-pro + vertex_project: your-gcp-project-id + vertex_location: us-central1 + litellm_credential_name: my-vertex-wif-credential # 👈 Reference credential stored in UI +``` + + + + +**WIF Credentials File Format** + +Your WIF credentials JSON file typically looks like this (for AWS federation): + +```json +{ + "type": "external_account", + "audience": "//iam.googleapis.com/projects/PROJECT_NUMBER/locations/global/workloadIdentityPools/POOL_ID/providers/PROVIDER_ID", + "subject_token_type": "urn:ietf:params:aws:token-type:aws4_request", + "service_account_impersonation_url": "https://iamcredentials.googleapis.com/v1/projects/-/serviceAccounts/SERVICE_ACCOUNT_EMAIL:generateAccessToken", + "token_url": "https://sts.googleapis.com/v1/token", + "credential_source": { + "environment_id": "aws1", + "region_url": "http://169.254.169.254/latest/meta-data/placement/availability-zone", + "url": "http://169.254.169.254/latest/meta-data/iam/security-credentials", + "regional_cred_verification_url": "https://sts.{region}.amazonaws.com?Action=GetCallerIdentity&Version=2011-06-15" + } +} +``` + +For more details on setting up Workload Identity Federation, see [Google Cloud WIF documentation](https://cloud.google.com/iam/docs/workload-identity-federation). + ### **Environment Variables** You can set: From 30c4a381795ee3feb76c5947769e20236468e89d Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Mon, 19 Jan 2026 12:03:26 +0900 Subject: [PATCH 56/77] test: const --- tests/mcp_tests/mcp_server.py | 2 -- tests/mcp_tests/test_proxy_mcp_e2e.py | 19 +++++++------------ 2 files changed, 7 insertions(+), 14 deletions(-) diff --git a/tests/mcp_tests/mcp_server.py b/tests/mcp_tests/mcp_server.py index c4daff4ab1c..bc6accbb721 100644 --- a/tests/mcp_tests/mcp_server.py +++ b/tests/mcp_tests/mcp_server.py @@ -1,6 +1,4 @@ # math_server.py -from __future__ import annotations - import argparse import os diff --git a/tests/mcp_tests/test_proxy_mcp_e2e.py b/tests/mcp_tests/test_proxy_mcp_e2e.py index b124fabe8b8..2b8cde54710 100644 --- a/tests/mcp_tests/test_proxy_mcp_e2e.py +++ b/tests/mcp_tests/test_proxy_mcp_e2e.py @@ -27,10 +27,7 @@ PROJECT_ROOT = Path(__file__).resolve().parents[2] PROXY_START_TIMEOUT = 30 -@pytest.fixture(scope="session") -def proxy_authorization_header() -> str: - """Shared Authorization header value for proxy calls.""" - return "Bearer sk-1234" +PROXY_AUTHORIZATION_HEADER = "Bearer sk-1234" @pytest.fixture(scope="session", autouse=True) @@ -153,14 +150,12 @@ def proxy_server_url( class TestProxyMcpSimpleConnections: @pytest.mark.asyncio - async def test_proxy_mcp_stdio_roundtrip( - self, proxy_server_url: str, proxy_authorization_header: str - ) -> None: + async def test_proxy_mcp_stdio_roundtrip(self, proxy_server_url: str) -> None: async with asyncio.timeout(20): async with streamablehttp_client( url=f"{proxy_server_url}/mcp", headers={ - "Authorization": proxy_authorization_header, + "Authorization": PROXY_AUTHORIZATION_HEADER, "x-mcp-servers": "math_stdio", }, ) as (read, write, _get_session_id): @@ -179,13 +174,13 @@ class TestProxyMcpSimpleConnections: @pytest.mark.asyncio async def test_proxy_mcp_streamable_http_roundtrip( - self, proxy_server_url: str, proxy_authorization_header: str + self, proxy_server_url: str ) -> None: async with asyncio.timeout(20): async with streamablehttp_client( url=f"{proxy_server_url}/mcp", headers={ - "Authorization": proxy_authorization_header, + "Authorization": PROXY_AUTHORIZATION_HEADER, "x-mcp-servers": "math_streamable_http", }, ) as (read, write, _get_session_id): @@ -204,12 +199,12 @@ class TestProxyMcpSimpleConnections: @pytest.mark.asyncio async def test_proxy_mcp_lists_all_servers_without_header( - self, proxy_server_url: str, proxy_authorization_header: str + self, proxy_server_url: str ) -> None: async with asyncio.timeout(20): async with streamablehttp_client( url=f"{proxy_server_url}/mcp", - headers={"Authorization": proxy_authorization_header}, + headers={"Authorization": PROXY_AUTHORIZATION_HEADER}, ) as (read, write, _get_session_id): async with ClientSession(read, write) as session: await session.initialize() From 1fbbe0a98328a4fa611a25e1bcaa32cb5ae208a0 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Mon, 19 Jan 2026 12:29:37 +0900 Subject: [PATCH 57/77] test: restore global MCP server manager after access-group test --- tests/mcp_tests/test_mcp_server.py | 51 ++++++++++++++++-------------- 1 file changed, 28 insertions(+), 23 deletions(-) diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index 66dd2bdc6b3..e8a1231c6fb 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -1001,34 +1001,39 @@ async def test_mcp_server_manager_access_groups_from_config(): MCPRequestHandler, ) - # Patch global_mcp_server_manager for this test + # Patch global_mcp_server_manager for this test and restore afterwards to + # avoid leaking state into other tests (e.g. the proxy MCP e2e suite). import litellm.proxy._experimental.mcp_server.mcp_server_manager as mcp_server_manager_mod + original_manager = mcp_server_manager_mod.global_mcp_server_manager mcp_server_manager_mod.global_mcp_server_manager = test_manager - # Should find config_server for group-a, both for group-b, other_server for group-c - import asyncio + try: + # Should find config_server for group-a, both for group-b, other_server for group-c + import asyncio - server_ids_a = await MCPRequestHandler._get_mcp_servers_from_access_groups([ - "group-a" - ]) - server_ids_b = await MCPRequestHandler._get_mcp_servers_from_access_groups([ - "group-b" - ]) - server_ids_c = await MCPRequestHandler._get_mcp_servers_from_access_groups([ - "group-c" - ]) - assert any(config_server.server_id == sid for sid in server_ids_a) - assert set(server_ids_b) == set( - [ - s.server_id + server_ids_a = await MCPRequestHandler._get_mcp_servers_from_access_groups([ + "group-a" + ]) + server_ids_b = await MCPRequestHandler._get_mcp_servers_from_access_groups([ + "group-b" + ]) + server_ids_c = await MCPRequestHandler._get_mcp_servers_from_access_groups([ + "group-c" + ]) + assert any(config_server.server_id == sid for sid in server_ids_a) + assert set(server_ids_b) == set( + [ + s.server_id + for s in test_manager.config_mcp_servers.values() + if "group-b" in s.access_groups + ] + ) + assert any( + s.name == "other_server" and s.server_id in server_ids_c for s in test_manager.config_mcp_servers.values() - if "group-b" in s.access_groups - ] - ) - assert any( - s.name == "other_server" and s.server_id in server_ids_c - for s in test_manager.config_mcp_servers.values() - ) + ) + finally: + mcp_server_manager_mod.global_mcp_server_manager = original_manager async def test_mcp_server_manager_config_integration_with_database(): From 0a15f1b66ad841b826b99171834236ef292eab64 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 19 Jan 2026 09:05:52 +0530 Subject: [PATCH 58/77] Fix: stability image optional para --- docs/my-website/docs/providers/stability.md | 52 ++++++++++- litellm/images/main.py | 6 +- .../base_llm/image_edit/transformation.py | 2 +- litellm/llms/bedrock/image_edit/handler.py | 12 +-- .../image_edit/stability_transformation.py | 93 +++++++++++-------- .../llms/gemini/image_edit/transformation.py | 12 ++- .../image_edit/dalle2_transformation.py | 22 +++-- .../llms/openai/image_edit/transformation.py | 20 ++-- .../llms/recraft/image_edit/transformation.py | 18 ++-- .../stability/image_edit/transformations.py | 9 +- .../vertex_gemini_transformation.py | 12 ++- .../vertex_imagen_transformation.py | 2 +- litellm/proxy/image_endpoints/endpoints.py | 10 +- 13 files changed, 175 insertions(+), 95 deletions(-) diff --git a/docs/my-website/docs/providers/stability.md b/docs/my-website/docs/providers/stability.md index 62a8ab43cd8..c4bc5376d1f 100644 --- a/docs/my-website/docs/providers/stability.md +++ b/docs/my-website/docs/providers/stability.md @@ -173,6 +173,14 @@ Stability AI returns images in base64 format. The response is OpenAI-compatible: Stability AI supports various image editing operations including inpainting, upscaling, outpainting, background removal, and more. +:::info Optional Parameters +**Important:** Different Stability models have different parameter requirements: +- Some models don't require a `prompt` (e.g., upscaling, background removal) +- The `style-transfer` model uses `init_image` and `style_image` instead of `image` +- The `outpaint` model requires numeric parameters (`left`, `right`, `up`, `down`) +LiteLLM automatically handles these differences for you. +::: + ### Usage - LiteLLM Python SDK #### Inpainting (Edit with Mask) @@ -217,11 +225,11 @@ response = image_edit( creativity=0.3, # 0-0.35, higher = more creative ) -# Fast upscaling - quick upscaling +# Fast upscaling - quick upscaling (no prompt needed) response = image_edit( model="stability/stable-fast-upscale-v1:0", image=open("low_res_image.png", "rb"), - prompt="Quickly upscale this image", + # No prompt required for fast upscale ) print(response) ``` @@ -259,7 +267,7 @@ os.environ['STABILITY_API_KEY'] = "your-api-key" response = image_edit( model="stability/stable-image-remove-background-v1:0", image=open("portrait.png", "rb"), - prompt="Remove the background", + # No prompt required for fast upscale ) print(response) ``` @@ -329,10 +337,29 @@ response = image_edit( model="stability/stable-image-erase-object-v1:0", image=open("scene.png", "rb"), mask=open("object_mask.png", "rb"), # Mask the object to erase - prompt="Remove the object", + # No prompt needed ) print(response) ``` +#### Style Transfer + +```python showLineNumbers +from litellm import image_edit +import os + +os.environ['STABILITY_API_KEY'] = "your-api-key" + +# Transfer style from one image to another +# Note: Uses init_image (via image param) and style_image +response = image_edit( + model="stability/stable-style-transfer-v1:0", + image=open("content_image.png", "rb"), # Maps to init_image + style_image=open("style_reference.png", "rb"), # Style to apply + fidelity=0.5, # 0-1, balance between content and style + # No prompt needed +) + +print(response) ### Supported Image Edit Models @@ -419,6 +446,23 @@ response = image_edit( ) print(response) ``` +# Fast upscale without prompt +response = image_edit( + model="bedrock/stability.stable-fast-upscale-v1:0", + image=open("low_res_image.png", "rb"), +) + +# Outpaint with numeric parameters +response = image_edit( + model="bedrock/stability.stable-outpaint-v1:0", + image=open("original_image.png", "rb"), + left=100, # Automatically converted to int + right=100, + up=50, + down=50, +) + +print(response) ### Supported Bedrock Stability Models diff --git a/litellm/images/main.py b/litellm/images/main.py index 1b09c20d350..9c2e1fa1389 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -714,8 +714,8 @@ def image_variation( @client def image_edit( # noqa: PLR0915 - image: Union[FileTypes, List[FileTypes]], - prompt: str, + image: Optional[Union[FileTypes, List[FileTypes]]], + prompt: Optional[str]= None, model: Optional[str] = None, mask: Optional[str] = None, n: Optional[int] = None, @@ -766,7 +766,7 @@ def image_edit( # noqa: PLR0915 _is_async = kwargs.pop("async_call", False) is True # add images / or return a single image - images = image if isinstance(image, list) else [image] + images = image if isinstance(image, list) else ([image] if image is not None else []) headers_from_kwargs = kwargs.get("headers") merged_extra_headers: Dict[str, Any] = {} diff --git a/litellm/llms/base_llm/image_edit/transformation.py b/litellm/llms/base_llm/image_edit/transformation.py index cc723480371..b088cdf37f6 100644 --- a/litellm/llms/base_llm/image_edit/transformation.py +++ b/litellm/llms/base_llm/image_edit/transformation.py @@ -93,7 +93,7 @@ class BaseImageEditConfig(ABC): self, model: str, prompt: Optional[str], - image: FileTypes, + image: Optional[FileTypes], image_edit_optional_request_params: Dict, litellm_params: GenericLiteLLMParams, headers: dict, diff --git a/litellm/llms/bedrock/image_edit/handler.py b/litellm/llms/bedrock/image_edit/handler.py index 0f1dcff6294..ef441fa5039 100644 --- a/litellm/llms/bedrock/image_edit/handler.py +++ b/litellm/llms/bedrock/image_edit/handler.py @@ -62,7 +62,7 @@ class BedrockImageEdit(BaseAWSLLM): self, model: str, image: list, - prompt: str, + prompt: Optional[str], model_response: ImageResponse, optional_params: dict, logging_obj: LitellmLogging, @@ -127,7 +127,7 @@ class BedrockImageEdit(BaseAWSLLM): timeout: Optional[Union[float, httpx.Timeout]], model: str, logging_obj: LitellmLogging, - prompt: str, + prompt: Optional[str], model_response: ImageResponse, client: Optional[AsyncHTTPHandler] = None, ) -> ImageResponse: @@ -163,7 +163,7 @@ class BedrockImageEdit(BaseAWSLLM): self, model: str, image: list, - prompt: str, + prompt: Optional[str], optional_params: dict, api_base: Optional[str], extra_headers: Optional[dict], @@ -176,7 +176,7 @@ class BedrockImageEdit(BaseAWSLLM): Args: model (str): The model to use for the image edit image (list): The images to edit - prompt (str): The prompt for the edit + prompt (Optional[str]): The prompt for the edit optional_params (dict): The optional parameters for the image edit api_base (Optional[str]): The base URL for the Bedrock API extra_headers (Optional[dict]): The extra headers to include in the request @@ -248,7 +248,7 @@ class BedrockImageEdit(BaseAWSLLM): self, model: str, image: list, - prompt: str, + prompt: Optional[str], optional_params: dict, ) -> dict: """ @@ -276,7 +276,7 @@ class BedrockImageEdit(BaseAWSLLM): model_response: ImageResponse, model: str, logging_obj: LitellmLogging, - prompt: str, + prompt: Optional[str], response: httpx.Response, data: dict, ) -> ImageResponse: diff --git a/litellm/llms/bedrock/image_edit/stability_transformation.py b/litellm/llms/bedrock/image_edit/stability_transformation.py index e8b77812988..c5794060272 100644 --- a/litellm/llms/bedrock/image_edit/stability_transformation.py +++ b/litellm/llms/bedrock/image_edit/stability_transformation.py @@ -154,7 +154,7 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig): self, model: str, prompt: Optional[str], - image: FileTypes, + image: Optional[FileTypes], image_edit_optional_request_params: Dict, litellm_params: GenericLiteLLMParams, headers: dict, @@ -164,32 +164,38 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig): Returns the request body dict that will be JSON-encoded by the handler. """ - if prompt is None: - raise ValueError("Bedrock Stability image edit requires a prompt.") - # Build Bedrock Stability request data: Dict[str, Any] = { - "prompt": prompt, "output_format": "png", # Default to PNG } - # Convert image to base64 - image_b64: str - if hasattr(image, 'read') and callable(getattr(image, 'read', None)): - # File-like object (e.g., BufferedReader from open()) - image_bytes = image.read() # type: ignore - image_b64 = base64.b64encode(image_bytes).decode('utf-8') # type: ignore - elif isinstance(image, bytes): - # Raw bytes - image_b64 = base64.b64encode(image).decode('utf-8') - elif isinstance(image, str): - # Already a base64 string - image_b64 = image - else: - # Try to handle as bytes - image_b64 = base64.b64encode(bytes(image)).decode('utf-8') # type: ignore + # Add prompt only if provided (some models don't require it) + if prompt is not None and prompt != "": + data["prompt"] = prompt + + # Convert image to base64 if provided + if image is not None: + image_b64: str + if hasattr(image, 'read') and callable(getattr(image, 'read', None)): + # File-like object (e.g., BufferedReader from open()) + image_bytes = image.read() # type: ignore + image_b64 = base64.b64encode(image_bytes).decode('utf-8') # type: ignore + elif isinstance(image, bytes): + # Raw bytes + image_b64 = base64.b64encode(image).decode('utf-8') + elif isinstance(image, str): + # Already a base64 string + image_b64 = image + else: + # Try to handle as bytes + image_b64 = base64.b64encode(bytes(image)).decode('utf-8') # type: ignore - data["image"] = image_b64 + # For style-transfer models, map image to init_image + model_lower = model.lower() + if "style-transfer" in model_lower: + data["init_image"] = image_b64 + else: + data["image"] = image_b64 # Add optional params (already mapped in map_openai_params) for key, value in image_edit_optional_request_params.items(): # type: ignore @@ -221,30 +227,43 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig): file_b64 = str(file_bytes) data[key] = file_b64 continue - - # Supported text fields - if key in [ - "negative_prompt", - "aspect_ratio", - "seed", - "output_format", - "model", - "mode", + + # Numeric fields that need to be converted to int/float + numeric_int_fields = ["left", "right", "up", "down", "seed"] + numeric_float_fields = [ "strength", - "style_preset", "creativity", "control_strength", "grow_mask", - "left", - "right", - "up", - "down", - "select_prompt", - "search_prompt", "fidelity", "composition_fidelity", "style_strength", "change_strength", + ] + + if key in numeric_int_fields: + # Convert to int (these are pixel values for outpaint) + try: + data[key] = int(value) # type: ignore + except (ValueError, TypeError): + data[key] = value # type: ignore + elif key in numeric_float_fields: + # Convert to float + try: + data[key] = float(value) # type: ignore + except (ValueError, TypeError): + data[key] = value # type: ignore + + # Supported text fields + elif key in [ + "negative_prompt", + "aspect_ratio", + "output_format", + "model", + "mode", + "style_preset", + "select_prompt", + "search_prompt", ]: data[key] = value # type: ignore diff --git a/litellm/llms/gemini/image_edit/transformation.py b/litellm/llms/gemini/image_edit/transformation.py index 0015155b47f..16541138217 100644 --- a/litellm/llms/gemini/image_edit/transformation.py +++ b/litellm/llms/gemini/image_edit/transformation.py @@ -81,21 +81,23 @@ class GeminiImageEditConfig(BaseImageEditConfig): self, model: str, prompt: Optional[str], - image: FileTypes, + image: Optional[FileTypes], image_edit_optional_request_params: Dict[str, Any], litellm_params: GenericLiteLLMParams, headers: dict, ) -> Tuple[Dict[str, Any], Optional[RequestFiles]]: - inline_parts = self._prepare_inline_image_parts(image) + inline_parts = self._prepare_inline_image_parts(image) if image else [] if not inline_parts: raise ValueError("Gemini image edit requires at least one image.") - if prompt is None: - raise ValueError("Gemini image edit requires a prompt.") + # Build parts list with image and prompt (if provided) + parts = inline_parts.copy() + if prompt is not None and prompt != "": + parts.append({"text": prompt}) contents = [ { - "parts": inline_parts + [{"text": prompt}], + "parts": parts, } ] diff --git a/litellm/llms/openai/image_edit/dalle2_transformation.py b/litellm/llms/openai/image_edit/dalle2_transformation.py index 13531546d2e..fd697b210ee 100644 --- a/litellm/llms/openai/image_edit/dalle2_transformation.py +++ b/litellm/llms/openai/image_edit/dalle2_transformation.py @@ -31,7 +31,7 @@ class DallE2ImageEditConfig(OpenAIImageEditConfig): self, model: str, prompt: Optional[str], - image: FileTypes, + image: Optional[FileTypes], image_edit_optional_request_params: Dict, litellm_params: GenericLiteLLMParams, headers: dict, @@ -40,18 +40,20 @@ class DallE2ImageEditConfig(OpenAIImageEditConfig): Transform image edit request for DALL-E-2. DALL-E-2 only accepts a single image with field name "image" (not "image[]"). - """ - if prompt is None: - raise ValueError("DALL-E-2 image edit requires a prompt.") - - request = ImageEditRequestParams( - model=model, - image=image, - prompt=prompt, + """ + request_params = { + "model": model, **image_edit_optional_request_params, - ) + } + if image is not None: + request_params["image"] = image + if prompt is not None: + request_params["prompt"] = prompt + + request = ImageEditRequestParams(**request_params) request_dict = cast(Dict, request) + ######################################################### # Separate images and masks as `files` and send other parameters as `data` ######################################################### diff --git a/litellm/llms/openai/image_edit/transformation.py b/litellm/llms/openai/image_edit/transformation.py index 9edad9ee2c9..a1e5375d098 100644 --- a/litellm/llms/openai/image_edit/transformation.py +++ b/litellm/llms/openai/image_edit/transformation.py @@ -80,7 +80,7 @@ class OpenAIImageEditConfig(BaseImageEditConfig): self, model: str, prompt: Optional[str], - image: FileTypes, + image: Optional[FileTypes], image_edit_optional_request_params: Dict, litellm_params: GenericLiteLLMParams, headers: dict, @@ -91,15 +91,17 @@ class OpenAIImageEditConfig(BaseImageEditConfig): Handles multipart/form-data for images. Uses "image[]" field name to support multiple images (e.g., for gpt-image-1). """ - if prompt is None: - raise ValueError("OpenAI image edit requires a prompt.") - - request = ImageEditRequestParams( - model=model, - image=image, - prompt=prompt, + # Build request params, only including non-None values + request_params = { + "model": model, **image_edit_optional_request_params, - ) + } + if image is not None: + request_params["image"] = image + if prompt is not None: + request_params["prompt"] = prompt + + request = ImageEditRequestParams(**request_params) request_dict = cast(Dict, request) ######################################################### diff --git a/litellm/llms/recraft/image_edit/transformation.py b/litellm/llms/recraft/image_edit/transformation.py index 9bf46704ed1..d2a56236819 100644 --- a/litellm/llms/recraft/image_edit/transformation.py +++ b/litellm/llms/recraft/image_edit/transformation.py @@ -102,7 +102,7 @@ class RecraftImageEditConfig(BaseImageEditConfig): self, model: str, prompt: Optional[str], - image: FileTypes, + image: Optional[FileTypes], image_edit_optional_request_params: Dict, litellm_params: GenericLiteLLMParams, headers: dict, @@ -114,15 +114,15 @@ class RecraftImageEditConfig(BaseImageEditConfig): https://www.recraft.ai/docs#image-to-image """ - if prompt is None: - raise ValueError("Recraft image edit requires a prompt.") - - request_body: RecraftImageEditRequestParams = RecraftImageEditRequestParams( - model=model, - prompt=prompt, - strength=image_edit_optional_request_params.pop("strength", self.DEFAULT_STRENGTH), + request_params = { + "model": model, + "strength": image_edit_optional_request_params.pop("strength", self.DEFAULT_STRENGTH), **image_edit_optional_request_params, - ) + } + if prompt is not None: + request_params["prompt"] = prompt + + request_body = RecraftImageEditRequestParams(**request_params) request_dict = cast(Dict, request_body) ######################################################### # Reuse OpenAI logic: Separate images as `files` and send other parameters as `data` diff --git a/litellm/llms/stability/image_edit/transformations.py b/litellm/llms/stability/image_edit/transformations.py index 013e3f27a02..53bdc825dd4 100644 --- a/litellm/llms/stability/image_edit/transformations.py +++ b/litellm/llms/stability/image_edit/transformations.py @@ -171,7 +171,7 @@ class StabilityImageEditConfig(BaseImageEditConfig): self, model: str, prompt: Optional[str], - image: FileTypes, + image: Optional[FileTypes], image_edit_optional_request_params: Dict, litellm_params: GenericLiteLLMParams, headers: dict, @@ -190,11 +190,14 @@ class StabilityImageEditConfig(BaseImageEditConfig): } # Add prompt only if provided (some Stability endpoints don't require it) - if prompt is not None: + if prompt is not None and prompt != "": data["prompt"] = prompt # Handle image parameter - could be a single file or list image_file = image[0] if isinstance(image, list) else image # type: ignore - files: Dict[str, Any] = {"image": image_file} + files: Dict[str, Any] = {} + if image is not None: + image_file = image[0] if isinstance(image, list) else image # type: ignore + files["image"] = image_file # Add optional params (already mapped in map_openai_params) for key, value in image_edit_optional_request_params.items(): # type: ignore diff --git a/litellm/llms/vertex_ai/image_edit/vertex_gemini_transformation.py b/litellm/llms/vertex_ai/image_edit/vertex_gemini_transformation.py index 154d5669eb8..8fcd285824d 100644 --- a/litellm/llms/vertex_ai/image_edit/vertex_gemini_transformation.py +++ b/litellm/llms/vertex_ai/image_edit/vertex_gemini_transformation.py @@ -152,22 +152,24 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM): self, model: str, prompt: Optional[str], - image: FileTypes, + image: Optional[FileTypes], image_edit_optional_request_params: Dict[str, Any], litellm_params: GenericLiteLLMParams, headers: dict, ) -> Tuple[Dict[str, Any], Optional[RequestFiles]]: - inline_parts = self._prepare_inline_image_parts(image) + inline_parts = self._prepare_inline_image_parts(image) if image else [] if not inline_parts: raise ValueError("Vertex AI Gemini image edit requires at least one image.") - if prompt is None: - raise ValueError("Vertex AI Gemini image edit requires a prompt.") + # Build parts list with image and prompt (if provided) + parts = inline_parts.copy() + if prompt is not None and prompt != "": + parts.append({"text": prompt}) # Correct format for Vertex AI Gemini image editing contents = { "role": "USER", - "parts": inline_parts + [{"text": prompt}] + "parts": parts } request_body: Dict[str, Any] = {"contents": contents} diff --git a/litellm/llms/vertex_ai/image_edit/vertex_imagen_transformation.py b/litellm/llms/vertex_ai/image_edit/vertex_imagen_transformation.py index 337a4bd4dd6..b58825e1faa 100644 --- a/litellm/llms/vertex_ai/image_edit/vertex_imagen_transformation.py +++ b/litellm/llms/vertex_ai/image_edit/vertex_imagen_transformation.py @@ -144,7 +144,7 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): self, model: str, prompt: Optional[str], - image: FileTypes, + image: Optional[FileTypes], image_edit_optional_request_params: Dict[str, Any], litellm_params: GenericLiteLLMParams, headers: dict, diff --git a/litellm/proxy/image_endpoints/endpoints.py b/litellm/proxy/image_endpoints/endpoints.py index a1453e10dbf..4a2c05f8590 100644 --- a/litellm/proxy/image_endpoints/endpoints.py +++ b/litellm/proxy/image_endpoints/endpoints.py @@ -244,8 +244,10 @@ async def image_edit_api( if mask is None and mask_array is not None: mask = mask_array - if image is None: - raise HTTPException(status_code=422, detail="Field required: image") + # if image is None: + # raise HTTPException(status_code=422, detail="Field required: image") + # Note: Image is optional for some models (e.g., Bedrock Stability style-transfer) + # The validation will be done at the model level if image is truly required from litellm.proxy.proxy_server import ( _read_request_body, @@ -272,6 +274,10 @@ async def image_edit_api( data["image"] = image_files if mask_files: data["mask"] = mask_files + + # Ensure prompt exists in data (default to None for models that don't require it) + if "prompt" not in data: + data["prompt"] = None data["model"] = ( model From 49b2886ef0f044f29d2938ee6b4fd07d253e42e0 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 19 Jan 2026 09:08:38 +0530 Subject: [PATCH 59/77] Add None as default image value --- litellm/images/main.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/images/main.py b/litellm/images/main.py index 9c2e1fa1389..6c4c502a7b0 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -714,7 +714,7 @@ def image_variation( @client def image_edit( # noqa: PLR0915 - image: Optional[Union[FileTypes, List[FileTypes]]], + image: Optional[Union[FileTypes, List[FileTypes]]] = None, prompt: Optional[str]= None, model: Optional[str] = None, mask: Optional[str] = None, From e1b73310ca22f2190e36c8e6bbbcc951ea668cbc Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 19 Jan 2026 09:16:12 +0530 Subject: [PATCH 60/77] Fix: llms/bedrock/image_edit/stability_transformation.py:153:9: PLR0915 Too many statements (55 > 50) --- litellm/llms/bedrock/image_edit/stability_transformation.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/llms/bedrock/image_edit/stability_transformation.py b/litellm/llms/bedrock/image_edit/stability_transformation.py index c5794060272..fc14b571a8c 100644 --- a/litellm/llms/bedrock/image_edit/stability_transformation.py +++ b/litellm/llms/bedrock/image_edit/stability_transformation.py @@ -150,7 +150,7 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig): return mapped_params - def transform_image_edit_request( + def transform_image_edit_request( #noqa: PLR0915 self, model: str, prompt: Optional[str], From fbf2d8337586657c71c343371715a415e1caaaed Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 19 Jan 2026 10:19:31 +0530 Subject: [PATCH 61/77] Fix: _handle_failure method getting called 2 times --- litellm/responses/streaming_iterator.py | 1 - ...t_base_responses_api_streaming_iterator.py | 50 ++++++++++++++++++- 2 files changed, 49 insertions(+), 2 deletions(-) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 0b838f916e2..540d9ad2642 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -170,7 +170,6 @@ class BaseResponsesAPIStreamingIterator: return None except Exception as e: # Ensure failures trigger failure hooks - self._handle_failure(e) raise def _handle_logging_completed_response(self): diff --git a/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py b/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py index 1103eaf92be..4a423305626 100644 --- a/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py +++ b/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py @@ -246,4 +246,52 @@ class TestBaseResponsesAPIStreamingIterator: # Test with None chunk result = iterator._process_chunk(None) - assert result is None \ No newline at end of file + assert result is None + + def test_process_chunk_exception_does_not_call_handle_failure(self): + """ + Test that _process_chunk raises exceptions without calling _handle_failure. + + This ensures _handle_failure is only called once in the outer exception handler + (in __next__ or __anext__), preventing duplicate failure logging. + + Previously, _handle_failure was called both in _process_chunk and in the outer + exception handler, causing duplicate logs. This test verifies the fix. + """ + # Mock dependencies + mock_response = Mock() + mock_response.headers = {} + mock_logging_obj = Mock(spec=LiteLLMLoggingObj) + mock_logging_obj.model_call_details = {"litellm_params": {}} + mock_config = Mock(spec=BaseResponsesAPIConfig) + + # Set up the mock transform method to raise an exception + test_exception = ValueError("Test exception in transform") + mock_config.transform_streaming_response.side_effect = test_exception + + # Create the iterator instance + iterator = BaseResponsesAPIStreamingIterator( + response=mock_response, + model="gpt-4", + responses_api_provider_config=mock_config, + logging_obj=mock_logging_obj + ) + + # Mock _handle_failure to track if it's called + with patch.object(iterator, '_handle_failure') as mock_handle_failure: + # Prepare valid JSON chunk that will trigger transform_streaming_response + test_chunk_data = { + "type": "response.output_text.delta", + "delta": "Hello" + } + + # _process_chunk should raise the exception without calling _handle_failure + with pytest.raises(ValueError) as exc_info: + iterator._process_chunk(json.dumps(test_chunk_data)) + + # Verify the exception was raised + assert str(exc_info.value) == "Test exception in transform" + + # Verify _handle_failure was NOT called in _process_chunk + # It should only be called by the outer exception handler in __next__/__anext__ + mock_handle_failure.assert_not_called() \ No newline at end of file From a141aa6026b911f6b74395a8b4b0452924572616 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Mon, 19 Jan 2026 13:57:40 +0900 Subject: [PATCH 62/77] test: temporary skip --- tests/mcp_tests/test_proxy_mcp_e2e.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/mcp_tests/test_proxy_mcp_e2e.py b/tests/mcp_tests/test_proxy_mcp_e2e.py index 2b8cde54710..cd073d79bf3 100644 --- a/tests/mcp_tests/test_proxy_mcp_e2e.py +++ b/tests/mcp_tests/test_proxy_mcp_e2e.py @@ -198,6 +198,7 @@ class TestProxyMcpSimpleConnections: assert text == "11" @pytest.mark.asyncio + @pytest.mark.skip async def test_proxy_mcp_lists_all_servers_without_header( self, proxy_server_url: str ) -> None: From 44a166a79274c90811b92de8829ef9a434ee614c Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Mon, 19 Jan 2026 14:27:00 +0900 Subject: [PATCH 63/77] fix: ci mcp version up --- .circleci/config.yml | 2 +- tests/mcp_tests/test_proxy_mcp_e2e.py | 1 - 2 files changed, 1 insertion(+), 2 deletions(-) diff --git a/.circleci/config.yml b/.circleci/config.yml index 133a7184f9b..e03d1086282 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -1153,7 +1153,7 @@ jobs: pip install "pytest-asyncio==0.21.1" pip install "respx==0.22.0" pip install "pydantic==2.10.2" - pip install "mcp==1.10.1" + pip install "mcp==1.21.2" # Run pytest and generate JUnit XML report - run: name: Run tests diff --git a/tests/mcp_tests/test_proxy_mcp_e2e.py b/tests/mcp_tests/test_proxy_mcp_e2e.py index cd073d79bf3..2b8cde54710 100644 --- a/tests/mcp_tests/test_proxy_mcp_e2e.py +++ b/tests/mcp_tests/test_proxy_mcp_e2e.py @@ -198,7 +198,6 @@ class TestProxyMcpSimpleConnections: assert text == "11" @pytest.mark.asyncio - @pytest.mark.skip async def test_proxy_mcp_lists_all_servers_without_header( self, proxy_server_url: str ) -> None: From 480cb9c0d8f9696e4bdab366811b6edf4766ba23 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 19 Jan 2026 11:58:32 +0530 Subject: [PATCH 64/77] Fix: upload pdfs for file endpoint --- litellm/llms/custom_httpx/llm_http_handler.py | 6 +- .../test_vertex_ai_binary_file_upload.py | 260 ++++++++++++++++++ 2 files changed, 262 insertions(+), 4 deletions(-) create mode 100644 tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_binary_file_upload.py diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index ab1e735fca7..6a87967c3aa 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -3080,10 +3080,8 @@ class BaseLLMHTTPHandler: transformed_request, bytes ): # Handle traditional file uploads - # Ensure transformed_request is a string for httpx compatibility - if isinstance(transformed_request, bytes): - transformed_request = transformed_request.decode("utf-8") - + # Note: transformed_request can be bytes (for binary files like PDFs) + # or str (for text files like JSONL). httpx handles both correctly. # Use the HTTP method specified by the provider config http_method = provider_config.file_upload_http_method.upper() if http_method == "PUT": diff --git a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_binary_file_upload.py b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_binary_file_upload.py new file mode 100644 index 00000000000..ceea3d0b16c --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_binary_file_upload.py @@ -0,0 +1,260 @@ +""" +Test Vertex AI binary file upload functionality + +This test ensures that binary files (like PDFs, images) are correctly handled +during upload without attempting UTF-8 decoding, which would cause errors. + +Regression test for: UTF-8 codec error when uploading binary files +""" + +import io +import pytest +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx + +from litellm.llms.custom_httpx.llm_http_handler import AsyncHTTPHandler +from litellm.llms.vertex_ai.files.transformation import VertexAIFilesConfig +from litellm.types.llms.openai import CreateFileRequest + + +class TestVertexAIBinaryFileUpload: + """Test binary file upload handling for Vertex AI""" + + def setup_method(self): + """Setup test method""" + self.http_handler = AsyncHTTPHandler() + self.vertex_config = VertexAIFilesConfig() + + @pytest.mark.asyncio + async def test_pdf_file_upload_bytes_handling(self): + """ + Test that PDF binary data is correctly handled without UTF-8 decoding. + + This is a regression test for the error: + 'utf-8' codec can't decode byte 0xc4 in position 10: invalid continuation byte + """ + # Create mock PDF binary data (with non-UTF-8 bytes) + # PDF files start with %PDF- and contain binary data + mock_pdf_content = b"%PDF-1.4\n%\xc4\xe5\xf2\xe5\xeb\xa7\xf3\xa0\xd0\xc4\xc6\n" + mock_pdf_content += b"\x00\x01\x02\x03\xff\xfe\xfd" * 100 # Add more binary data + + # Create file object + file_obj = io.BytesIO(mock_pdf_content) + file_obj.name = "test_document.pdf" + + # Create file request + create_file_data: CreateFileRequest = { + "file": file_obj, + "purpose": "user_data", + } + + # Transform the request + transformed_request = self.vertex_config.transform_create_file_request( + model="vertex_ai/gemini-flash", + create_file_data=create_file_data, + optional_params={}, + litellm_params={}, + ) + + # Verify the transformation returns bytes (not string) + assert isinstance(transformed_request, bytes), ( + f"Expected bytes for binary file, got {type(transformed_request)}" + ) + + # Verify the bytes match the original content + assert transformed_request == mock_pdf_content, ( + "Transformed request should preserve binary content exactly" + ) + + # Verify that the bytes contain non-UTF-8 characters + # This should raise UnicodeDecodeError if we try to decode + with pytest.raises(UnicodeDecodeError): + transformed_request.decode("utf-8") + + @pytest.mark.asyncio + async def test_image_file_upload_bytes_handling(self): + """Test that image binary data (PNG) is correctly handled""" + # Create mock PNG binary data (PNG signature + some binary data) + mock_png_content = b"\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR" + mock_png_content += b"\x00\x01\x02\x03\xff\xfe\xfd" * 50 + + file_obj = io.BytesIO(mock_png_content) + file_obj.name = "test_image.png" + + create_file_data: CreateFileRequest = { + "file": file_obj, + "purpose": "user_data", + } + + transformed_request = self.vertex_config.transform_create_file_request( + model="vertex_ai/gemini-flash", + create_file_data=create_file_data, + optional_params={}, + litellm_params={}, + ) + + # Verify bytes are preserved + assert isinstance(transformed_request, bytes) + assert transformed_request == mock_png_content + + @pytest.mark.asyncio + async def test_http_handler_accepts_bytes_without_decoding(self): + """ + Test that httpx correctly accepts binary data without decoding. + + This test verifies that bytes can be passed to httpx's post/put methods + without needing UTF-8 decoding, which is the core of our fix. + """ + # Create mock binary data with non-UTF-8 bytes + mock_binary_data = b"\x00\x01\x02\x03\xff\xfe\xfd\xc4\xe5\xf2" + + # Test that httpx accepts bytes in the data parameter + # We're testing the behavior, not making an actual request + + # Verify that attempting to decode would fail (proving it's binary) + with pytest.raises(UnicodeDecodeError): + mock_binary_data.decode("utf-8") + + # Verify that httpx Request accepts bytes + try: + request = httpx.Request( + method="POST", + url="https://example.com/upload", + data=mock_binary_data, + headers={"Content-Type": "application/octet-stream"}, + ) + # If we get here, httpx accepts bytes - which is what we need + assert request.content == mock_binary_data + except Exception as e: + pytest.fail(f"httpx should accept bytes in data parameter: {e}") + + # Document the expected behavior + assert isinstance(mock_binary_data, bytes), ( + "Binary file data should remain as bytes" + ) + + @pytest.mark.asyncio + async def test_jsonl_file_upload_returns_string(self): + """ + Test that JSONL files (text) are correctly transformed to strings. + + This ensures we handle both binary and text files correctly. + """ + # Create mock JSONL content + mock_jsonl_content = ( + '{"custom_id": "req-1", "method": "POST", "url": "/v1/chat/completions", ' + '"body": {"model": "gemini-flash", "messages": [{"role": "user", "content": "Hello"}]}}\n' + ) + + file_obj = io.BytesIO(mock_jsonl_content.encode("utf-8")) + file_obj.name = "batch_requests.jsonl" + + create_file_data: CreateFileRequest = { + "file": file_obj, + "purpose": "batch", + } + + transformed_request = self.vertex_config.transform_create_file_request( + model="vertex_ai/gemini-flash", + create_file_data=create_file_data, + optional_params={}, + litellm_params={}, + ) + + # JSONL files should be transformed to string + assert isinstance(transformed_request, str), ( + f"Expected string for JSONL file, got {type(transformed_request)}" + ) + + @pytest.mark.asyncio + async def test_mixed_file_types_in_sequence(self): + """ + Test uploading different file types in sequence to ensure no state pollution. + """ + # Test 1: Upload binary file + binary_content = b"\x00\x01\x02\x03\xff\xfe\xfd" + binary_file = io.BytesIO(binary_content) + binary_file.name = "binary.dat" + + binary_request: CreateFileRequest = { + "file": binary_file, + "purpose": "user_data", + } + + result1 = self.vertex_config.transform_create_file_request( + model="vertex_ai/gemini-flash", + create_file_data=binary_request, + optional_params={}, + litellm_params={}, + ) + assert isinstance(result1, bytes) + + # Test 2: Upload JSONL file + jsonl_content = '{"test": "data"}\n' + jsonl_file = io.BytesIO(jsonl_content.encode("utf-8")) + jsonl_file.name = "batch.jsonl" + + jsonl_request: CreateFileRequest = { + "file": jsonl_file, + "purpose": "batch", + } + + result2 = self.vertex_config.transform_create_file_request( + model="vertex_ai/gemini-flash", + create_file_data=jsonl_request, + optional_params={}, + litellm_params={}, + ) + assert isinstance(result2, str) + + # Test 3: Upload another binary file + binary_content2 = b"\xc4\xe5\xf2\xe5\xeb" + binary_file2 = io.BytesIO(binary_content2) + binary_file2.name = "binary2.dat" + + binary_request2: CreateFileRequest = { + "file": binary_file2, + "purpose": "user_data", + } + + result3 = self.vertex_config.transform_create_file_request( + model="vertex_ai/gemini-flash", + create_file_data=binary_request2, + optional_params={}, + litellm_params={}, + ) + assert isinstance(result3, bytes) + + def test_bytes_type_preservation_documentation(self): + """ + Documentation test: Verify that bytes are the correct type for binary uploads. + + This test documents the expected behavior: + - Binary files (PDF, images, etc.) should remain as bytes + - Text files (JSONL) should be strings + - httpx accepts both bytes and strings in the 'data' parameter + - bytes should NEVER be decoded to UTF-8 for binary files + """ + # This is a documentation test - it always passes + # but serves as a reference for the expected behavior + + expected_behavior = { + "binary_files": { + "input_type": "bytes", + "output_type": "bytes", + "examples": ["PDF", "PNG", "JPEG", "binary data"], + "http_method": "POST or PUT", + "encoding": "none - preserve raw bytes", + }, + "text_files": { + "input_type": "str or bytes", + "output_type": "str", + "examples": ["JSONL", "CSV", "TXT"], + "http_method": "POST", + "encoding": "UTF-8", + }, + } + + assert expected_behavior["binary_files"]["encoding"] == "none - preserve raw bytes" + assert expected_behavior["text_files"]["encoding"] == "UTF-8" From 514ebb0d96c13967ab040611cf426f92157992b0 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 19 Jan 2026 13:17:08 +0530 Subject: [PATCH 65/77] Fix: vector store sync issues --- .../management_endpoints.py | 152 ++++++++++++--- .../test_vector_store_endpoints.py | 183 ++++++++++++++++++ 2 files changed, 307 insertions(+), 28 deletions(-) diff --git a/litellm/proxy/vector_store_endpoints/management_endpoints.py b/litellm/proxy/vector_store_endpoints/management_endpoints.py index 661f94e5f04..bc61a60fe5a 100644 --- a/litellm/proxy/vector_store_endpoints/management_endpoints.py +++ b/litellm/proxy/vector_store_endpoints/management_endpoints.py @@ -245,6 +245,7 @@ async def list_vector_stores( """ List all available vector stores with optional filtering and pagination. Combines both in-memory vector stores and those stored in the database. + Database is the source of truth - deleted stores are removed from memory, updated stores sync to memory. Parameters: - page: int - Page number for pagination (default: 1) @@ -252,29 +253,65 @@ async def list_vector_stores( """ from litellm.proxy.proxy_server import prisma_client - seen_vector_store_ids = set() + vector_store_map: Dict[str, LiteLLM_ManagedVectorStore] = {} + db_vector_store_ids: set = set() try: - # Get in-memory vector stores - in_memory_vector_stores: List[LiteLLM_ManagedVectorStore] = [] + # Get vector stores from database first (source of truth) + vector_stores_from_db = await VectorStoreRegistry._get_vector_stores_from_db( + prisma_client=prisma_client + ) + + # Build map from database vector stores + for vector_store in vector_stores_from_db: + vector_store_id = vector_store.get("vector_store_id", None) + if vector_store_id: + vector_store_map[vector_store_id] = vector_store + db_vector_store_ids.add(vector_store_id) + + # Process in-memory vector stores if litellm.vector_store_registry is not None: in_memory_vector_stores = copy.deepcopy( litellm.vector_store_registry.vector_stores ) + + vector_stores_to_delete_from_memory: List[str] = [] + + for vector_store in in_memory_vector_stores: + vector_store_id = vector_store.get("vector_store_id", None) + if not vector_store_id: + continue + + # If vector store is in memory but NOT in database, it was deleted + if vector_store_id not in db_vector_store_ids: + verbose_proxy_logger.info( + f"Vector store {vector_store_id} exists in memory but not in database - marking for deletion from cache" + ) + vector_stores_to_delete_from_memory.append(vector_store_id) + # If not in our map yet, add it (only in-memory, not in DB) + elif vector_store_id not in vector_store_map: + vector_store_map[vector_store_id] = vector_store + + # Synchronize in-memory registry with database + # 1. Remove deleted vector stores from memory + for vs_id in vector_stores_to_delete_from_memory: + litellm.vector_store_registry.delete_vector_store_from_registry( + vector_store_id=vs_id + ) + verbose_proxy_logger.debug( + f"Removed deleted vector store {vs_id} from in-memory registry" + ) + + # 2. Update in-memory registry with database versions (for updates) + for vector_store in vector_stores_from_db: + vector_store_id = vector_store.get("vector_store_id", None) + if vector_store_id: + litellm.vector_store_registry.update_vector_store_in_registry( + vector_store_id=vector_store_id, + updated_data=vector_store + ) - # Get vector stores from database - vector_stores_from_db = await VectorStoreRegistry._get_vector_stores_from_db( - prisma_client=prisma_client - ) - - # Combine in-memory and database vector stores - combined_vector_stores: List[LiteLLM_ManagedVectorStore] = [] - for vector_store in in_memory_vector_stores + vector_stores_from_db: - vector_store_id = vector_store.get("vector_store_id", None) - if vector_store_id not in seen_vector_store_ids: - combined_vector_stores.append(vector_store) - seen_vector_store_ids.add(vector_store_id) - + combined_vector_stores = list(vector_store_map.values()) total_count = len(combined_vector_stores) total_pages = (total_count + page_size - 1) // page_size @@ -303,7 +340,7 @@ async def delete_vector_store( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ - Delete a vector store. + Delete a vector store from both database and in-memory registry. Parameters: - vector_store_id: str - ID of the vector store to delete @@ -314,31 +351,53 @@ async def delete_vector_store( raise HTTPException(status_code=500, detail="Database not connected") try: - # Check if vector store exists + # Check if vector store exists in database or in-memory registry + db_vector_store_exists = False + memory_vector_store_exists = False + existing_vector_store = ( await prisma_client.db.litellm_managedvectorstorestable.find_unique( where={"vector_store_id": data.vector_store_id} ) ) - if existing_vector_store is None: + if existing_vector_store is not None: + db_vector_store_exists = True + + # Check in-memory registry + if litellm.vector_store_registry is not None: + memory_vector_store = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry( + vector_store_id=data.vector_store_id + ) + if memory_vector_store is not None: + memory_vector_store_exists = True + + # If not found in either location, raise 404 + if not db_vector_store_exists and not memory_vector_store_exists: raise HTTPException( status_code=404, detail=f"Vector store with ID {data.vector_store_id} not found", ) - # Delete vector store - await prisma_client.db.litellm_managedvectorstorestable.delete( - where={"vector_store_id": data.vector_store_id} - ) + # Delete from database if exists + if db_vector_store_exists: + await prisma_client.db.litellm_managedvectorstorestable.delete( + where={"vector_store_id": data.vector_store_id} + ) - # Delete vector store from registry - if litellm.vector_store_registry is not None: + # Delete from in-memory registry if exists + if memory_vector_store_exists and litellm.vector_store_registry is not None: litellm.vector_store_registry.delete_vector_store_from_registry( vector_store_id=data.vector_store_id ) - return {"message": f"Vector store {data.vector_store_id} deleted successfully"} + return { + "status": "success", + "message": f"Vector store {data.vector_store_id} deleted successfully" + } + except HTTPException: + raise except Exception as e: + verbose_proxy_logger.exception(f"Error deleting vector store: {str(e)}") raise HTTPException(status_code=500, detail=str(e)) @@ -415,8 +474,12 @@ async def update_vector_store( data: VectorStoreUpdateRequest, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): - """Update vector store details""" + """ + Update vector store details in both database and in-memory registry. + The updated data is immediately synchronized to the in-memory registry. + """ from litellm.proxy.proxy_server import prisma_client + from litellm.types.router import GenericLiteLLMParams if prisma_client is None: raise HTTPException(status_code=500, detail="Database not connected") @@ -424,11 +487,36 @@ async def update_vector_store( try: update_data = data.model_dump(exclude_unset=True) vector_store_id = update_data.pop("vector_store_id") + + # Handle metadata serialization if update_data.get("vector_store_metadata") is not None: update_data["vector_store_metadata"] = safe_dumps( update_data["vector_store_metadata"] ) + + # Handle litellm_params if provided + if "litellm_params" in update_data: + _input_litellm_params: dict = update_data.get("litellm_params", {}) or {} + + # Auto-resolve embedding config if embedding model is provided but config is not + embedding_model = _input_litellm_params.get("litellm_embedding_model") + if embedding_model and not _input_litellm_params.get("litellm_embedding_config"): + resolved_config = await _resolve_embedding_config_from_db( + embedding_model=embedding_model, + prisma_client=prisma_client + ) + if resolved_config: + _input_litellm_params["litellm_embedding_config"] = resolved_config + verbose_proxy_logger.info( + f"Auto-resolved embedding config for model {embedding_model}" + ) + + litellm_params_dict = GenericLiteLLMParams( + **_input_litellm_params + ).model_dump(exclude_none=True) + update_data["litellm_params"] = safe_dumps(litellm_params_dict) + # Update in database updated = await prisma_client.db.litellm_managedvectorstorestable.update( where={"vector_store_id": vector_store_id}, data=update_data, @@ -436,13 +524,21 @@ async def update_vector_store( updated_vs = LiteLLM_ManagedVectorStore(**updated.model_dump()) + # Immediately update in-memory registry to keep it in sync if litellm.vector_store_registry is not None: litellm.vector_store_registry.update_vector_store_in_registry( vector_store_id=vector_store_id, updated_data=updated_vs, ) + verbose_proxy_logger.debug( + f"Updated vector store {vector_store_id} in both database and in-memory registry" + ) - return {"vector_store": updated_vs} + return { + "status": "success", + "message": f"Vector store {vector_store_id} updated successfully", + "vector_store": updated_vs + } except Exception as e: verbose_proxy_logger.exception(f"Error updating vector store: {str(e)}") raise HTTPException(status_code=500, detail=str(e)) diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py index 352e84719f1..558fe18ae38 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py @@ -1051,6 +1051,189 @@ async def test_vector_store_synchronization_across_instances(): ) +@pytest.mark.asyncio +async def test_vector_store_update_and_list_synchronization(): + """ + Test that vector store updates are properly synchronized across multiple instances. + + This test simulates the scenario where: + 1. Instance 1 creates a vector store + 2. Instance 2 caches it in memory + 3. Instance 1 updates the vector store in the database + 4. Instance 2 should see the updated data when listing (database is source of truth) + + This is a regression test to prevent the bug where Instance 2 would show + stale cached data instead of the updated database version. + """ + from datetime import datetime, timezone + from unittest.mock import AsyncMock, MagicMock + + from litellm.types.vector_stores import LiteLLM_ManagedVectorStore + from litellm.vector_stores.vector_store_registry import VectorStoreRegistry + + # Simulate two instances with separate in-memory registries + instance_1_registry = VectorStoreRegistry(vector_stores=[]) + instance_2_registry = VectorStoreRegistry(vector_stores=[]) + + # Mock database that both instances share + mock_db_vector_stores = [] + + async def mock_find_many(order=None): + """Mock find_many for listing vector stores""" + result = [] + for vs in mock_db_vector_stores: + class MockVectorStore: + def __init__(self, data): + for key, value in data.items(): + setattr(self, key, value) + self._data = data + + def __iter__(self): + return iter(self._data.items()) + result.append(MockVectorStore(vs)) + return result + + async def mock_create(data): + """Mock create for adding vector store to DB""" + vector_store = data.copy() + mock_db_vector_stores.append(vector_store) + mock_obj = MagicMock() + mock_obj.model_dump.return_value = vector_store + return mock_obj + + async def mock_update(where, data): + """Mock update for modifying vector store in DB""" + vector_store_id = where.get("vector_store_id") + for i, vs in enumerate(mock_db_vector_stores): + if vs.get("vector_store_id") == vector_store_id: + # Update the vector store + mock_db_vector_stores[i].update(data) + mock_obj = MagicMock() + mock_obj.model_dump.return_value = mock_db_vector_stores[i] + return mock_obj + raise Exception(f"Vector store {vector_store_id} not found") + + # Create mock prisma client + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_managedvectorstorestable.find_many = AsyncMock( + side_effect=mock_find_many + ) + mock_prisma_client.db.litellm_managedvectorstorestable.create = AsyncMock( + side_effect=mock_create + ) + mock_prisma_client.db.litellm_managedvectorstorestable.update = AsyncMock( + side_effect=mock_update + ) + + # Test vector store data + test_vector_store_id = "test-update-store-001" + original_name = "Original Name" + updated_name = "Updated Name" + + test_vector_store: LiteLLM_ManagedVectorStore = { + "vector_store_id": test_vector_store_id, + "custom_llm_provider": "bedrock", + "vector_store_name": original_name, + "vector_store_description": "Testing update synchronization", + "litellm_params": { + "vector_store_id": test_vector_store_id, + "custom_llm_provider": "bedrock", + "region_name": "us-east-1" + }, + "created_at": datetime.now(timezone.utc), + "updated_at": datetime.now(timezone.utc), + } + + # Step 1: Create vector store on Instance 1 + await mock_prisma_client.db.litellm_managedvectorstorestable.create( + data=test_vector_store + ) + instance_1_registry.add_vector_store_to_registry(vector_store=test_vector_store) + + # Step 2: Instance 2 fetches and caches the vector store + vector_stores_from_db = await VectorStoreRegistry._get_vector_stores_from_db( + prisma_client=mock_prisma_client + ) + for vs in vector_stores_from_db: + if vs.get("vector_store_id") == test_vector_store_id: + instance_2_registry.add_vector_store_to_registry(vector_store=vs) + + # Verify both instances have the original data + instance_1_vs = instance_1_registry.get_litellm_managed_vector_store_from_registry( + test_vector_store_id + ) + instance_2_vs = instance_2_registry.get_litellm_managed_vector_store_from_registry( + test_vector_store_id + ) + assert instance_1_vs.get("vector_store_name") == original_name + assert instance_2_vs.get("vector_store_name") == original_name + + # Step 3: Instance 1 updates the vector store in the database + # (Simulating what happens in update_vector_store endpoint) + update_data = {"vector_store_name": updated_name} + await mock_prisma_client.db.litellm_managedvectorstorestable.update( + where={"vector_store_id": test_vector_store_id}, + data=update_data + ) + + # Instance 1 updates its own cache + updated_vs_instance_1 = test_vector_store.copy() + updated_vs_instance_1["vector_store_name"] = updated_name + instance_1_registry.update_vector_store_in_registry( + vector_store_id=test_vector_store_id, + updated_data=updated_vs_instance_1 + ) + + # Verify Instance 1 has the updated data + instance_1_vs_after_update = instance_1_registry.get_litellm_managed_vector_store_from_registry( + test_vector_store_id + ) + assert instance_1_vs_after_update.get("vector_store_name") == updated_name + + # Verify Instance 2 still has stale data in cache + instance_2_vs_before_list = instance_2_registry.get_litellm_managed_vector_store_from_registry( + test_vector_store_id + ) + assert instance_2_vs_before_list.get("vector_store_name") == original_name, ( + "Instance 2 should still have stale cached data before list operation" + ) + + # Step 4: Instance 2 calls list endpoint (which should sync with database) + # This simulates what list_vector_stores endpoint does + vector_stores_from_db_after_update = await VectorStoreRegistry._get_vector_stores_from_db( + prisma_client=mock_prisma_client + ) + + # Build map from database vector stores (database is source of truth) + vector_store_map = {} + for vector_store in vector_stores_from_db_after_update: + vector_store_id = vector_store.get("vector_store_id") + if vector_store_id: + vector_store_map[vector_store_id] = vector_store + + # Update in-memory registry with database versions (this is the key fix) + instance_2_registry.update_vector_store_in_registry( + vector_store_id=vector_store_id, + updated_data=vector_store + ) + + # Step 5: Verify Instance 2 now has the updated data + instance_2_vs_after_list = instance_2_registry.get_litellm_managed_vector_store_from_registry( + test_vector_store_id + ) + assert instance_2_vs_after_list.get("vector_store_name") == updated_name, ( + "Instance 2 should have updated data after list operation syncs with database" + ) + + # Verify the list returned the correct data + combined_vector_stores = list(vector_store_map.values()) + assert len(combined_vector_stores) == 1 + assert combined_vector_stores[0].get("vector_store_id") == test_vector_store_id + assert combined_vector_stores[0].get("vector_store_name") == updated_name, ( + "List should return updated data from database" + ) + + @pytest.mark.asyncio async def test_resolve_embedding_config_from_db(): """Test that _resolve_embedding_config_from_db correctly resolves embedding config from database.""" From eea24978b990f46dc4a5b38c93675e6720683fb1 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 19 Jan 2026 15:54:04 +0530 Subject: [PATCH 66/77] Add managed files support when load_balancing is True --- .../openai_files_endpoints/files_endpoints.py | 31 +++--- .../test_files_endpoint.py | 94 +++++++++++++++++++ 2 files changed, 111 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 7e3f5820814..2c6b378ae38 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -146,8 +146,8 @@ async def route_create_file( Priority: 1. If target_storage is specified and not "default" -> use storage backend 2. If model parameter provided -> use model credentials and encode ID - 3. If enable_loadbalancing_on_batch_endpoints -> deprecated loadbalancing - 4. If target_model_names_list -> managed files (requires DB) + 3. If target_model_names_list -> managed files (requires DB, supports loadbalancing) + 4. If enable_loadbalancing_on_batch_endpoints -> deprecated loadbalancing 5. Else -> use custom_llm_provider with files_settings """ @@ -202,18 +202,9 @@ async def route_create_file( return response - # EXISTING: Deprecated loadbalancing approach - if ( - litellm.enable_loadbalancing_on_batch_endpoints is True - and is_router_model - and router_model is not None - ): - response = await _deprecated_loadbalanced_create_file( - llm_router=llm_router, - router_model=router_model, - _create_file_request=_create_file_request, - ) - elif target_model_names_list: + # Handle managed files (supports loadbalancing via llm_router.acreate_file) + # Priority: Check for managed files BEFORE deprecated loadbalancing + if target_model_names_list: managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files") if managed_files_obj is None: raise ProxyException( @@ -236,6 +227,7 @@ async def route_create_file( param="None", code=500, ) + # Managed files internally calls llm_router.acreate_file() which includes loadbalancing response = await managed_files_obj.acreate_file( llm_router=llm_router, create_file_request=_create_file_request, @@ -243,6 +235,17 @@ async def route_create_file( litellm_parent_otel_span=user_api_key_dict.parent_otel_span, user_api_key_dict=user_api_key_dict, ) + # EXISTING: Deprecated loadbalancing approach (for backwards compatibility when not using managed files) + elif ( + litellm.enable_loadbalancing_on_batch_endpoints is True + and is_router_model + and router_model is not None + ): + response = await _deprecated_loadbalanced_create_file( + llm_router=llm_router, + router_model=router_model, + _create_file_request=_create_file_request, + ) else: # get configs for custom_llm_provider llm_provider_config = get_files_provider_config( diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 4651bf59b40..b86f927ea00 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -856,3 +856,97 @@ def test_create_file_without_expires_after(mocker: MockerFixture, monkeypatch, l result = response.json() assert result["id"] == "file-abc123" assert result["purpose"] == "fine-tune" + + +def test_managed_files_with_loadbalancing(mocker: MockerFixture, monkeypatch, llm_router: Router): + """ + Test that managed files work with loadbalancing when both target_model_names + and enable_loadbalancing_on_batch_endpoints are enabled. + + This ensures that the priority order is correct: + - managed files should take precedence over deprecated loadbalancing + - managed files internally use llm_router.acreate_file() which provides loadbalancing + """ + from litellm.llms.base_llm.files.transformation import BaseFileEndpoints + from litellm.types.llms.openai import OpenAIFileObject + + # Enable loadbalancing on batch endpoints + monkeypatch.setattr("litellm.enable_loadbalancing_on_batch_endpoints", True) + + proxy_logging_obj = ProxyLogging( + user_api_key_cache=DualCache(default_in_memory_ttl=1) + ) + proxy_logging_obj._add_proxy_hooks(llm_router) + + # Track calls to verify loadbalancing through router + router_acreate_file_calls = [] + + class ManagedFilesWithLoadbalancing(BaseFileEndpoints): + async def acreate_file(self, llm_router, create_file_request, target_model_names_list, litellm_parent_otel_span, user_api_key_dict): + # Verify we receive the target model names + assert len(target_model_names_list) > 0, "Should have target_model_names_list" + + # Simulate what managed files does - call llm_router.acreate_file for each model + # This is where loadbalancing happens internally + for model in target_model_names_list: + router_acreate_file_calls.append({ + "model": model, + "via_router": True + }) + + # Return a managed file ID (base64 encoded) + return OpenAIFileObject( + id="litellm_managed_file_abc123", + object="file", + bytes=100, + created_at=1234567890, + filename="batch_data.jsonl", + purpose="batch", + status="uploaded", + ) + + async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router): + raise NotImplementedError("Not implemented for test") + + async def afile_list(self, purpose, litellm_parent_otel_span): + raise NotImplementedError("Not implemented for test") + + async def afile_delete(self, file_id, litellm_parent_otel_span, llm_router, **data): + raise NotImplementedError("Not implemented for test") + + async def afile_content(self, file_id, litellm_parent_otel_span, llm_router, **data): + raise NotImplementedError("Not implemented for test") + + proxy_logging_obj.proxy_hook_mapping["managed_files"] = ManagedFilesWithLoadbalancing() + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router) + monkeypatch.setattr( + "litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj + ) + + # Create batch file content + test_file_content = b'{"custom_id": "request-1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hello"}]}}' + test_file = ("batch_data.jsonl", test_file_content, "application/jsonl") + + # Make request with both target_model_names AND enable_loadbalancing_on_batch_endpoints + response = client.post( + "/v1/files", + files={"file": test_file}, + data={ + "purpose": "batch", + "target_model_names": "azure-gpt-3-5-turbo,gpt-3.5-turbo", # Multiple models + }, + headers={"Authorization": "Bearer test-key"}, + ) + + # Verify success + assert response.status_code == 200 + result = response.json() + assert result["id"] == "litellm_managed_file_abc123" + assert result["purpose"] == "batch" + + # Verify that managed files was called (via router for loadbalancing) + # This proves that managed files took precedence over deprecated loadbalancing + assert len(router_acreate_file_calls) == 2, "Should have called router for both models" + assert router_acreate_file_calls[0]["model"] == "azure-gpt-3-5-turbo" + assert router_acreate_file_calls[1]["model"] == "gpt-3.5-turbo" + assert all(call["via_router"] for call in router_acreate_file_calls), "All calls should go through router" From d7b103158a2b30132c15052c8b72b06cdb79554d Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 19 Jan 2026 16:40:02 +0530 Subject: [PATCH 67/77] Fix: anthropic-beta is getting overriden and set to anthropic-beta': 'structured-outputs-2025-11-13', --- litellm/llms/anthropic/chat/transformation.py | 27 +++++--- .../test_anthropic_chat_transformation.py | 67 +++++++++++++++++++ 2 files changed, 86 insertions(+), 8 deletions(-) diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 5b1b663e855..86378b97d2e 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -934,8 +934,15 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) return tools - def _ensure_context_management_beta_header(self, headers: dict) -> None: - beta_value = ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value + def _ensure_beta_header(self, headers: dict, beta_value: str) -> None: + """ + Ensure a beta header value is present in the anthropic-beta header. + Merges with existing values instead of overriding them. + + Args: + headers: Dictionary of headers to update + beta_value: The beta header value to add + """ existing_beta = headers.get("anthropic-beta") if existing_beta is None: headers["anthropic-beta"] = beta_value @@ -944,6 +951,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if beta_value not in existing_values: headers["anthropic-beta"] = f"{existing_beta}, {beta_value}" + def _ensure_context_management_beta_header(self, headers: dict) -> None: + beta_value = ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value + self._ensure_beta_header(headers, beta_value) + def update_headers_with_optional_anthropic_beta( self, headers: dict, optional_params: dict ) -> dict: @@ -960,20 +971,20 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if tool.get("type", None) and tool.get("type").startswith( ANTHROPIC_HOSTED_TOOLS.WEB_FETCH.value ): - headers["anthropic-beta"] = ( - ANTHROPIC_BETA_HEADER_VALUES.WEB_FETCH_2025_09_10.value + self._ensure_beta_header( + headers, ANTHROPIC_BETA_HEADER_VALUES.WEB_FETCH_2025_09_10.value ) elif tool.get("type", None) and tool.get("type").startswith( ANTHROPIC_HOSTED_TOOLS.MEMORY.value ): - headers["anthropic-beta"] = ( - ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value + self._ensure_beta_header( + headers, ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value ) if optional_params.get("context_management") is not None: self._ensure_context_management_beta_header(headers) if optional_params.get("output_format") is not None: - headers["anthropic-beta"] = ( - ANTHROPIC_BETA_HEADER_VALUES.STRUCTURED_OUTPUT_2025_09_25.value + self._ensure_beta_header( + headers, ANTHROPIC_BETA_HEADER_VALUES.STRUCTURED_OUTPUT_2025_09_25.value ) return headers diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py index e7a123aa8c3..7e96c4634fd 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py @@ -681,6 +681,73 @@ def test_anthropic_chat_headers_add_context_management_beta(): assert headers["anthropic-beta"] == "context-management-2025-06-27" +def test_anthropic_beta_header_merging_with_output_format(): + """ + Test that anthropic-beta headers from extra_headers are merged with + output_format beta headers instead of being overridden. + + This is a regression test for: https://github.com/BerriAI/litellm/issues/... + When using response_format with a Pydantic model AND extra_headers with + anthropic-beta (e.g., for context-1m extension), both beta headers should + be present in the final request. + """ + config = AnthropicConfig() + + # Simulate headers that already have the context-1m beta header from extra_headers + headers = {"anthropic-beta": "context-1m-2025-08-07"} + + # Simulate output_format being set (happens when using response_format with Sonnet 4.5) + optional_params = { + "output_format": { + "type": "json_schema", + "schema": {"type": "object", "properties": {}} + } + } + + result_headers = config.update_headers_with_optional_anthropic_beta( + headers, optional_params + ) + + # Both beta headers should be present + beta_value = result_headers["anthropic-beta"] + assert "context-1m-2025-08-07" in beta_value, \ + f"User's context-1m beta header missing from: {beta_value}" + assert "structured-outputs-2025-11-13" in beta_value, \ + f"Structured output beta header missing from: {beta_value}" + + +def test_anthropic_beta_header_merging_with_multiple_features(): + """ + Test that multiple beta headers can be merged when using multiple features. + """ + config = AnthropicConfig() + + # Start with a user-provided beta header + headers = {"anthropic-beta": "context-1m-2025-08-07"} + + # Use multiple features that require beta headers + optional_params = { + "output_format": { + "type": "json_schema", + "schema": {"type": "object", "properties": {}} + }, + "context_management": _sample_context_management_payload(), + "tools": [{"type": "web_fetch_20250910", "name": "web_fetch"}] + } + + result_headers = config.update_headers_with_optional_anthropic_beta( + headers, optional_params + ) + + beta_value = result_headers["anthropic-beta"] + + # All beta headers should be present + assert "context-1m-2025-08-07" in beta_value + assert "structured-outputs-2025-11-13" in beta_value + assert "context-management-2025-06-27" in beta_value + assert "web-fetch-2025-09-10" in beta_value + + def test_anthropic_chat_transform_request_includes_context_management(): config = AnthropicConfig() headers = {} From 58daf3eabf3923cda45dadc50834561fbc426a4a Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 19 Jan 2026 17:22:06 +0530 Subject: [PATCH 68/77] Fix Output None for replicate handler --- litellm/llms/replicate/chat/handler.py | 20 ++++++++++++++------ 1 file changed, 14 insertions(+), 6 deletions(-) diff --git a/litellm/llms/replicate/chat/handler.py b/litellm/llms/replicate/chat/handler.py index 4c75db5abc6..c37473b3183 100644 --- a/litellm/llms/replicate/chat/handler.py +++ b/litellm/llms/replicate/chat/handler.py @@ -83,19 +83,27 @@ async def async_handle_prediction_response_streaming( await asyncio.sleep( REPLICATE_POLLING_DELAY_SECONDS ) # prevent being rate limited by replicate - print_verbose(f"replicate: polling endpoint: {prediction_url}") response = await http_client.get(prediction_url, headers=headers) if response.status_code == 200: response_data = response.json() - status = response_data["status"] - if "output" in response_data: + status = response_data.get("status", "") + # Check that "output" exists and is not None or empty + output_present = "output" in response_data and response_data["output"] is not None + if output_present: try: - output_string = "".join(response_data["output"]) + # If output is None or not a list, treat as empty string + if isinstance(response_data["output"], list): + output_string = "".join(response_data["output"]) + elif response_data["output"] is None: + output_string = "" + else: + # fallback for other types; convert to string safely + output_string = str(response_data["output"]) except Exception: raise ReplicateError( status_code=422, message="Unable to parse response. Got={}".format( - response_data["output"] + response_data.get("output", None) ), headers=response.headers, ) @@ -103,7 +111,7 @@ async def async_handle_prediction_response_streaming( print_verbose(f"New chunk: {new_output}") yield {"output": new_output, "status": status} previous_output = output_string - status = response_data["status"] + status = response_data.get("status", "") if status == "failed": replicate_error = response_data.get("error", "") raise ReplicateError( From d5293af053250953394e577e65b772a3a14bd054 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 19 Jan 2026 17:57:00 +0530 Subject: [PATCH 69/77] Fix: update the doc --- README.md | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/README.md b/README.md index 75a23faa5c1..58ffa12c5e1 100644 --- a/README.md +++ b/README.md @@ -374,7 +374,9 @@ Support for more providers. Missing a provider or LLM Platform, raise a [feature 1. (In root) create virtual environment `python -m venv .venv` 2. Activate virtual environment `source .venv/bin/activate` 3. Install dependencies `pip install -e ".[all]"` -4. Start proxy backend `python litellm/proxy_cli.py` +4. `pip install prisma` +5. `prisma generate` +6. Start proxy backend `python litellm/proxy/proxy_cli.py` ### Frontend 1. Navigate to `ui/litellm-dashboard` From 896d1a7dad96d1c21b08b65634196d98da6fdfdd Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 19 Jan 2026 18:18:24 +0530 Subject: [PATCH 70/77] Fix Error: Found packages that need verification: --- ...model_prices_and_context_window_backup.json | 18 ++++++++++++++++++ tests/code_coverage_tests/liccheck.ini | 1 + 2 files changed, 19 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 470d598a25f..43f7bde5da3 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -7857,6 +7857,24 @@ "supports_tool_choice": true, "supports_vision": true }, + "dall-e-2": { + "input_cost_per_image": 0.02, + "litellm_provider": "openai", + "mode": "image_generation", + "supported_endpoints": [ + "/v1/images/generations", + "/v1/images/edits", + "/v1/images/variations" + ] + }, + "dall-e-3": { + "input_cost_per_image": 0.04, + "litellm_provider": "openai", + "mode": "image_generation", + "supported_endpoints": [ + "/v1/images/generations" + ] + }, "deepseek-chat": { "cache_read_input_token_cost": 2.8e-08, "input_cost_per_token": 2.8e-07, diff --git a/tests/code_coverage_tests/liccheck.ini b/tests/code_coverage_tests/liccheck.ini index feb182921db..ea87b56ff30 100644 --- a/tests/code_coverage_tests/liccheck.ini +++ b/tests/code_coverage_tests/liccheck.ini @@ -89,6 +89,7 @@ tokenizers: >=0.20.2 # Apache 2.0 License jinja2: >=3.1.4 # BSD 3-Clause License litellm-proxy-extras: >=0.1.1 # MIT License litellm-enterprise: >=0.1.1 # LiteLLM Enterprise License +a2a-sdk: >=0.3.22 # Apache 2.0 license anyio: >=4.5.0 # Unknown license httpx-aiohttp: >=0.1.4 # Unknown license backoff: >=2.2.1 # Unknown license From 2e6ce1469e9af016072f6e870f7dae002c0b39ab Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 19 Jan 2026 18:47:28 +0530 Subject: [PATCH 71/77] Fix module import for opentelemetry --- litellm/integrations/opentelemetry.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index a223925d59a..c8fc6e96147 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -986,8 +986,7 @@ class OpenTelemetry(CustomLogger): # See: https://github.com/open-telemetry/opentelemetry-python/pull/4676 # TODO: Refactor to use the proper OTEL Logs API instead of directly creating SDK LogRecords - from opentelemetry._logs import SeverityNumber, get_logger, get_logger_provider - from opentelemetry.sdk._logs import LogRecord as SdkLogRecord + from opentelemetry._logs import SeverityNumber, get_logger, get_logger_provider, LogRecord otel_logger = get_logger(LITELLM_LOGGER_NAME) From 6eb3f579d77e2b4b69f8201635d2ab244016452b Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 19 Jan 2026 18:50:18 +0530 Subject: [PATCH 72/77] Fix module import for opentelemetry --- litellm/integrations/opentelemetry.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index c8fc6e96147..9fbccac68dd 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -986,7 +986,11 @@ class OpenTelemetry(CustomLogger): # See: https://github.com/open-telemetry/opentelemetry-python/pull/4676 # TODO: Refactor to use the proper OTEL Logs API instead of directly creating SDK LogRecords - from opentelemetry._logs import SeverityNumber, get_logger, get_logger_provider, LogRecord + from opentelemetry._logs import SeverityNumber, get_logger, get_logger_provider + try: + from opentelemetry.sdk._logs import LogRecord as SdkLogRecord # OTEL < 1.39.0 + except ImportError: + from opentelemetry.sdk._logs._internal import LogRecord as SdkLogRecord # OTEL >= 1.39.0 otel_logger = get_logger(LITELLM_LOGGER_NAME) From 574391c118c331ade273c4fa21c79ff618358404 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 19 Jan 2026 18:51:08 +0530 Subject: [PATCH 73/77] Revert "Fix audio cost per second override (#19158)" This reverts commit 2a0f87bde048e406e6dbf1d5465e77d7e0282bc3. --- docs/my-website/docs/sdk_custom_pricing.md | 5 +--- litellm/llms/openai/cost_calculation.py | 22 ++++++++------ tests/test_litellm/test_cost_calculator.py | 34 ---------------------- 3 files changed, 14 insertions(+), 47 deletions(-) diff --git a/docs/my-website/docs/sdk_custom_pricing.md b/docs/my-website/docs/sdk_custom_pricing.md index ac956db472a..c8577115109 100644 --- a/docs/my-website/docs/sdk_custom_pricing.md +++ b/docs/my-website/docs/sdk_custom_pricing.md @@ -2,10 +2,7 @@ Register custom pricing for sagemaker completion model. -For cost per second pricing, register `input_cost_per_second`. If your provider -charges for audio output duration (e.g., TTS), also set `output_cost_per_second`. -Values of `0` are treated as not billable, so `output_cost_per_second: 0` will -not override `input_cost_per_second`. +For cost per second pricing, you **just** need to register `input_cost_per_second`. ```python # !pip install boto3 diff --git a/litellm/llms/openai/cost_calculation.py b/litellm/llms/openai/cost_calculation.py index 339ce9caedd..e5349db3af7 100644 --- a/litellm/llms/openai/cost_calculation.py +++ b/litellm/llms/openai/cost_calculation.py @@ -105,21 +105,25 @@ def cost_per_second( prompt_cost = 0.0 completion_cost = 0.0 ## Speech / Audio cost calculation - output_cost_per_second = model_info.get("output_cost_per_second") - input_cost_per_second = model_info.get("input_cost_per_second") - - if output_cost_per_second is not None and output_cost_per_second > 0: + if ( + "output_cost_per_second" in model_info + and model_info["output_cost_per_second"] is not None + ): verbose_logger.debug( - f"For model={model} - output_cost_per_second: {output_cost_per_second}; duration: {duration}" + f"For model={model} - output_cost_per_second: {model_info.get('output_cost_per_second')}; duration: {duration}" ) ## COST PER SECOND ## - completion_cost = output_cost_per_second * duration - if input_cost_per_second is not None and input_cost_per_second > 0: + completion_cost = model_info["output_cost_per_second"] * duration + elif ( + "input_cost_per_second" in model_info + and model_info["input_cost_per_second"] is not None + ): verbose_logger.debug( - f"For model={model} - input_cost_per_second: {input_cost_per_second}; duration: {duration}" + f"For model={model} - input_cost_per_second: {model_info.get('input_cost_per_second')}; duration: {duration}" ) ## COST PER SECOND ## - prompt_cost = input_cost_per_second * duration + prompt_cost = model_info["input_cost_per_second"] * duration + completion_cost = 0.0 return prompt_cost, completion_cost diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 42bcfca2311..4d6599fc1b5 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -192,40 +192,6 @@ def test_transcription_cost_falls_back_to_duration(): assert pytest.approx(cost, rel=1e-6) == expected_cost -def test_transcription_cost_prefers_input_when_output_zero(monkeypatch): - from litellm import completion_cost - - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - - model_name = "custom-whisper-input-only" - custom_model_info = { - "input_cost_per_second": 0.00005, - "output_cost_per_second": 0.0, - "litellm_provider": "openai", - "mode": "audio_transcription", - "supported_endpoints": ["/v1/audio/transcriptions"], - } - monkeypatch.setattr( - litellm, - "model_cost", - {**litellm.model_cost, model_name: custom_model_info}, - ) - - response = TranscriptionResponse(text="demo text") - response.duration = 300.0 - - cost = completion_cost( - completion_response=response, - model=model_name, - custom_llm_provider="openai", - call_type="atranscription", - ) - - expected_cost = 300.0 * 0.00005 - assert pytest.approx(cost, rel=1e-6) == expected_cost - - def test_handle_realtime_stream_cost_calculation(): from litellm.cost_calculator import RealtimeAPITokenUsageProcessor From e126e98a7fd1ca21bf3c7581ecdd4977dccfc2e1 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 19 Jan 2026 19:10:43 +0530 Subject: [PATCH 74/77] Fix azure image mypy issues --- litellm/llms/azure_ai/image_edit/flux2_transformation.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/litellm/llms/azure_ai/image_edit/flux2_transformation.py b/litellm/llms/azure_ai/image_edit/flux2_transformation.py index 87bae59ba0f..77d46ff9179 100644 --- a/litellm/llms/azure_ai/image_edit/flux2_transformation.py +++ b/litellm/llms/azure_ai/image_edit/flux2_transformation.py @@ -88,7 +88,7 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig): self, model: str, prompt: Optional[str], - image: FileTypes, + image: Optional[FileTypes], image_edit_optional_request_params: Dict, litellm_params: GenericLiteLLMParams, headers: dict, @@ -102,6 +102,9 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig): if prompt is None: raise ValueError("FLUX 2 image edit requires a prompt.") + if image is None: + raise ValueError("FLUX 2 image edit requires an image.") + image_b64 = self._convert_image_to_base64(image) # Build request body with required params From b49f0a91e41cb607d733be12cdeb8ab5d9c29012 Mon Sep 17 00:00:00 2001 From: Cesar Garcia <128240629+Chesars@users.noreply.github.com> Date: Mon, 19 Jan 2026 10:44:20 -0300 Subject: [PATCH 75/77] fix(responses): resolve deepcopy error with tool_choice ValidatorIterator (#17192) (#17205) Replace copy.deepcopy with model_dump + model_validate in streaming iterator logging to handle Pydantic ValidatorIterator objects that cannot be pickled when tool_choice uses allowed_tools mode. Co-authored-by: Krish Dholakia --- litellm/responses/streaming_iterator.py | 30 ++++++-- ...t_base_responses_api_streaming_iterator.py | 71 +++++++++++++++++-- 2 files changed, 91 insertions(+), 10 deletions(-) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 540d9ad2642..5d17107320d 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -382,11 +382,20 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): def _handle_logging_completed_response(self): """Handle logging for completed responses in async context""" - # Create a deep copy for logging to avoid modifying the response object that will be returned to the user + # Create a copy for logging to avoid modifying the response object that will be returned to the user # The logging handlers may transform usage from Responses API format (input_tokens/output_tokens) # to chat completion format (prompt_tokens/completion_tokens) for internal logging - import copy - logging_response = copy.deepcopy(self.completed_response) + # Use model_dump + model_validate instead of deepcopy to avoid pickle errors with + # Pydantic ValidatorIterator when response contains tool_choice with allowed_tools (fixes #17192) + logging_response = self.completed_response + if self.completed_response is not None and hasattr(self.completed_response, 'model_dump'): + try: + logging_response = type(self.completed_response).model_validate( + self.completed_response.model_dump() + ) + except Exception: + # Fallback to original if serialization fails + pass asyncio.create_task( self.logging_obj.async_success_handler( @@ -468,11 +477,20 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): def _handle_logging_completed_response(self): """Handle logging for completed responses in sync context""" - # Create a deep copy for logging to avoid modifying the response object that will be returned to the user + # Create a copy for logging to avoid modifying the response object that will be returned to the user # The logging handlers may transform usage from Responses API format (input_tokens/output_tokens) # to chat completion format (prompt_tokens/completion_tokens) for internal logging - import copy - logging_response = copy.deepcopy(self.completed_response) + # Use model_dump + model_validate instead of deepcopy to avoid pickle errors with + # Pydantic ValidatorIterator when response contains tool_choice with allowed_tools (fixes #17192) + logging_response = self.completed_response + if self.completed_response is not None and hasattr(self.completed_response, 'model_dump'): + try: + logging_response = type(self.completed_response).model_validate( + self.completed_response.model_dump() + ) + except Exception: + # Fallback to original if serialization fails + pass run_async_function( async_function=self.logging_obj.async_success_handler, diff --git a/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py b/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py index 4a423305626..a0abe9d3a12 100644 --- a/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py +++ b/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py @@ -231,7 +231,7 @@ class TestBaseResponsesAPIStreamingIterator: mock_logging_obj = Mock(spec=LiteLLMLoggingObj) mock_logging_obj.model_call_details = {"litellm_params": {}} mock_config = Mock(spec=BaseResponsesAPIConfig) - + # Create the iterator instance iterator = BaseResponsesAPIStreamingIterator( response=mock_response, @@ -239,15 +239,78 @@ class TestBaseResponsesAPIStreamingIterator: responses_api_provider_config=mock_config, logging_obj=mock_logging_obj ) - + # Test with empty chunk result = iterator._process_chunk("") assert result is None - + # Test with None chunk result = iterator._process_chunk(None) assert result is None + def test_handle_logging_completed_response_with_unpickleable_objects(self): + """ + Test that _handle_logging_completed_response handles responses containing + objects that cannot be pickled (like Pydantic ValidatorIterator). + + This test verifies the fix for issue #17192 where streaming with tool_choice + containing allowed_tools would fail with: + "cannot pickle 'pydantic_core._pydantic_core.ValidatorIterator' object" + + The fix uses model_dump + model_validate instead of copy.deepcopy. + """ + import asyncio + from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator + + # Mock dependencies + mock_response = Mock() + mock_response.headers = {} + mock_response.aiter_lines = Mock() + mock_logging_obj = Mock(spec=LiteLLMLoggingObj) + mock_logging_obj.model_call_details = {"litellm_params": {}} + mock_logging_obj.async_success_handler = Mock() + mock_logging_obj.success_handler = Mock() + mock_config = Mock(spec=BaseResponsesAPIConfig) + + # Create the iterator instance + iterator = ResponsesAPIStreamingIterator( + response=mock_response, + model="gpt-4", + responses_api_provider_config=mock_config, + logging_obj=mock_logging_obj, + litellm_metadata={"model_info": {"id": "model_123"}}, + custom_llm_provider="openai" + ) + + # Create a ResponseCompletedEvent with tool_choice that has model_dump + mock_completed_response = Mock() + mock_completed_response.model_dump.return_value = { + "type": "response.completed", + "response": { + "id": "resp_123", + "output": [{"type": "function_call", "name": "search_web"}], + "tool_choice": {"type": "function", "name": "search_web"} + } + } + # model_validate should return a new mock (the copy) + type(mock_completed_response).model_validate = Mock(return_value=Mock()) + + iterator.completed_response = mock_completed_response + + # This should NOT raise an exception + # Previously it would fail with: TypeError: cannot pickle 'ValidatorIterator' + # Mock asyncio.create_task and executor.submit since we're not in async context + with patch('asyncio.create_task') as mock_create_task, \ + patch('litellm.responses.streaming_iterator.executor') as mock_executor: + try: + iterator._handle_logging_completed_response() + except TypeError as e: + if "pickle" in str(e): + pytest.fail(f"_handle_logging_completed_response failed with pickle error: {e}") + raise + + # Verify model_dump was called (our fix uses this instead of deepcopy) + mock_completed_response.model_dump.assert_called_once() def test_process_chunk_exception_does_not_call_handle_failure(self): """ Test that _process_chunk raises exceptions without calling _handle_failure. @@ -294,4 +357,4 @@ class TestBaseResponsesAPIStreamingIterator: # Verify _handle_failure was NOT called in _process_chunk # It should only be called by the outer exception handler in __next__/__anext__ - mock_handle_failure.assert_not_called() \ No newline at end of file + mock_handle_failure.assert_not_called() From 4d6a430adc87e4629b109564c1d8b5b889209f4a Mon Sep 17 00:00:00 2001 From: Cesar Garcia <128240629+Chesars@users.noreply.github.com> Date: Mon, 19 Jan 2026 11:18:45 -0300 Subject: [PATCH 76/77] docs: update UI contributing guide (#19353) * docs: update UI contributing guide with correct commands - Replace outdated proxy_cli.py command with poetry run litellm - Add config.yaml example with required settings - Clarify that UI comes pre-built in the repo - Add two development options: Build Mode and Dev Mode (hot reload) - Note about redirect issues in Dev Mode * docs: add hot reload login flow and PR submission section - Document the 3000 -> 4000 -> 3000 login flow for hot reload - Reorder: Hot Reload as Option A, Build Mode as Option B - Add section 4 on submitting PRs - Add note that UI changes don't require tests * Update login flow navigation URL in contributing.md --- docs/my-website/docs/contributing.md | 101 +++++++++++++++++++++------ 1 file changed, 78 insertions(+), 23 deletions(-) diff --git a/docs/my-website/docs/contributing.md b/docs/my-website/docs/contributing.md index a88013ff1b3..be7222f6cb8 100644 --- a/docs/my-website/docs/contributing.md +++ b/docs/my-website/docs/contributing.md @@ -1,45 +1,100 @@ # Contributing - UI -Here's how to run the LiteLLM UI locally for making changes: +Thanks for contributing to the LiteLLM UI! This guide will help you set up your local development environment. + + +## 1. Clone the repo -## 1. Clone the repo ```bash git clone https://github.com/BerriAI/litellm.git +cd litellm ``` -## 2. Start the UI + Proxy +## 2. Start the Proxy -**2.1 Start the proxy on port 4000** +Create a config file (e.g., `config.yaml`): -Tell the proxy where the UI is located -```bash -DATABASE_URL = "postgresql://:@:/" -LITELLM_MASTER_KEY = "sk-1234" -STORE_MODEL_IN_DB = "True" +```yaml +model_list: + - model_name: gpt-4o + litellm_params: + model: openai/gpt-4o + +general_settings: + master_key: sk-1234 + database_url: postgresql://:@:/ + store_model_in_db: true ``` +Start the proxy on port 4000: + ```bash -cd litellm/litellm/proxy -python3 proxy_cli.py --config /path/to/config.yaml --port 4000 +poetry run litellm --config config.yaml --port 4000 ``` -**2.2 Start the UI** +The UI comes pre-built in the repo. Access it at `http://localhost:4000/ui` -Set the mode as development (this will assume the proxy is running on localhost:4000) -```bash -npm install # install dependencies -``` +## 3. UI Development + +There are two options for UI development: + +### Option A: Development Mode (Hot Reload) + +This runs the UI on port 3000 with hot reload. The proxy runs on port 4000. ```bash -cd litellm/ui/litellm-dashboard - +cd ui/litellm-dashboard +npm install npm run dev - -# starts on http://0.0.0.0:3000 ``` -## 3. Go to local UI +**Login flow:** +1. Go to `http://localhost:3000` +2. You'll be redirected to `http://localhost:4000/ui` for login +3. After logging in, manually navigate back to `http://localhost:3000/` +4. You're now authenticated and can develop with hot reload + +:::note +If you experience redirect loops or authentication issues, clear your browser cookies for localhost or use Build Mode instead. +::: + +### Option B: Build Mode + +This builds the UI and copies it to the proxy. Changes require rebuilding. + +1. Make your code changes in `ui/litellm-dashboard/src/` + +2. Build the UI +```bash +cd ui/litellm-dashboard +npm install +npm run build +``` + +After building, copy the output to the proxy: ```bash -http://0.0.0.0:3000 -``` \ No newline at end of file +cp -r out/* ../../litellm/proxy/_experimental/out/ +``` + +Then restart the proxy and access the UI at `http://localhost:4000/ui` + +## 4. Submitting a PR + +1. Create a new branch for your changes: +```bash +git checkout -b feat/your-feature-name +``` + +2. Stage and commit your changes: +```bash +git add . +git commit -m "feat: description of your changes" +``` + +3. Push to your fork: +```bash +git push origin feat/your-feature-name +``` + +4. Create a Pull Request on GitHub following the [PR template](https://github.com/BerriAI/litellm/blob/main/.github/pull_request_template.md) From 1be7e877838f1adb923a0c4e0600ca3d1fa5eeff Mon Sep 17 00:00:00 2001 From: superpoussin22 Date: Mon, 19 Jan 2026 15:20:08 +0100 Subject: [PATCH 77/77] Fix HTML entity in survey description text (#19307) --- ui/litellm-dashboard/src/components/survey/ClaudeCodeModal.tsx | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ui/litellm-dashboard/src/components/survey/ClaudeCodeModal.tsx b/ui/litellm-dashboard/src/components/survey/ClaudeCodeModal.tsx index 94527c160c1..eac5e8b7a41 100644 --- a/ui/litellm-dashboard/src/components/survey/ClaudeCodeModal.tsx +++ b/ui/litellm-dashboard/src/components/survey/ClaudeCodeModal.tsx @@ -45,7 +45,7 @@ export function ClaudeCodeModal({ isOpen, onClose, onComplete }: ClaudeCodeModal Help us improve your experience

- We'd love to hear about your experience using LiteLLM with Claude Code. Your feedback helps us improve the product for everyone. + We'd love to hear about your experience using LiteLLM with Claude Code. Your feedback helps us improve the product for everyone.

This brief survey takes about 2-3 minutes to complete.