From b72a03050140e8514ccefdab15ea9b277569ba60 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 14:56:52 -0700 Subject: [PATCH 001/218] test: take keys out of the legacy proxy, enterprise and mcp unit tests before moving them (#42901) * ci: fix the litellm-tests unit job with sysmon coverage, an env allowlist and coverage upload on failure * test: replace key-dependent proxy, enterprise and mcp unit tests with synthetic values and integration and e2e coverage * test: drop key reads at the legacy proxy, enterprise and mcp paths and wire the gemini pass-through split * ci: fail the unit shard when circleci tests split errors * test: drop restating comments from the gemini pass-through split * ci: exit the unit shard cleanly when circleci tests split assigns it no files --------- Co-authored-by: yuneng --- .github/workflows/test-unit-proxy-db.yml | 2 +- .../test_token_counter_gemini_contents_e2e.py | 78 ++++ .../test_prometheus_unit_tests.py | 10 +- .../routing/test_user_config_routing.py | 103 +++++ .../mcp_tests/test_aresponses_api_with_mcp.py | 377 +---------------- .../test_aresponses_api_with_mcp_providers.py | 389 ++++++++++++++++++ .../test_proxy_custom_auth.py | 5 +- .../test_proxy_pass_user_config.py | 114 ----- tests/proxy_unit_tests/test_proxy_server.py | 53 --- .../test_proxy_server_gemini_pass_through.py | 51 +++ .../test_proxy_token_counter.py | 159 +------ tests/proxy_unit_tests/test_proxy_utils.py | 4 +- 12 files changed, 652 insertions(+), 693 deletions(-) create mode 100644 tests/e2e/llm_translation/test_token_counter_gemini_contents_e2e.py create mode 100644 tests/integration/routing/test_user_config_routing.py create mode 100644 tests/mcp_tests/test_aresponses_api_with_mcp_providers.py delete mode 100644 tests/proxy_unit_tests/test_proxy_pass_user_config.py create mode 100644 tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py diff --git a/.github/workflows/test-unit-proxy-db.yml b/.github/workflows/test-unit-proxy-db.yml index 4af7a161984..73015ac6e02 100644 --- a/.github/workflows/test-unit-proxy-db.yml +++ b/.github/workflows/test-unit-proxy-db.yml @@ -106,6 +106,7 @@ jobs: - test-group: proxy-server-core test-path: >- tests/proxy_unit_tests/test_proxy_server.py + tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py tests/proxy_unit_tests/test_aproxy_startup.py workers: 4 dist: loadscope @@ -115,7 +116,6 @@ jobs: tests/proxy_unit_tests/test_proxy_config_unit_test.py tests/proxy_unit_tests/test_proxy_routes.py tests/proxy_unit_tests/test_server_root_path.py - tests/proxy_unit_tests/test_proxy_pass_user_config.py tests/proxy_unit_tests/test_proxy_token_counter.py tests/proxy_unit_tests/test_request_size_limit_middleware.py tests/proxy_unit_tests/test_multipart_bypass_repro.py diff --git a/tests/e2e/llm_translation/test_token_counter_gemini_contents_e2e.py b/tests/e2e/llm_translation/test_token_counter_gemini_contents_e2e.py new file mode 100644 index 00000000000..b2b8f46ab0a --- /dev/null +++ b/tests/e2e/llm_translation/test_token_counter_gemini_contents_e2e.py @@ -0,0 +1,78 @@ +"""Live e2e: `/utils/token_counter?call_endpoint=true` counts Gemini `contents` upstream. + +Google's countTokens API is the only tokenizer that knows Gemini's real token +boundaries, so the proxy must forward `contents` to it for both the AI Studio and +Vertex deployments and hand back the provider's `promptTokensDetails`. Claude on +Vertex is covered by `/v1/messages/count_tokens`; this is the Gemini `contents` +route the claude_code rows never reach +""" + +from __future__ import annotations + +import pytest +from e2e_config import unique_marker +from e2e_http import require_successful_call +from proxy_client import ProxyClient +from pydantic import BaseModel + +pytestmark = pytest.mark.e2e + +GEMINI_DEPLOYMENTS = ("gemini-2.5-flash", "gemini-2.5-flash-vertex") + + +class _Part(BaseModel): + text: str + + +class _Content(BaseModel): + parts: tuple[_Part, ...] + + +class _TokenCountBody(BaseModel): + model: str + contents: tuple[_Content, ...] + + +class _CallEndpoint(BaseModel): + call_endpoint: bool = True + + +class _ModalityTokens(BaseModel): + modality: str + tokenCount: int + + +class _CountTokensUpstream(BaseModel): + totalTokens: int + promptTokensDetails: tuple[_ModalityTokens, ...] + + +class _TokenCountResponse(BaseModel): + total_tokens: int + request_model: str + model_used: str + tokenizer_type: str + original_response: _CountTokensUpstream + + +class TestGeminiContentsTokenCounting: + @pytest.mark.parametrize("model", GEMINI_DEPLOYMENTS) + def test_contents_are_counted_by_the_provider_endpoint( + self, proxy: ProxyClient, scoped_key: str, model: str + ) -> None: + text = f"Hello world, how are you doing today? {unique_marker()}" + body = _TokenCountBody(model=model, contents=(_Content(parts=(_Part(text=text),)),)) + + result = proxy.transport.send( + "/utils/token_counter", + headers=proxy.transport.bearer(scoped_key), + json=body, + params=_CallEndpoint(), + ) + + require_successful_call(result) + counted = _TokenCountResponse.model_validate_json(result.body) + assert counted.request_model == model, counted + assert counted.original_response.totalTokens == counted.total_tokens > 0, counted + assert counted.original_response.promptTokensDetails, counted + assert all(detail.tokenCount > 0 for detail in counted.original_response.promptTokensDetails), counted diff --git a/tests/enterprise/litellm_enterprise/integrations/test_prometheus_unit_tests.py b/tests/enterprise/litellm_enterprise/integrations/test_prometheus_unit_tests.py index 28fd03daf37..5b26dad269e 100644 --- a/tests/enterprise/litellm_enterprise/integrations/test_prometheus_unit_tests.py +++ b/tests/enterprise/litellm_enterprise/integrations/test_prometheus_unit_tests.py @@ -13,7 +13,6 @@ import asyncio from dotenv import load_dotenv load_dotenv() -import os from unittest.mock import MagicMock @@ -165,9 +164,9 @@ async def test_prometheus_metric_tracking(): "model_name": "gpt-5-mini", # openai model name "litellm_params": { # params for litellm completion/embedding call "model": "azure/gpt-4.1-mini", - "api_key": os.getenv("AZURE_AI_API_KEY"), - "api_version": os.getenv("AZURE_AI_API_VERSION"), - "api_base": os.getenv("AZURE_AI_API_BASE"), + "api_key": "sk-azure-unit-test", + "api_version": "2025-01-01-preview", + "api_base": "https://unit-test.openai.azure.com", }, "model_info": {"id": "azure-model-id"}, }, @@ -180,9 +179,6 @@ async def test_prometheus_metric_tracking(): }, ], provider_budget_config=provider_budget_config, - redis_host=os.getenv("REDIS_HOST"), - redis_port=int(os.getenv("REDIS_PORT", 6379)), - redis_password=os.getenv("REDIS_PASSWORD"), ) try: diff --git a/tests/integration/routing/test_user_config_routing.py b/tests/integration/routing/test_user_config_routing.py new file mode 100644 index 00000000000..1c50a78201b --- /dev/null +++ b/tests/integration/routing/test_user_config_routing.py @@ -0,0 +1,103 @@ +import json +import uuid +from pathlib import Path +from typing import Final + +from integration._support.client import Gateway +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server + +USER_KEY: Final = "sk-user-supplied-" + uuid.uuid4().hex + + +def _completion(request: Request) -> Reply: + if request.target != "/v1/chat/completions": + return Reply(status=404, body=b"{}") + body: Final = json.loads(request.body) + return Reply( + body=json.dumps( + { + "id": "chatcmpl-user-config", + "object": "chat.completion", + "created": 0, + "model": body["model"], + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "routed"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 3, "completion_tokens": 1, "total_tokens": 4}, + } + ).encode() + ) + + +def _user_config(upstream_url: str) -> dict[str, object]: + return { + "model_list": [ + { + "model_name": "user-config-deployment", + "litellm_params": { + "model": "openai/gpt-4.1-mini", + "api_base": upstream_url + "/v1", + "api_key": USER_KEY, + }, + } + ], + "num_retries": 0, + } + + +def _opt_in_config(directory: Path, upstream_url: str) -> Path: + config: Final = directory / "allow_client_side_credentials_config.yaml" + config.write_text( + json.dumps( + { + "model_list": [ + { + "model_name": "admin-deployment", + "litellm_params": { + "model": "openai/gpt-4.1-mini", + "api_base": upstream_url + "/v1", + "api_key": "sk-admin-configured", + }, + } + ], + "general_settings": { + "master_key": "os.environ/LITELLM_MASTER_KEY", + "database_url": "os.environ/DATABASE_URL", + "store_model_in_db": True, + "allow_client_side_credentials": True, + }, + } + ) + ) + return config + + +def _request_body(upstream_url: str) -> dict[str, object]: + return { + "model": "user-config-deployment", + "messages": [{"role": "user", "content": "user config control"}], + "user_config": _user_config(upstream_url), + } + + +def test_user_config_routes_to_the_user_supplied_deployment_when_opted_in(gateway: Gateway, tmp_path: Path) -> None: + with wire_server(_completion) as upstream: + config: Final = _opt_in_config(tmp_path, upstream.url) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate: + response: Final = candidate.request("POST", "/v1/chat/completions", _request_body(upstream.url)) + assert response.status_code == 200, response.text + assert response.json()["choices"][0]["message"]["content"] == "routed" + outbound: Final = tuple(upstream.received.get_nowait() for _ in range(upstream.received.qsize())) + completions: Final = tuple(request for request in outbound if request.target == "/v1/chat/completions") + assert len(completions) == 1, outbound + assert completions[0].headers["authorization"] == f"Bearer {USER_KEY}" + assert json.loads(completions[0].body)["model"] == "gpt-4.1-mini" + + +def test_user_config_is_rejected_without_the_opt_in(gateway: Gateway) -> None: + with wire_server(_completion) as upstream: + response: Final = gateway.request("POST", "/v1/chat/completions", _request_body(upstream.url)) + assert response.status_code == 401, response.text + assert "user_config is not allowed in request body" in response.text + assert upstream.received.empty() diff --git a/tests/mcp_tests/test_aresponses_api_with_mcp.py b/tests/mcp_tests/test_aresponses_api_with_mcp.py index eb6f78b57a1..de0dc78af43 100644 --- a/tests/mcp_tests/test_aresponses_api_with_mcp.py +++ b/tests/mcp_tests/test_aresponses_api_with_mcp.py @@ -1,5 +1,3 @@ -import logging -import os import pytest from mcp.types import Tool as MCPTool from typing import List, Any, cast @@ -846,161 +844,6 @@ async def test_streaming_mcp_events_validation(): assert mock_get_tools.called, "MCP tools should have been fetched" -@pytest.mark.asyncio -@pytest.mark.parametrize( - "model", - [ - pytest.param("gpt-4o-mini", id="openai"), - pytest.param("claude-haiku-4-5", id="anthropic"), - ], -) -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" - - Under the hood the follow occurs - - - MCP: responses called litellm MCP manager.list_tools (MOCKED) - - 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 model under test with fetched tools and tool results (REAL LLM CALL) - - Return the user the result of request 2 - """ - # Skip test if API keys are not set for the respective models - if ("claude" in model.lower() or "anthropic" in model.lower()) and not os.getenv( - "ANTHROPIC_API_KEY" - ): - pytest.skip("ANTHROPIC_API_KEY not set, skipping anthropic model test") - if ("gpt" in model.lower() or "openai" in model.lower()) and not os.getenv( - "OPENAI_API_KEY" - ): - pytest.skip("OPENAI_API_KEY not set, skipping openai model test") - - from unittest.mock import AsyncMock, patch - - print("๐Ÿงช Testing basic streaming with MCP tools...") - - # Mock MCP tools that would be returned from the manager - mock_mcp_tools = [ - MCPTool.model_validate({ - "name": "search_repo", - "description": "Search BerriAI/litellm repository for information", - "inputSchema": { - "type": "object", - "properties": { - "query": {"type": "string", "description": "Search query"} - }, - "required": ["query"], - }, - }, by_name=False) - ] - - # Only mock the MCP-specific operations, let LLM responses be real - 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, **kwargs - ): - """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=[ - { - "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 async def test_mcp_parameter_preparation_helpers(): """ @@ -1215,7 +1058,7 @@ async def test_no_duplicate_mcp_tools_in_streaming_e2e(): The test mocks the MCP manager response but validates the actual tools sent to the LLM to ensure no duplication occurs. """ - from unittest.mock import AsyncMock, patch, call + from unittest.mock import AsyncMock, patch from litellm.responses.mcp.litellm_proxy_mcp_handler import ( LiteLLM_Proxy_MCP_Handler, ) @@ -1432,221 +1275,3 @@ async def test_no_duplicate_mcp_tools_in_streaming_e2e(): "tools_per_call": [len(tools) for tools in llm_call_tools], "duplicate_tools_found": False, } - - -@pytest.mark.asyncio -@pytest.mark.parametrize("model", ["gpt-4o-mini"]) -async def test_streaming_mcp_event_order_and_response_id_consistency( - model: str, caplog: pytest.LogCaptureFixture -): - """ - Test that: - 1. Streaming events are emitted in correct order (response.created, response.in_progress, response.output_item.added before MCP events) - 2. All response lifecycle events share the same response ID within a cycle - """ - if ("gpt" in model.lower() or "openai" in model.lower()) and not os.getenv( - "OPENAI_API_KEY" - ): - pytest.skip("OPENAI_API_KEY not set, skipping openai model test") - - from unittest.mock import AsyncMock, patch - - mock_mcp_tools = [ - MCPTool.model_validate({ - "name": "get_weather", - "description": "Get weather for a city", - "inputSchema": { - "type": "object", - "properties": { - "city": {"type": "string", "description": "City name"} - }, - "required": ["city"], - }, - }, by_name=False) - ] - - 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, - ): - mock_get_tools.return_value = (mock_mcp_tools, ["litellm_proxy"]) - - def mock_execute_side_effect(tool_calls, user_api_key_auth, **kwargs): - results = [] - for tool_call in tool_calls: - 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": "Sunny, 72ยฐF", - } - ) - return results - - mock_execute_tools.side_effect = mock_execute_side_effect - - mcp_tool_config = cast( - Any, - { - "type": "mcp", - "server_url": "litellm_proxy", - "require_approval": "never", - }, - ) - - response = await litellm.aresponses( - model=model, - tools=[mcp_tool_config], - input=[ - { - "role": "user", - "type": "message", - "content": "What's the weather in San Francisco?", - } - ], - stream=True, - ) - - events = [] - async for chunk in response: - events.append(chunk) - - assert len(events) > 0, "Should receive streaming events" - - created_idx = next( - ( - i - for i, e in enumerate(events) - if getattr(e, "type", None) == "response.created" - ), - None, - ) - in_progress_idx = next( - ( - i - for i, e in enumerate(events) - if getattr(e, "type", None) == "response.in_progress" - ), - None, - ) - output_item_added_idx = next( - ( - i - for i, e in enumerate(events) - if getattr(e, "type", None) == "response.output_item.added" - ), - None, - ) - mcp_in_progress_idx = next( - ( - i - for i, e in enumerate(events) - if "mcp_list_tools.in_progress" in str(getattr(e, "type", "")) - ), - None, - ) - completed_idx = next( - ( - i - for i, e in enumerate(events) - if getattr(e, "type", None) == "response.completed" - ), - None, - ) - - assert created_idx is not None, "response.created event should be present" - assert ( - in_progress_idx is not None - ), "response.in_progress event should be present" - assert ( - output_item_added_idx is not None - ), "response.output_item.added event should be present" - - assert ( - created_idx < in_progress_idx - ), "response.created should come before response.in_progress" - assert ( - in_progress_idx < output_item_added_idx - ), "response.in_progress should come before response.output_item.added" - - if mcp_in_progress_idx is not None: - assert ( - output_item_added_idx < mcp_in_progress_idx - ), "response.output_item.added should come before response.mcp_list_tools.in_progress" - - response_ids = [] - for i, event in enumerate(events): - event_type = getattr(event, "type", None) - if hasattr(event, "response"): - response_obj = getattr(event, "response", None) - if response_obj and hasattr(response_obj, "id"): - event_type_value = ( - event_type.value - if hasattr(event_type, "value") - else str(event_type) - ) - if any( - x in event_type_value - for x in [ - "response.created", - "response.in_progress", - "response.completed", - ] - ): - response_ids.append((i, event_type_value, response_obj.id)) - - assert ( - len(response_ids) >= 2 - ), f"Should have at least 2 response lifecycle events. Found {len(response_ids)}" - - cycles = [] - current_cycle = [] - current_id = None - - for idx, event_type, resp_id in response_ids: - if current_id is None or resp_id == current_id: - current_cycle.append((idx, event_type, resp_id)) - current_id = resp_id - else: - if current_cycle: - cycles.append(current_cycle) - current_cycle = [(idx, event_type, resp_id)] - current_id = resp_id - if current_cycle: - cycles.append(current_cycle) - - for cycle_num, cycle in enumerate(cycles): - cycle_ids = set(resp_id for _, _, resp_id in cycle) - assert ( - len(cycle_ids) == 1 - ), f"Cycle {cycle_num + 1} should have consistent response ID. Found {len(cycle_ids)} unique IDs" - - assert ( - completed_idx is not None - ), "response.completed event should be present" - - 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 - ) diff --git a/tests/mcp_tests/test_aresponses_api_with_mcp_providers.py b/tests/mcp_tests/test_aresponses_api_with_mcp_providers.py new file mode 100644 index 00000000000..72a0415cf88 --- /dev/null +++ b/tests/mcp_tests/test_aresponses_api_with_mcp_providers.py @@ -0,0 +1,389 @@ +import logging +import os +import pytest +from mcp.types import Tool as MCPTool +from typing import Any, cast + +import litellm +from litellm.responses.mcp.litellm_proxy_mcp_handler import LiteLLM_Proxy_MCP_Handler + + +class MockUserAPIKeyAuth: + """Mock UserAPIKeyAuth for testing""" + + def __init__(self): + self.api_key = "test_key" + self.user_id = "test_user" + self.team_id = "test_team" + self.user_email = "test@example.com" + self.max_budget = 100.0 + self.spend = 0.0 + self.models = [] + self.aliases = {} + self.config = {} + self.permissions = {} + self.metadata = {} + self.object_permission_id = "test_permission_id" + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "model", + [ + pytest.param("gpt-4o-mini", id="openai"), + pytest.param("claude-haiku-4-5", id="anthropic"), + ], +) +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" + + Under the hood the follow occurs + + - MCP: responses called litellm MCP manager.list_tools (MOCKED) + - 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 model under test with fetched tools and tool results (REAL LLM CALL) + + Return the user the result of request 2 + """ + if ("claude" in model.lower() or "anthropic" in model.lower()) and not os.getenv( + "ANTHROPIC_API_KEY" + ): + pytest.skip("ANTHROPIC_API_KEY not set, skipping anthropic model test") + if ("gpt" in model.lower() or "openai" in model.lower()) and not os.getenv( + "OPENAI_API_KEY" + ): + pytest.skip("OPENAI_API_KEY not set, skipping openai model test") + + from unittest.mock import AsyncMock, patch + + print("๐Ÿงช Testing basic streaming with MCP tools...") + + mock_mcp_tools = [ + MCPTool.model_validate({ + "name": "search_repo", + "description": "Search BerriAI/litellm repository for information", + "inputSchema": { + "type": "object", + "properties": { + "query": {"type": "string", "description": "Search query"} + }, + "required": ["query"], + }, + }, by_name=False) + ] + + 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, + ): + mock_get_tools.return_value = (mock_mcp_tools, ["litellm_proxy"]) + + 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: + 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 + + 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=[ + { + "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" + + chunks = [] + async for chunk in response: + chunks.append(chunk) + print(f"๐Ÿ“ฆ Chunk type: {getattr(chunk, 'type', 'unknown')}") + + print(f"๐Ÿ“Š Total chunks received: {len(chunks)}") + + 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)}") + + 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 +@pytest.mark.parametrize("model", ["gpt-4o-mini"]) +async def test_streaming_mcp_event_order_and_response_id_consistency( + model: str, caplog: pytest.LogCaptureFixture +): + """ + Test that: + 1. Streaming events are emitted in correct order (response.created, response.in_progress, response.output_item.added before MCP events) + 2. All response lifecycle events share the same response ID within a cycle + """ + if ("gpt" in model.lower() or "openai" in model.lower()) and not os.getenv( + "OPENAI_API_KEY" + ): + pytest.skip("OPENAI_API_KEY not set, skipping openai model test") + + from unittest.mock import AsyncMock, patch + + mock_mcp_tools = [ + MCPTool.model_validate({ + "name": "get_weather", + "description": "Get weather for a city", + "inputSchema": { + "type": "object", + "properties": { + "city": {"type": "string", "description": "City name"} + }, + "required": ["city"], + }, + }, by_name=False) + ] + + 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, + ): + mock_get_tools.return_value = (mock_mcp_tools, ["litellm_proxy"]) + + def mock_execute_side_effect(tool_calls, user_api_key_auth, **kwargs): + results = [] + for tool_call in tool_calls: + 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": "Sunny, 72ยฐF", + } + ) + return results + + mock_execute_tools.side_effect = mock_execute_side_effect + + mcp_tool_config = cast( + Any, + { + "type": "mcp", + "server_url": "litellm_proxy", + "require_approval": "never", + }, + ) + + response = await litellm.aresponses( + model=model, + tools=[mcp_tool_config], + input=[ + { + "role": "user", + "type": "message", + "content": "What's the weather in San Francisco?", + } + ], + stream=True, + ) + + events = [] + async for chunk in response: + events.append(chunk) + + assert len(events) > 0, "Should receive streaming events" + + created_idx = next( + ( + i + for i, e in enumerate(events) + if getattr(e, "type", None) == "response.created" + ), + None, + ) + in_progress_idx = next( + ( + i + for i, e in enumerate(events) + if getattr(e, "type", None) == "response.in_progress" + ), + None, + ) + output_item_added_idx = next( + ( + i + for i, e in enumerate(events) + if getattr(e, "type", None) == "response.output_item.added" + ), + None, + ) + mcp_in_progress_idx = next( + ( + i + for i, e in enumerate(events) + if "mcp_list_tools.in_progress" in str(getattr(e, "type", "")) + ), + None, + ) + completed_idx = next( + ( + i + for i, e in enumerate(events) + if getattr(e, "type", None) == "response.completed" + ), + None, + ) + + assert created_idx is not None, "response.created event should be present" + assert ( + in_progress_idx is not None + ), "response.in_progress event should be present" + assert ( + output_item_added_idx is not None + ), "response.output_item.added event should be present" + + assert ( + created_idx < in_progress_idx + ), "response.created should come before response.in_progress" + assert ( + in_progress_idx < output_item_added_idx + ), "response.in_progress should come before response.output_item.added" + + if mcp_in_progress_idx is not None: + assert ( + output_item_added_idx < mcp_in_progress_idx + ), "response.output_item.added should come before response.mcp_list_tools.in_progress" + + response_ids = [] + for i, event in enumerate(events): + event_type = getattr(event, "type", None) + if hasattr(event, "response"): + response_obj = getattr(event, "response", None) + if response_obj and hasattr(response_obj, "id"): + event_type_value = ( + event_type.value + if hasattr(event_type, "value") + else str(event_type) + ) + if any( + x in event_type_value + for x in [ + "response.created", + "response.in_progress", + "response.completed", + ] + ): + response_ids.append((i, event_type_value, response_obj.id)) + + assert ( + len(response_ids) >= 2 + ), f"Should have at least 2 response lifecycle events. Found {len(response_ids)}" + + cycles = [] + current_cycle = [] + current_id = None + + for idx, event_type, resp_id in response_ids: + if current_id is None or resp_id == current_id: + current_cycle.append((idx, event_type, resp_id)) + current_id = resp_id + else: + if current_cycle: + cycles.append(current_cycle) + current_cycle = [(idx, event_type, resp_id)] + current_id = resp_id + if current_cycle: + cycles.append(current_cycle) + + for cycle_num, cycle in enumerate(cycles): + cycle_ids = set(resp_id for _, _, resp_id in cycle) + assert ( + len(cycle_ids) == 1 + ), f"Cycle {cycle_num + 1} should have consistent response ID. Found {len(cycle_ids)} unique IDs" + + assert ( + completed_idx is not None + ), "response.completed event should be present" + + 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 + ) diff --git a/tests/proxy_unit_tests/test_proxy_custom_auth.py b/tests/proxy_unit_tests/test_proxy_custom_auth.py index b575e4c85c6..dbbad0dab1e 100644 --- a/tests/proxy_unit_tests/test_proxy_custom_auth.py +++ b/tests/proxy_unit_tests/test_proxy_custom_auth.py @@ -53,8 +53,7 @@ def test_custom_auth(client): "max_tokens": 10, } # Your bearer token - token = os.getenv("PROXY_MASTER_KEY") - print(f"token: {token}") + token = "sk-unit-test-master" headers = {"Authorization": f"Bearer {token}"} with pytest.raises(Exception, match="Authentication Error, Failed custom auth") as exc_info: client.post("/chat/completions", json=test_data, headers=headers) @@ -71,7 +70,7 @@ def test_custom_auth_bearer(client): "max_tokens": 10, } # Your bearer token - token = os.getenv("PROXY_MASTER_KEY") + token = "sk-unit-test-master" headers = {"Authorization": f"WITHOUT BEAR Er {token}"} with pytest.raises(Exception, match="CustomAuth - Malformed API Key passed in") as exc_info: diff --git a/tests/proxy_unit_tests/test_proxy_pass_user_config.py b/tests/proxy_unit_tests/test_proxy_pass_user_config.py deleted file mode 100644 index 91911c142ea..00000000000 --- a/tests/proxy_unit_tests/test_proxy_pass_user_config.py +++ /dev/null @@ -1,114 +0,0 @@ -import sys, os -import traceback -from dotenv import load_dotenv - -load_dotenv() -import io - -# this file is to test litellm/proxy - -import pytest, logging, asyncio -import litellm -from litellm import embedding, completion, completion_cost, Timeout -from litellm import RateLimitError - -# Configure logging -logging.basicConfig( - level=logging.DEBUG, # Set the desired logging level - format="%(asctime)s - %(levelname)s - %(message)s", -) - -# test /chat/completion request to the proxy -from fastapi.testclient import TestClient -from fastapi import FastAPI -from litellm.proxy.proxy_server import ( - router, - save_worker_config, - initialize, -) # Replace with the actual module where your FastAPI router is defined - -# Your bearer token -token = "sk-1234" - -headers = {"Authorization": f"Bearer {token}"} - - -@pytest.fixture(scope="function") -def client_no_auth(): - # Assuming litellm.proxy.proxy_server is an object - from litellm.proxy.proxy_server import cleanup_router_config_variables - - cleanup_router_config_variables() - filepath = os.path.dirname(os.path.abspath(__file__)) - config_fp = f"{filepath}/test_configs/test_config_no_auth.yaml" - # initialize can get run in parallel, it sets specific variables for the fast api app, sinc eit gets run in parallel different tests use the wrong variables - asyncio.run(initialize(config=config_fp, debug=True)) - app = FastAPI() - app.include_router(router) # Include your router in the test app - - return TestClient(app) - - -@pytest.mark.skipif( - os.environ.get("AZURE_AI_API_KEY") is None - or os.environ.get("OPENAI_API_KEY") is None, - reason="AZURE_AI_API_KEY or OPENAI_API_KEY not set - skipping integration test", -) -def test_chat_completion(client_no_auth): - global headers - - from litellm.types.router import RouterConfig, ModelConfig - from litellm.types.completion import CompletionRequest - - user_config = RouterConfig( - model_list=[ - ModelConfig( - model_name="user-azure-instance", - litellm_params=CompletionRequest( - model="azure/gpt-4.1-mini", - api_key=os.getenv("AZURE_AI_API_KEY"), - api_version=os.getenv("AZURE_API_VERSION"), - api_base=os.getenv("AZURE_AI_API_BASE"), - timeout=10, - ), - tpm=240000, - rpm=1800, - ), - ModelConfig( - model_name="user-openai-instance", - litellm_params=CompletionRequest( - model="gpt-3.5-turbo", - api_key=os.getenv("OPENAI_API_KEY"), - timeout=10, - ), - tpm=240000, - rpm=1800, - ), - ], - num_retries=2, - allowed_fails=3, - fallbacks=[{"user-azure-instance": ["user-openai-instance"]}], - ).dict() - - try: - # Your test data - test_data = { - "model": "user-azure-instance", - "messages": [ - {"role": "user", "content": "hi"}, - ], - "max_tokens": 10, - "user_config": user_config, - } - - print("testing proxy server with chat completions") - response = client_no_auth.post("/v1/chat/completions", json=test_data) - print(f"response - {response.text}") - assert response.status_code == 200 - result = response.json() - print(f"Received response: {result}") - except Exception as e: - pytest.fail(f"LiteLLM Proxy test failed. Exception - {str(e)}") - - -# Run the test diff --git a/tests/proxy_unit_tests/test_proxy_server.py b/tests/proxy_unit_tests/test_proxy_server.py index ed0380058a5..5be27b3ad72 100644 --- a/tests/proxy_unit_tests/test_proxy_server.py +++ b/tests/proxy_unit_tests/test_proxy_server.py @@ -2065,59 +2065,6 @@ async def test_add_callback_via_key_litellm_pre_call_utils_langsmith( assert new_data["failure_callback"] == expected_failure_callbacks -@pytest.mark.skipif( - not os.getenv("GEMINI_API_KEY") and not os.getenv("GOOGLE_API_KEY"), - reason="Requires GEMINI_API_KEY or GOOGLE_API_KEY.", -) -@pytest.mark.asyncio -async def test_gemini_pass_through_endpoint(): - from starlette.datastructures import URL - - from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( - Request, - Response, - gemini_proxy_route, - ) - - body = b""" - { - "contents": [{ - "parts":[{ - "text": "The quick brown fox jumps over the lazy dog." - }] - }] - } - """ - - # Construct the scope dictionary - scope = { - "type": "http", - "method": "POST", - "path": "/gemini/v1beta/models/gemini-2.5-flash:countTokens", - "query_string": b"key=sk-1234", - "headers": [ - (b"content-type", b"application/json"), - ], - } - - # Create a new Request object - async def async_receive(): - return {"type": "http.request", "body": body, "more_body": False} - - request = Request( - scope=scope, - receive=async_receive, - ) - - resp = await gemini_proxy_route( - endpoint="v1beta/models/gemini-2.5-flash:countTokens?key=sk-1234", - request=request, - fastapi_response=Response(), - ) - - print(resp.body) - - @pytest.mark.parametrize("hidden", [True, False]) @pytest.mark.asyncio async def test_model_info_alias_without_prisma(hidden): diff --git a/tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py b/tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py new file mode 100644 index 00000000000..2453ec3bfe3 --- /dev/null +++ b/tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py @@ -0,0 +1,51 @@ +import os + +import pytest + + +@pytest.mark.skipif( + not os.getenv("GEMINI_API_KEY") and not os.getenv("GOOGLE_API_KEY"), + reason="Requires GEMINI_API_KEY or GOOGLE_API_KEY.", +) +@pytest.mark.asyncio +async def test_gemini_pass_through_endpoint(): + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + Request, + Response, + gemini_proxy_route, + ) + + body = b""" + { + "contents": [{ + "parts":[{ + "text": "The quick brown fox jumps over the lazy dog." + }] + }] + } + """ + + scope = { + "type": "http", + "method": "POST", + "path": "/gemini/v1beta/models/gemini-2.5-flash:countTokens", + "query_string": b"key=sk-1234", + "headers": [ + (b"content-type", b"application/json"), + ], + } + + async def async_receive(): + return {"type": "http.request", "body": body, "more_body": False} + + request = Request( + scope=scope, + receive=async_receive, + ) + + await gemini_proxy_route( + endpoint="v1beta/models/gemini-2.5-flash:countTokens?key=sk-1234", + request=request, + fastapi_response=Response(), + ) + diff --git a/tests/proxy_unit_tests/test_proxy_token_counter.py b/tests/proxy_unit_tests/test_proxy_token_counter.py index 39ec4bb1887..8590e959961 100644 --- a/tests/proxy_unit_tests/test_proxy_token_counter.py +++ b/tests/proxy_unit_tests/test_proxy_token_counter.py @@ -2,10 +2,7 @@ # 1. Generate a Key, and use it to make a call -import json import logging -import os -import tempfile from unittest.mock import AsyncMock, MagicMock, patch import httpx @@ -35,79 +32,6 @@ from litellm.types.utils import TokenCountResponse verbose_proxy_logger.setLevel(level=logging.DEBUG) -def get_vertex_ai_creds_json() -> dict: - # Define the path to the vertex_key.json file - print("loading vertex ai credentials") - filepath = os.path.dirname(os.path.abspath(__file__)) - vertex_key_path = filepath + "/vertex_key.json" - # Read the existing content of the file or create an empty dictionary - try: - with open(vertex_key_path, "r") as file: - # Read the file content - print("Read vertexai file path") - content = file.read() - - # If the file is empty or not valid JSON, create an empty dictionary - if not content or not content.strip(): - service_account_key_data = {} - else: - # Attempt to load the existing JSON content - file.seek(0) - service_account_key_data = json.load(file) - except FileNotFoundError: - # If the file doesn't exist, create an empty dictionary - service_account_key_data = {} - - # Update the service_account_key_data with environment variables - private_key_id = os.environ.get("VERTEX_AI_PRIVATE_KEY_ID", "") - private_key = os.environ.get("VERTEX_AI_PRIVATE_KEY", "") - private_key = private_key.replace("\\n", "\n") - service_account_key_data["private_key_id"] = private_key_id - service_account_key_data["private_key"] = private_key - - return service_account_key_data - - -def load_vertex_ai_credentials(): - # Define the path to the vertex_key.json file - print("loading vertex ai credentials") - filepath = os.path.dirname(os.path.abspath(__file__)) - vertex_key_path = filepath + "/vertex_key.json" - - # Read the existing content of the file or create an empty dictionary - try: - with open(vertex_key_path, "r") as file: - # Read the file content - print("Read vertexai file path") - content = file.read() - - # If the file is empty or not valid JSON, create an empty dictionary - if not content or not content.strip(): - service_account_key_data = {} - else: - # Attempt to load the existing JSON content - file.seek(0) - service_account_key_data = json.load(file) - except FileNotFoundError: - # If the file doesn't exist, create an empty dictionary - service_account_key_data = {} - - # Update the service_account_key_data with environment variables - private_key_id = os.environ.get("VERTEX_AI_PRIVATE_KEY_ID", "") - private_key = os.environ.get("VERTEX_AI_PRIVATE_KEY", "") - private_key = private_key.replace("\\n", "\n") - service_account_key_data["private_key_id"] = private_key_id - service_account_key_data["private_key"] = private_key - - # Create a temporary file - with tempfile.NamedTemporaryFile(mode="w+", delete=False) as temp_file: - # Write the updated content to the temporary files - json.dump(service_account_key_data, temp_file, indent=2) - - # Export the temporary file as GOOGLE_APPLICATION_CREDENTIALS - os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = os.path.abspath(temp_file.name) - - @pytest.mark.asyncio async def test_vLLM_token_counting(): """ @@ -223,10 +147,12 @@ async def test_anthropic_messages_count_tokens_endpoint(): - Should return response in Anthropic format: {"input_tokens": } - Should work as wrapper around internal token_counter function """ - from litellm.proxy.anthropic_endpoints.endpoints import count_tokens - from fastapi import Request from unittest.mock import MagicMock + from fastapi import Request + + from litellm.proxy.anthropic_endpoints.endpoints import count_tokens + # Mock request object mock_request = MagicMock(spec=Request) mock_request_data = { @@ -295,10 +221,12 @@ async def test_anthropic_messages_count_tokens_with_non_anthropic_model(): - Should still work and return Anthropic format - Should call internal token_counter with from_anthropic_endpoint=True """ - from litellm.proxy.anthropic_endpoints.endpoints import count_tokens - from fastapi import Request from unittest.mock import MagicMock + from fastapi import Request + + from litellm.proxy.anthropic_endpoints.endpoints import count_tokens + # Mock request object mock_request = MagicMock(spec=Request) mock_request_data = { @@ -435,10 +363,12 @@ async def test_anthropic_endpoint_error_handling(): """ Test error handling in the /v1/messages/count_tokens endpoint """ - from litellm.proxy.anthropic_endpoints.endpoints import count_tokens - from fastapi import Request, HTTPException from unittest.mock import MagicMock + from fastapi import HTTPException, Request + + from litellm.proxy.anthropic_endpoints.endpoints import count_tokens + # Mock request object mock_request = MagicMock(spec=Request) mock_user_api_key_dict = MagicMock() @@ -474,8 +404,10 @@ async def test_anthropic_endpoint_error_handling(): @pytest.mark.asyncio async def test_factory_anthropic_endpoint_calls_anthropic_counter(): """Test that /v1/messages/count_tokens with Anthropic model uses Anthropic counter.""" - from unittest.mock import patch, AsyncMock, MagicMock + from unittest.mock import AsyncMock, MagicMock, patch + from fastapi.testclient import TestClient + from litellm.proxy.proxy_server import app # Mock the global handler instance in token_counter module @@ -531,8 +463,10 @@ async def test_factory_anthropic_endpoint_calls_anthropic_counter(): @pytest.mark.asyncio async def test_factory_gpt4_endpoint_does_not_call_anthropic_counter(): """Test that /v1/messages/count_tokens with GPT-4 does NOT use Anthropic counter.""" - from unittest.mock import patch, AsyncMock, MagicMock + from unittest.mock import AsyncMock, MagicMock, patch + from fastapi.testclient import TestClient + from litellm.proxy.proxy_server import app # Mock the global handler instance in token_counter module @@ -590,8 +524,10 @@ async def test_factory_gpt4_endpoint_does_not_call_anthropic_counter(): @pytest.mark.asyncio async def test_factory_normal_token_counter_endpoint_does_not_call_anthropic(): """Test that /utils/token_counter does NOT use Anthropic counter even with Anthropic model.""" - from unittest.mock import patch, AsyncMock, MagicMock + from unittest.mock import AsyncMock, MagicMock, patch + from fastapi.testclient import TestClient + from litellm.proxy.proxy_server import app # Mock the global handler instance in token_counter module @@ -678,57 +614,6 @@ async def test_factory_registration(): assert not counter.should_use_token_counting_api(custom_llm_provider=None) -@pytest.mark.skip( - reason="Requires Google/Vertex AI credentials (GEMINI_API_KEY or VERTEX_AI_PRIVATE_KEY)." -) -@pytest.mark.asyncio -@pytest.mark.parametrize("model_name", ["gemini-2.5-pro", "vertex-ai-gemini-2.5-pro"]) -async def test_vertex_ai_gemini_token_counting_with_contents(model_name): - """ - Test token counting for Vertex AI Gemini model using contents format with call_endpoint=True - """ - load_vertex_ai_credentials() - llm_router = Router( - model_list=[ - { - "model_name": "gemini-2.5-pro", - "litellm_params": { - "model": "gemini/gemini-2.5-pro", - }, - }, - { - "model_name": "vertex-ai-gemini-2.5-pro", - "litellm_params": { - "model": "vertex_ai/gemini-2.5-pro", - }, - }, - ] - ) - - setattr(litellm.proxy.proxy_server, "llm_router", llm_router) - - # Test with contents format and call_endpoint=True - response = await token_counter( - request=TokenCountRequest( - model=model_name, - contents=[ - {"parts": [{"text": "Hello world, how are you doing today? i am ij"}]} - ], - ), - call_endpoint=True, - ) - - print("Vertex AI Gemini token counting response:", response) - - # validate we have original response - assert response.original_response is not None - assert response.original_response.get("totalTokens") is not None - assert response.original_response.get("promptTokensDetails") is not None - - prompt_tokens_details = response.original_response.get("promptTokensDetails") - assert prompt_tokens_details is not None - - @pytest.mark.asyncio async def test_bedrock_count_tokens_endpoint(): """ @@ -779,7 +664,7 @@ async def test_vertex_ai_anthropic_token_counting(): This tests the token counting implementation for Vertex AI partner models without making actual API calls. Mocks at the handler level to test the full flow. """ - from unittest.mock import AsyncMock, patch, MagicMock + from unittest.mock import patch # Mock the Vertex AI partner models token counter response mock_token_response = { diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index 1134f41a940..ab787bbbe27 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -438,7 +438,7 @@ def test_is_request_body_safe_global_enabled( "model_name": "gpt-3.5-turbo", "litellm_params": { "model": "gpt-3.5-turbo", - "api_key": os.getenv("OPENAI_API_KEY"), + "api_key": "sk-openai-unit-test", }, } ] @@ -475,7 +475,7 @@ def test_is_request_body_safe_model_enabled( "model_name": "fireworks_ai/*", "litellm_params": { "model": "fireworks_ai/*", - "api_key": os.getenv("FIREWORKS_API_KEY"), + "api_key": "sk-fireworks-unit-test", "configurable_clientside_auth_params": ( ["api_base"] if allow_client_side_credentials else [] ), From e3f087315de3eac8ec5c78b31fca218b1f846892 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 17:06:30 -0500 Subject: [PATCH 002/218] feat(terraform): add display_name to litellm_model resource and model data sources (#42987) * feat(terraform): add display_name to litellm_model resource and model data sources Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(terraform): persist display_name on update and read /model/info data envelope Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore(terraform): drop PATCH /model/{model_id}/update from endpoint audit allowlist Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(terraform): surface external display_name removal as drift on refresh Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * ci(terraform): rerun after uv download timeout Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- terraform/provider/CHANGELOG.md | 2 + terraform/provider/docs/data-sources/model.md | 1 + .../provider/docs/data-sources/models.md | 1 + terraform/provider/docs/resources/model.md | 2 + .../provider/litellm/data_source_model.go | 24 +- .../litellm/data_source_model_test.go | 12 +- terraform/provider/litellm/resource_model.go | 5 + .../provider/litellm/resource_model_crud.go | 35 ++- .../provider/litellm/resource_model_test.go | 244 ++++++++++++++++++ terraform/provider/litellm/types.go | 23 +- terraform/provider/litellm/utils.go | 7 + .../endpointaudit/coverage_allowlist.txt | 1 - 12 files changed, 333 insertions(+), 24 deletions(-) create mode 100644 terraform/provider/litellm/resource_model_test.go diff --git a/terraform/provider/CHANGELOG.md b/terraform/provider/CHANGELOG.md index 079bb7d8667..4123b7e6b51 100644 --- a/terraform/provider/CHANGELOG.md +++ b/terraform/provider/CHANGELOG.md @@ -16,6 +16,7 @@ longer signal it. ### Added +- **model**: Optional `display_name` argument on `litellm_model`, sent as `model_info.display_name` and returned as `display_name` by `/v1/models`, so client model pickers show a readable name; changes are persisted through `/model/{id}/update` since `/model/update` ignores `model_info`; also exported by the `litellm_model` and `litellm_models` data sources - **key**: Computed `server_metadata` attribute on `litellm_key` exposing every metadata entry the proxy stores, so metadata created outside Terraform is visible in state and drift on it shows on refresh, while `metadata` keeps tracking only the declared entries and updates keep preserving undeclared ones - **team_member_add**: `tpm_limit`, `rpm_limit`, `budget_duration`, and `allowed_models` attributes on `litellm_team_member_add`, applied to every member of the resource; `budget_duration` and `allowed_models` ride on `/team/member_add`, while the limits are sent through `/team/member_update`, which is where the proxy accepts them - **team**: Optional `team_id` argument on `litellm_team`, so teams can be created with a stable, human-readable ID instead of a provider-generated UUID; changing it forces replacement @@ -48,6 +49,7 @@ longer signal it. ### Fixed +- **model**: `litellm_model` refresh now reads the `{"data": [...]}` envelope `/model/info` returns, so `model_info` fields changed outside Terraform show up as drift instead of silently keeping the previous state - **key**: An update that changes `team_id` and fails because the key was already cascade-deleted along with its previous team now recovers by recreating the key under the new team, instead of aborting the apply. The key's absence is confirmed against the proxy first, so an unrelated failure still errors out, and a `team_id` change between two teams that both still exist stays a plain in-place update - **credential**: create now reports a `credential_name` collision as a clear error naming the `terraform import` command that adopts the existing credential, instead of surfacing the proxy's raw 500 with a Prisma `Unique constraint failed` message. New `adopt_existing` argument (default `false`) opts into taking the existing credential over during create, which makes `apply` idempotent again once state loses track of a credential that still exists on the proxy. Requires a proxy that answers 409 on the collision; older proxies are still detected by their 500 message - **credential**: credential names and `model_id` are now percent-encoded in request URLs, so a name containing `/`, `?`, `#` or spaces reaches the proxy intact instead of being cut at the first reserved character and read, updated or deleted as a different credential diff --git a/terraform/provider/docs/data-sources/model.md b/terraform/provider/docs/data-sources/model.md index 6976ff1523a..587652bcbd8 100644 --- a/terraform/provider/docs/data-sources/model.md +++ b/terraform/provider/docs/data-sources/model.md @@ -43,6 +43,7 @@ In addition to all arguments above, the following attributes are exported: * `tier` - Model tier (`free` or `paid`). * `mode` - Model mode, e.g. `chat` or `embedding`. * `team_id` - Team the deployment is scoped to, if any. +* `display_name` - Human-readable name returned by `/v1/models`, if configured. * `db_model` - Whether the deployment is stored in the database (as opposed to config). ## Security Note diff --git a/terraform/provider/docs/data-sources/models.md b/terraform/provider/docs/data-sources/models.md index 7862dc30ab7..1cf0fd36ab3 100644 --- a/terraform/provider/docs/data-sources/models.md +++ b/terraform/provider/docs/data-sources/models.md @@ -41,4 +41,5 @@ In addition to all arguments above, the following attributes are exported: * `tier` - Model tier (`free` or `paid`). * `mode` - Model mode, e.g. `chat` or `embedding`. * `team_id` - Team the deployment is scoped to, if any. + * `display_name` - Human-readable name returned by `/v1/models`, if configured. * `db_model` - Whether the deployment is stored in the database. diff --git a/terraform/provider/docs/resources/model.md b/terraform/provider/docs/resources/model.md index 0409b48b391..5bb68bfe918 100644 --- a/terraform/provider/docs/resources/model.md +++ b/terraform/provider/docs/resources/model.md @@ -126,6 +126,8 @@ The following arguments are supported: * `team_id` - (Optional) string. Associate the model with a specific team. +* `display_name` - (Optional) string. Human-readable name stored in `model_info.display_name` and returned as `display_name` by `/v1/models`, so clients such as Claude Code and Claude Desktop show it in their model picker instead of `model_name`. When unset, clients fall back to `model_name`. + * `mode` - (Optional) string. The intended use of the model. Valid values are: * `completion` * `embedding` diff --git a/terraform/provider/litellm/data_source_model.go b/terraform/provider/litellm/data_source_model.go index 78af04ac160..6b9993d267e 100644 --- a/terraform/provider/litellm/data_source_model.go +++ b/terraform/provider/litellm/data_source_model.go @@ -23,14 +23,15 @@ type modelInfoParams struct { } type modelInfoMeta struct { - ID string `json:"id"` - DBModel bool `json:"db_model"` - BaseModel string `json:"base_model"` - Tier string `json:"tier"` - Mode string `json:"mode"` - TeamID string `json:"team_id"` - CreatedAt string `json:"created_at"` - UpdatedAt string `json:"updated_at"` + ID string `json:"id"` + DBModel bool `json:"db_model"` + BaseModel string `json:"base_model"` + Tier string `json:"tier"` + Mode string `json:"mode"` + TeamID string `json:"team_id"` + DisplayName string `json:"display_name"` + CreatedAt string `json:"created_at"` + UpdatedAt string `json:"updated_at"` } type modelInfoEntry struct { @@ -112,6 +113,10 @@ func dataSourceLiteLLMModel() *schema.Resource { Type: schema.TypeString, Computed: true, }, + "display_name": { + Type: schema.TypeString, + Computed: true, + }, "db_model": { Type: schema.TypeBool, Computed: true, @@ -161,6 +166,7 @@ func dataSourceLiteLLMModelRead(d *schema.ResourceData, m interface{}) error { d.Set("tier", entry.ModelInfo.Tier) d.Set("mode", entry.ModelInfo.Mode) d.Set("team_id", entry.ModelInfo.TeamID) + d.Set("display_name", entry.ModelInfo.DisplayName) d.Set("db_model", entry.ModelInfo.DBModel) log.Printf("[INFO] Successfully read model with ID: %s", modelID) @@ -197,6 +203,7 @@ func dataSourceLiteLLMModels() *schema.Resource { "tier": {Type: schema.TypeString, Computed: true}, "mode": {Type: schema.TypeString, Computed: true}, "team_id": {Type: schema.TypeString, Computed: true}, + "display_name": {Type: schema.TypeString, Computed: true}, "db_model": {Type: schema.TypeBool, Computed: true}, }, }, @@ -247,6 +254,7 @@ func dataSourceLiteLLMModelsRead(d *schema.ResourceData, m interface{}) error { "tier": entry.ModelInfo.Tier, "mode": entry.ModelInfo.Mode, "team_id": entry.ModelInfo.TeamID, + "display_name": entry.ModelInfo.DisplayName, "db_model": entry.ModelInfo.DBModel, }) } diff --git a/terraform/provider/litellm/data_source_model_test.go b/terraform/provider/litellm/data_source_model_test.go index 97d7f07dcd8..b46a355853a 100644 --- a/terraform/provider/litellm/data_source_model_test.go +++ b/terraform/provider/litellm/data_source_model_test.go @@ -35,7 +35,8 @@ func TestDataSourceModelReadSingleObject(t *testing.T) { "base_model": "gpt-4o", "tier": "paid", "mode": "chat", - "team_id": "team-1" + "team_id": "team-1", + "display_name": "GPT-4o" } } }`)) @@ -66,6 +67,7 @@ func TestDataSourceModelReadSingleObject(t *testing.T) { "tier": "paid", "mode": "chat", "team_id": "team-1", + "display_name": "GPT-4o", "db_model": true, } for attr, want := range checks { @@ -115,7 +117,7 @@ func TestDataSourceModelsRead(t *testing.T) { w.Header().Set("Content-Type", "application/json") w.Write([]byte(`{ "data": [ - {"model_name": "a", "litellm_params": {"model": "openai/a", "custom_llm_provider": "openai"}, "model_info": {"id": "id-1", "db_model": true}}, + {"model_name": "a", "litellm_params": {"model": "openai/a", "custom_llm_provider": "openai"}, "model_info": {"id": "id-1", "db_model": true, "display_name": "Model A"}}, {"model_name": "b", "litellm_params": {"model": "anthropic/b", "custom_llm_provider": "anthropic"}, "model_info": {"id": "id-2"}} ] }`)) @@ -143,7 +145,11 @@ func TestDataSourceModelsRead(t *testing.T) { t.Fatalf("expected 2 models, got %d", len(models)) } first := models[0].(map[string]interface{}) - if first["model_name"] != "a" || first["custom_llm_provider"] != "openai" || first["db_model"] != true { + if first["model_name"] != "a" || first["custom_llm_provider"] != "openai" || first["db_model"] != true || first["display_name"] != "Model A" { t.Errorf("unexpected first model: %v", first) } + second := models[1].(map[string]interface{}) + if second["display_name"] != "" { + t.Errorf("expected empty display_name for model without one, got %v", second["display_name"]) + } } diff --git a/terraform/provider/litellm/resource_model.go b/terraform/provider/litellm/resource_model.go index b0a7304718b..85cb1d038bf 100644 --- a/terraform/provider/litellm/resource_model.go +++ b/terraform/provider/litellm/resource_model.go @@ -93,6 +93,11 @@ func resourceLiteLLMModel() *schema.Resource { Type: schema.TypeString, Optional: true, }, + "display_name": { + Type: schema.TypeString, + Optional: true, + Description: "Human-readable name returned as display_name by /v1/models, shown in client model pickers instead of model_name", + }, "mode": { Type: schema.TypeString, Optional: true, diff --git a/terraform/provider/litellm/resource_model_crud.go b/terraform/provider/litellm/resource_model_crud.go index fc5d5b09dd5..c7db2a673e5 100644 --- a/terraform/provider/litellm/resource_model_crud.go +++ b/terraform/provider/litellm/resource_model_crud.go @@ -4,6 +4,7 @@ import ( "encoding/json" "fmt" "log" + "net/url" "strconv" "strings" "time" @@ -53,6 +54,7 @@ func retryModelRead(d *schema.ResourceData, m interface{}, maxRetries int) error const ( endpointModelNew = "/model/new" endpointModelUpdate = "/model/update" + endpointModelPatch = "/model/%s/update" endpointModelInfo = "/model/info" endpointModelDelete = "/model/delete" ) @@ -246,12 +248,13 @@ func createOrUpdateModel(d *schema.ResourceData, m interface{}, isUpdate bool) e ModelName: d.Get("model_name").(string), LiteLLMParams: litellmParams, ModelInfo: ModelInfo{ - ID: modelID, - DBModel: true, - BaseModel: pricingBaseModel, - Tier: d.Get("tier").(string), - Mode: d.Get("mode").(string), - TeamID: d.Get("team_id").(string), + ID: modelID, + DBModel: true, + BaseModel: pricingBaseModel, + Tier: d.Get("tier").(string), + Mode: d.Get("mode").(string), + TeamID: d.Get("team_id").(string), + DisplayName: d.Get("display_name").(string), }, Additional: make(map[string]interface{}), } @@ -275,6 +278,12 @@ func createOrUpdateModel(d *schema.ResourceData, m interface{}, isUpdate bool) e return fmt.Errorf("failed to %s model: %w", map[bool]string{true: "update", false: "create"}[isUpdate], err) } + if isUpdate && d.HasChange("display_name") { + if err := patchModelDisplayName(client, modelID, d.Get("display_name").(string)); err != nil { + return fmt.Errorf("failed to update model display_name: %w", err) + } + } + d.SetId(modelID) log.Printf("[INFO] Model created with ID %s. Starting retry mechanism to read the model...", modelID) @@ -282,6 +291,19 @@ func createOrUpdateModel(d *schema.ResourceData, m interface{}, isUpdate bool) e return retryModelRead(d, m, 5) } +// /model/update only merges litellm_params, so model_info changes go through the PATCH endpoint. +func patchModelDisplayName(client *Client, modelID, displayName string) error { + resp, err := MakeRequest(client, "PATCH", fmt.Sprintf(endpointModelPatch, url.PathEscape(modelID)), ModelInfoPatch{ + ModelInfo: ModelInfoPatchFields{ID: modelID, DisplayName: displayName}, + }) + if err != nil { + return err + } + defer resp.Body.Close() + _, err = handleAPIResponse(resp, nil, client) + return err +} + func resourceLiteLLMModelCreate(d *schema.ResourceData, m interface{}) error { return createOrUpdateModel(d, m, false) } @@ -327,6 +349,7 @@ func resourceLiteLLMModelRead(d *schema.ResourceData, m interface{}) error { d.Set("tier", GetStringValue(modelResp.ModelInfo.Tier, d.Get("tier").(string))) d.Set("mode", GetStringValue(modelResp.ModelInfo.Mode, d.Get("mode").(string))) d.Set("team_id", GetStringValue(modelResp.ModelInfo.TeamID, d.Get("team_id").(string))) + d.Set("display_name", modelResp.ModelInfo.DisplayName) // Preserve credential name from state since it might not be returned by API d.Set("litellm_credential_name", d.Get("litellm_credential_name").(string)) diff --git a/terraform/provider/litellm/resource_model_test.go b/terraform/provider/litellm/resource_model_test.go new file mode 100644 index 00000000000..0be3c39a3c5 --- /dev/null +++ b/terraform/provider/litellm/resource_model_test.go @@ -0,0 +1,244 @@ +package litellm + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/hashicorp/terraform-plugin-sdk/v2/helper/schema" + "github.com/hashicorp/terraform-plugin-sdk/v2/terraform" +) + +func modelInfoBody(displayName string) string { + modelInfo := map[string]interface{}{ + "id": "model-123", + "db_model": true, + "base_model": "claude-sonnet-4-5", + "tier": "free", + "mode": "chat", + } + if displayName != "" { + modelInfo["display_name"] = displayName + } + body, _ := json.Marshal(map[string]interface{}{ + "model_name": "sonnet-4-5-anthropic", + "litellm_params": map[string]interface{}{"model": "anthropic/claude-sonnet-4-5", "custom_llm_provider": "anthropic"}, + "model_info": modelInfo, + }) + return string(body) +} + +func modelInfoDataEnvelope(displayName string) string { + return `{"data": [` + modelInfoBody(displayName) + `]}` +} + +func TestResourceLiteLLMModelCreateSendsDisplayName(t *testing.T) { + var createPayload map[string]interface{} + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/model/new": + if err := json.NewDecoder(r.Body).Decode(&createPayload); err != nil { + t.Errorf("failed to decode create payload: %v", err) + } + w.Write([]byte(modelInfoBody("Claude Sonnet 4.5"))) + case "/model/info": + w.Write([]byte(modelInfoBody("Claude Sonnet 4.5"))) + default: + t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path) + w.WriteHeader(http.StatusNotFound) + } + })) + defer srv.Close() + + d := schema.TestResourceDataRaw(t, resourceLiteLLMModel().Schema, map[string]interface{}{ + "model_name": "sonnet-4-5-anthropic", + "custom_llm_provider": "anthropic", + "base_model": "claude-sonnet-4-5", + "model_api_key": "sk-ant-test", + "mode": "chat", + "display_name": "Claude Sonnet 4.5", + }) + + if err := resourceLiteLLMModelCreate(d, NewClient(srv.URL, "test-key", true)); err != nil { + t.Fatalf("create failed: %v", err) + } + + modelInfo, ok := createPayload["model_info"].(map[string]interface{}) + if !ok { + t.Fatalf("expected model_info object in create payload, got %v", createPayload["model_info"]) + } + if modelInfo["display_name"] != "Claude Sonnet 4.5" { + t.Errorf("expected model_info.display_name 'Claude Sonnet 4.5', got %v", modelInfo["display_name"]) + } + if got := d.Get("display_name").(string); got != "Claude Sonnet 4.5" { + t.Errorf("expected state display_name 'Claude Sonnet 4.5', got %q", got) + } +} + +func TestResourceLiteLLMModelCreateOmitsUnsetDisplayName(t *testing.T) { + var createPayload map[string]interface{} + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/model/new": + if err := json.NewDecoder(r.Body).Decode(&createPayload); err != nil { + t.Errorf("failed to decode create payload: %v", err) + } + w.Write([]byte(modelInfoBody(""))) + case "/model/info": + w.Write([]byte(modelInfoBody(""))) + default: + t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path) + w.WriteHeader(http.StatusNotFound) + } + })) + defer srv.Close() + + d := schema.TestResourceDataRaw(t, resourceLiteLLMModel().Schema, map[string]interface{}{ + "model_name": "sonnet-4-5-anthropic", + "custom_llm_provider": "anthropic", + "base_model": "claude-sonnet-4-5", + "model_api_key": "sk-ant-test", + }) + + if err := resourceLiteLLMModelCreate(d, NewClient(srv.URL, "test-key", true)); err != nil { + t.Fatalf("create failed: %v", err) + } + + modelInfo := createPayload["model_info"].(map[string]interface{}) + if _, present := modelInfo["display_name"]; present { + t.Errorf("expected display_name to be omitted from model_info when unset, got %v", modelInfo["display_name"]) + } + if got := d.Get("display_name").(string); got != "" { + t.Errorf("expected empty state display_name, got %q", got) + } +} + +func TestResourceLiteLLMModelReadDisplayName(t *testing.T) { + cases := map[string]struct { + serverBody string + want string + }{ + "server value wins inside data envelope": {serverBody: modelInfoDataEnvelope("Renamed In Admin UI"), want: "Renamed In Admin UI"}, + "server value wins unwrapped": {serverBody: modelInfoBody("Renamed In Admin UI"), want: "Renamed In Admin UI"}, + "external removal clears state": {serverBody: modelInfoDataEnvelope(""), want: ""}, + } + for name, tc := range cases { + t.Run(name, func(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/model/info" { + t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path) + } + w.Write([]byte(tc.serverBody)) + })) + defer srv.Close() + + d := schema.TestResourceDataRaw(t, resourceLiteLLMModel().Schema, map[string]interface{}{ + "model_name": "sonnet-4-5-anthropic", + "custom_llm_provider": "anthropic", + "base_model": "claude-sonnet-4-5", + "display_name": "Claude Sonnet 4.5", + }) + d.SetId("model-123") + + if err := resourceLiteLLMModelRead(d, NewClient(srv.URL, "test-key", true)); err != nil { + t.Fatalf("read failed: %v", err) + } + if got := d.Get("display_name").(string); got != tc.want { + t.Errorf("expected display_name %q, got %q", tc.want, got) + } + }) + } +} + +func updateResourceData(t *testing.T, oldDisplayName, newDisplayName string) *schema.ResourceData { + t.Helper() + res := resourceLiteLLMModel() + attrs := map[string]string{ + "model_name": "sonnet-4-5-anthropic", + "custom_llm_provider": "anthropic", + "base_model": "claude-sonnet-4-5", + } + if oldDisplayName != "" { + attrs["display_name"] = oldDisplayName + } + state := &terraform.InstanceState{ID: "model-123", Attributes: attrs} + diff, err := res.Diff(context.Background(), state, &terraform.ResourceConfig{Config: map[string]interface{}{ + "model_name": "sonnet-4-5-anthropic", + "custom_llm_provider": "anthropic", + "base_model": "claude-sonnet-4-5", + "display_name": newDisplayName, + }}, nil) + if err != nil { + t.Fatalf("diff failed: %v", err) + } + d, err := schema.InternalMap(res.Schema).Data(state, diff) + if err != nil { + t.Fatalf("data failed: %v", err) + } + return d +} + +func TestResourceLiteLLMModelUpdatePatchesDisplayName(t *testing.T) { + cases := map[string]struct { + newName string + }{ + "changed name is patched": {newName: "Claude Sonnet 4.5 v2"}, + "cleared name is patched": {newName: ""}, + } + for name, tc := range cases { + t.Run(name, func(t *testing.T) { + var patchPayload map[string]interface{} + var patchPath string + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch { + case r.Method == http.MethodPost && r.URL.Path == "/model/update": + w.Write([]byte(modelInfoBody("Claude Sonnet 4.5"))) + case r.Method == http.MethodPatch: + patchPath = r.URL.Path + if err := json.NewDecoder(r.Body).Decode(&patchPayload); err != nil { + t.Errorf("failed to decode patch payload: %v", err) + } + w.Write([]byte(modelInfoBody(tc.newName))) + case r.URL.Path == "/model/info": + w.Write([]byte(modelInfoDataEnvelope(tc.newName))) + default: + t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path) + w.WriteHeader(http.StatusNotFound) + } + })) + defer srv.Close() + + d := updateResourceData(t, "Claude Sonnet 4.5", tc.newName) + if err := resourceLiteLLMModelUpdate(d, NewClient(srv.URL, "test-key", true)); err != nil { + t.Fatalf("update failed: %v", err) + } + if patchPath != "/model/model-123/update" { + t.Fatalf("expected PATCH /model/model-123/update, got %q", patchPath) + } + modelInfo := patchPayload["model_info"].(map[string]interface{}) + if modelInfo["display_name"] != tc.newName { + t.Errorf("expected patched display_name %q, got %v", tc.newName, modelInfo["display_name"]) + } + if got := d.Get("display_name").(string); got != tc.newName { + t.Errorf("expected state display_name %q, got %q", tc.newName, got) + } + }) + } +} + +func TestResourceLiteLLMModelUpdateSkipsPatchWhenDisplayNameUnchanged(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodPatch { + t.Errorf("unexpected PATCH %s", r.URL.Path) + } + w.Write([]byte(modelInfoDataEnvelope("Claude Sonnet 4.5"))) + })) + defer srv.Close() + + d := updateResourceData(t, "Claude Sonnet 4.5", "Claude Sonnet 4.5") + if err := resourceLiteLLMModelUpdate(d, NewClient(srv.URL, "test-key", true)); err != nil { + t.Fatalf("update failed: %v", err) + } +} diff --git a/terraform/provider/litellm/types.go b/terraform/provider/litellm/types.go index 8bcf7dc4fe3..a8784b8a6a9 100644 --- a/terraform/provider/litellm/types.go +++ b/terraform/provider/litellm/types.go @@ -25,6 +25,16 @@ type ModelResponse struct { Additional map[string]interface{} `json:"additional"` } +// ModelInfoPatch is the body for PATCH /model/{id}/update; display_name is sent even when empty so it can be cleared. +type ModelInfoPatch struct { + ModelInfo ModelInfoPatchFields `json:"model_info"` +} + +type ModelInfoPatchFields struct { + ID string `json:"id"` + DisplayName string `json:"display_name"` +} + // ModelRequest represents a request to create or update a model. type ModelRequest struct { ModelName string `json:"model_name"` @@ -108,12 +118,13 @@ type LiteLLMParams struct { // ModelInfo represents information about a model. type ModelInfo struct { - ID string `json:"id"` - DBModel bool `json:"db_model"` - BaseModel string `json:"base_model"` - Tier string `json:"tier"` - Mode string `json:"mode"` - TeamID string `json:"team_id,omitempty"` + ID string `json:"id"` + DBModel bool `json:"db_model"` + BaseModel string `json:"base_model"` + Tier string `json:"tier"` + Mode string `json:"mode"` + TeamID string `json:"team_id,omitempty"` + DisplayName string `json:"display_name,omitempty"` } // Key represents a LiteLLM API key. diff --git a/terraform/provider/litellm/utils.go b/terraform/provider/litellm/utils.go index f8f66afba3c..ce1dae55f59 100644 --- a/terraform/provider/litellm/utils.go +++ b/terraform/provider/litellm/utils.go @@ -55,6 +55,13 @@ func handleAPIResponse(resp *http.Response, reqBody interface{}, client *Client) resp.Status, client.redactSensitiveData(string(bodyBytes)), client.redactSensitiveData(string(reqBodyBytes))) } + var envelope struct { + Data []json.RawMessage `json:"data"` + } + if err := json.Unmarshal(bodyBytes, &envelope); err == nil && len(envelope.Data) > 0 { + bodyBytes = envelope.Data[0] + } + var modelResp ModelResponse if err := json.Unmarshal(bodyBytes, &modelResp); err != nil { return nil, fmt.Errorf("failed to parse response: %v", err) diff --git a/terraform/provider/tools/endpointaudit/coverage_allowlist.txt b/terraform/provider/tools/endpointaudit/coverage_allowlist.txt index 7aacf7ceab9..e4574031d86 100644 --- a/terraform/provider/tools/endpointaudit/coverage_allowlist.txt +++ b/terraform/provider/tools/endpointaudit/coverage_allowlist.txt @@ -97,7 +97,6 @@ GET /guardrails/{guardrail_id} GET /prompts/{prompt_id} GET /prompts/{prompt_id}/versions PATCH /guardrails/{guardrail_id} -PATCH /model/{model_id}/update PATCH /prompts/{prompt_id} PATCH /team/{team_id} POST /team/model/add From de8aeff6c6d2f360b4ade529ecb90159c896bf1e Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 17:07:10 -0500 Subject: [PATCH 003/218] feat(proxy_cli): add --validate_config dry-run flag (#41705) * feat(proxy_cli): add --validate_config dry-run flag Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(tests): format --validate_config CliRunner calls Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy_cli): run --validate_config before the ollama auto-start Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy_cli): restore file and add ollama validate_config regression test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: yassin --- litellm/proxy/proxy_cli.py | 29 ++++++ tests/test_litellm/proxy/test_proxy_cli.py | 103 +++++++++++++++++++++ 2 files changed, 132 insertions(+) diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index ac69f2e3894..27c03d2d5d7 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -210,6 +210,25 @@ class ProxyInitializationHelpers: response: Final = httpx.get(url=f"http://{host}:{port}/health") print(json.dumps(response.json(), indent=4)) + @staticmethod + def _run_config_validation(config: str | None) -> None: + if config is None: + raise click.UsageError("--validate_config requires --config ") + import asyncio + + from litellm.proxy.proxy_server import ProxyConfig + + async def _load() -> int: + _, model_list, _ = await ProxyConfig().load_config(router=None, config_file_path=config) + return len(model_list) + + try: + model_count: Final = asyncio.run(_load()) + except Exception as error: + click.echo(f"LiteLLM: config validation failed: {error}", err=True) + raise click.exceptions.Exit(1) from error + click.echo(f"LiteLLM: config OK ({model_count} models)") + @staticmethod def _run_test_chat_completion( host: str, @@ -887,6 +906,12 @@ class ProxyInitializationHelpers: default=False, help="Skip starting the server after setup (useful for migrations only)", ) +@click.option( + "--validate_config", + is_flag=True, + default=False, + help="Load and validate the config file (including mcp_servers) without starting the server, then exit. Exit code 1 on any config error.", +) @click.option( "--keepalive_timeout", default=None, @@ -1027,6 +1052,7 @@ def run_server( log_config, use_prisma_db_push: bool, skip_server_startup, + validate_config: bool, keepalive_timeout, timeout_worker_healthcheck, max_requests_before_restart, @@ -1069,6 +1095,9 @@ def run_server( if version is True: ProxyInitializationHelpers._echo_litellm_version() return + if validate_config is True: + ProxyInitializationHelpers._run_config_validation(config) + return if model and "ollama" in model and api_base is None: ProxyInitializationHelpers._run_ollama_serve() if health is True: diff --git a/tests/test_litellm/proxy/test_proxy_cli.py b/tests/test_litellm/proxy/test_proxy_cli.py index b84efda3308..a275dd62400 100644 --- a/tests/test_litellm/proxy/test_proxy_cli.py +++ b/tests/test_litellm/proxy/test_proxy_cli.py @@ -3224,3 +3224,106 @@ class TestLibpqSslParamTranslation: assert query["sslmode"] == ["require"] assert query["sslcert"] == ["/certs/rds-bundle.pem"] assert query["sslaccept"] == ["strict"] + + +@pytest.mark.xdist_group("proxy_cli") +class TestValidateConfigFlag: + def test_validate_config_valid_config_exits_zero(self, tmp_path, monkeypatch): + from click.testing import CliRunner + + from litellm.proxy.proxy_cli import run_server + + monkeypatch.delenv("DATABASE_URL", raising=False) + monkeypatch.delenv("DIRECT_URL", raising=False) + config_path = tmp_path / "config.yaml" + config_path.write_text( + yaml.safe_dump( + { + "model_list": [ + { + "model_name": "gpt-4o", + "litellm_params": { + "model": "openai/gpt-4o", + "api_key": "sk-fake", + }, + } + ] + } + ) + ) + + result = CliRunner().invoke(run_server, ["--config", str(config_path), "--validate_config"]) + + assert result.exit_code == 0, f"exit_code={result.exit_code}, output={result.output}" + assert "config OK" in result.output + + def test_validate_config_invalid_mcp_server_exits_one(self, tmp_path, monkeypatch): + from click.testing import CliRunner + + from litellm.proxy.proxy_cli import run_server + + monkeypatch.delenv("DATABASE_URL", raising=False) + monkeypatch.delenv("DIRECT_URL", raising=False) + config_path = tmp_path / "config.yaml" + config_path.write_text( + yaml.safe_dump( + { + "mcp_servers": { + "zapier": { + "url": "https://example.com/mcp", + "transport": "http", + "per_server_oauth_discovery": "yes", + } + } + } + ) + ) + + result = CliRunner().invoke(run_server, ["--config", str(config_path), "--validate_config"]) + + assert result.exit_code == 1, f"exit_code={result.exit_code}, output={result.output}" + assert "per_server_oauth_discovery must be a boolean" in result.output + + def test_validate_config_without_config_is_usage_error(self, monkeypatch): + from click.testing import CliRunner + + from litellm.proxy.proxy_cli import run_server + + result = CliRunner().invoke(run_server, ["--validate_config"]) + + assert result.exit_code != 0 + assert "--validate_config requires --config" in result.output + + @patch("subprocess.Popen") + def test_validate_config_with_ollama_model_does_not_start_ollama(self, mock_popen, tmp_path, monkeypatch): + from click.testing import CliRunner + + from litellm.proxy.proxy_cli import run_server + + monkeypatch.delenv("DATABASE_URL", raising=False) + monkeypatch.delenv("DIRECT_URL", raising=False) + config_path = tmp_path / "config.yaml" + config_path.write_text( + yaml.safe_dump( + { + "model_list": [ + { + "model_name": "gpt-4o", + "litellm_params": { + "model": "openai/gpt-4o", + "api_key": "sk-fake", + }, + } + ] + } + ) + ) + + result = CliRunner().invoke( + run_server, + ["--config", str(config_path), "--model", "ollama/llama3", "--validate_config"], + ) + + assert result.exit_code == 0, f"exit_code={result.exit_code}, output={result.output}" + assert "config OK" in result.output + mock_popen.assert_not_called() From 7faeb15ff3dfa5e227e32fb69b3a2b6945b3ad72 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 15:17:31 -0700 Subject: [PATCH 004/218] fix(e2e): skip unpublished npm versions in the Claude Code PR-gate resolver (#43053) npm keeps an unpublished version's timestamp in the packument's time map but drops it from versions, so the resolver could hand npm install a version it refuses with ETARGET. Only versions still present in versions are candidates now. Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- .../_pr_gate_unit_tests/__init__.py | 0 .../test_pr_gate_version_resolver.py | 60 +++++++++++++++++++ .../claude_code/pr_gate_version_resolver.py | 7 ++- 3 files changed, 66 insertions(+), 1 deletion(-) create mode 100644 tests/e2e/claude_code/_pr_gate_unit_tests/__init__.py create mode 100644 tests/e2e/claude_code/_pr_gate_unit_tests/test_pr_gate_version_resolver.py diff --git a/tests/e2e/claude_code/_pr_gate_unit_tests/__init__.py b/tests/e2e/claude_code/_pr_gate_unit_tests/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/e2e/claude_code/_pr_gate_unit_tests/test_pr_gate_version_resolver.py b/tests/e2e/claude_code/_pr_gate_unit_tests/test_pr_gate_version_resolver.py new file mode 100644 index 00000000000..0ed7e2bf083 --- /dev/null +++ b/tests/e2e/claude_code/_pr_gate_unit_tests/test_pr_gate_version_resolver.py @@ -0,0 +1,60 @@ +"""Unit tests for the Claude Code PR-gate version resolver. + +Markerless harness tests: they feed the resolver a hand-built packument and a +fixed clock, so they run without a proxy, never reach the npm registry, and +carry no `e2e` marker. +""" + +from __future__ import annotations + +from datetime import datetime, timezone +from typing import Final, Mapping + +import pytest + +from claude_code.pr_gate_version_resolver import NoEligibleVersionError, resolve_pr_gate_version + +NOW: Final = datetime(2026, 4, 25, 12, 0, tzinfo=timezone.utc) +INSIDE_THE_2_1_88_WINDOW: Final = datetime(2026, 4, 3, 12, 0, tzinfo=timezone.utc) + + +def _packument(times: Mapping[str, str], unpublished: frozenset[str] = frozenset()) -> dict[str, object]: + return { + "name": "@anthropic-ai/claude-code", + "time": {"created": "2024-01-01T00:00:00.000Z", "modified": "2026-04-25T00:00:00.000Z", **times}, + "versions": {version: {"version": version} for version in times if version not in unpublished}, + } + + +def test_skips_a_version_npm_has_unpublished() -> None: + metadata: Final = _packument( + { + "2.1.87": "2026-03-28T20:00:00.000Z", + "2.1.88": "2026-03-30T22:36:48.424Z", + "2.1.89": "2026-03-31T23:32:40.000Z", + }, + unpublished=frozenset({"2.1.88"}), + ) + assert resolve_pr_gate_version(metadata=metadata, as_of=INSIDE_THE_2_1_88_WINDOW) == "2.1.87" + + +def test_raises_when_the_only_old_enough_version_is_unpublished() -> None: + metadata: Final = _packument( + {"2.1.88": "2026-03-30T22:36:48.424Z", "2.1.89": "2026-03-31T23:32:40.000Z"}, + unpublished=frozenset({"2.1.88"}), + ) + with pytest.raises(NoEligibleVersionError): + resolve_pr_gate_version(metadata=metadata, as_of=INSIDE_THE_2_1_88_WINDOW) + + +def test_picks_the_newest_published_version_at_least_min_age_old() -> None: + metadata: Final = _packument( + { + "2.1.118": "2026-04-15T10:00:00.000Z", + "2.1.119": "2026-04-21T10:00:00.000Z", + "2.2.0-rc.1": "2026-04-22T10:00:00.000Z", + "2.1.120": "2026-04-23T10:00:00.000Z", + "2.1.121": "2026-04-25T11:00:00.000Z", + } + ) + assert resolve_pr_gate_version(metadata=metadata, as_of=NOW) == "2.1.119" diff --git a/tests/e2e/claude_code/pr_gate_version_resolver.py b/tests/e2e/claude_code/pr_gate_version_resolver.py index 82e12a2bf15..756dafb7d73 100644 --- a/tests/e2e/claude_code/pr_gate_version_resolver.py +++ b/tests/e2e/claude_code/pr_gate_version_resolver.py @@ -80,7 +80,9 @@ def resolve_pr_gate_version( "Newest" means newest by **publish time**, not semver string order โ€” if a patch lands on an older major after a newer release, the - patched line is the eligible one. + patched line is the eligible one. A version npm has unpublished keeps + its ``time`` entry but drops out of ``versions``, so only versions + still present in ``versions`` are candidates. Args: metadata: Pre-fetched npm packument (skips the HTTP call). Useful @@ -101,6 +103,7 @@ def resolve_pr_gate_version( metadata = fetch(package_name) times = metadata.get("time") or {} + versions = metadata.get("versions") or {} if as_of is None: as_of = datetime.now(timezone.utc) cutoff = as_of - min_age @@ -109,6 +112,8 @@ def resolve_pr_gate_version( for version, raw_ts in times.items(): if version in _TIME_META_KEYS: continue + if version not in versions: + continue if not isinstance(raw_ts, str): continue if "-" in version: From 040b37fa49b99be78b997f76ad46bd6f10c07f27 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 15:18:55 -0700 Subject: [PATCH 005/218] chore(cost-map): move azure gpt-realtime-2.1-mini deprecation date to the later Models API date (#43058) Co-authored-by: kerry Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/model_prices_and_context_window_backup.json | 2 +- model_prices_and_context_window.json | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 680b6cceab6..9bec83f6b08 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -6092,7 +6092,7 @@ "cache_creation_input_audio_token_cost": 3e-07, "cache_read_input_audio_token_cost": 3e-07, "cache_read_input_token_cost": 6e-08, - "deprecation_date": "2027-06-25", + "deprecation_date": "2027-07-31", "input_cost_per_audio_token": 1e-05, "input_cost_per_image_token": 8e-07, "input_cost_per_token": 6e-07, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 680b6cceab6..9bec83f6b08 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -6092,7 +6092,7 @@ "cache_creation_input_audio_token_cost": 3e-07, "cache_read_input_audio_token_cost": 3e-07, "cache_read_input_token_cost": 6e-08, - "deprecation_date": "2027-06-25", + "deprecation_date": "2027-07-31", "input_cost_per_audio_token": 1e-05, "input_cost_per_image_token": 8e-07, "input_cost_per_token": 6e-07, From 4aa3ff47fe1fd525c54edd9adc4e206cf439baea Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 15:23:05 -0700 Subject: [PATCH 006/218] docs(github): require UI before/after screenshots and intentional UX change note in PR template (#43021) * docs(github): require UI before/after screenshots and intentional UX change note in PR template Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * docs(github): move intentional change note into TLDR rules and dedupe screenshots Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * docs(github): refresh user flow screenshots with new commits Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: kerry Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .github/pull_request_template.md | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index 4b3878bed11..1fe0c602036 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -6,7 +6,8 @@ ## TLDR - + Problem this solves: @@ -28,7 +29,8 @@ How it solves it: No LiteLLM internals: never name functions, files, DB tables, config classes, hooks, callbacks, or code paths. "The upload hands back an ID that looks like OpenAI's own `file-abc123` instead of the scrambled one the gateway returned" is right, "no managed-file row was registered" is wrong Keep the two lists step-for-step identical until they diverge, so the changed step is obvious If the bug had a security or authorization consequence, end each list with what another user could or could no longer do - Regenerate this section whenever new commits change the PR's behavior, so it never describes an older revision + Regenerate this section, screenshots included, whenever new commits change the PR's behavior, so it never describes an older revision + If the PR changes what an Admin UI page shows, embed a before and an after screenshot of that page right after its list, taken at the same URL on the same data, with the rows, fields, or controls that changed boxed in red so a reader spots the difference without reading the steps. These are the UI screenshots for Screenshots / Proof of Fix too: embed them once here and have that section's Before and After steps point back to them instead of repeating the images Example: From bf0187072bb360153400a17895e65dc10f4a4110 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 15:49:59 -0700 Subject: [PATCH 007/218] ci: move caching, proxy-extras, gateway and enterprise tests into tests/unit and run them from litellm-tests (#42902) * ci: fix the litellm-tests unit job with sysmon coverage, an env allowlist and coverage upload on failure * test: replace key-dependent proxy, enterprise and mcp unit tests with synthetic values and integration and e2e coverage * test: drop key reads at the legacy proxy, enterprise and mcp paths and wire the gemini pass-through split * ci: move caching, proxy-extras, gateway and enterprise tests into tests/unit and run them from litellm-tests under their legacy flags * ci: move caching, proxy-extras, gateway and enterprise tests into tests/unit and run them from litellm-tests under their legacy flags * ci: fail the unit shard when circleci tests split errors * test: drop restating comments from the gemini pass-through split * ci: exit the unit shard cleanly when circleci tests split assigns it no files --------- Co-authored-by: yuneng --- .circleci/scripts/unit_selection.sh | 63 +++++++++++++++++++ .circleci/tests.yml | 33 +++++++++- .github/scripts/assert_ci_coverage.py | 12 +++- .github/workflows/_test-unit-base.yml | 33 ++++++++-- .github/workflows/test-unit.yml | 22 ++++--- Makefile | 2 +- tests/test_litellm/test_assert_ci_coverage.py | 2 +- .../proxy => unit/caching}/__init__.py | 0 .../caching}/test_cache_preset_key.py | 0 .../caching}/test_caching_handler.py | 0 .../test_responses_stream_cache_keys.py | 0 .../caching}/test_unit_test_caching.py | 0 tests/unit/conftest.py | 2 +- tests/{ => unit}/enterprise/conftest.py | 0 .../send_emails/__init__.py | 0 .../send_emails/test_base_email.py | 0 .../send_emails/test_endpoints.py | 0 .../send_emails/test_resend_email.py | 0 .../send_emails/test_sendgrid_email.py | 0 .../test_prometheus_logging_callbacks.py | 0 .../unit/enterprise/integrations/__init__.py | 0 .../integrations/test_custom_guardrail.py | 0 .../integrations/test_prometheus.py | 0 .../test_prometheus_unit_tests.py | 0 tests/unit/enterprise/proxy/__init__.py | 0 tests/unit/enterprise/proxy/auth/__init__.py | 0 .../proxy/auth/test_route_checks.py | 0 .../proxy/auth/test_user_api_key_auth.py | 0 .../enterprise/proxy/guardrails/__init__.py | 0 .../enterprise}/proxy/guardrails/conftest.py | 0 .../test_apply_guardrail_endpoint.py | 0 .../test_bedrock_apply_guardrail.py | 0 tests/unit/enterprise/proxy/hooks/__init__.py | 0 .../proxy/hooks/test_managed_files.py | 0 .../proxy/management_endpoints/__init__.py | 0 .../test_internal_user_endpoints.py | 0 .../test_project_endpoints_prisma.py | 0 .../test_afile_retrieve_returns_unified_id.py | 0 .../proxy/test_audit_logging_endpoints.py | 0 .../test_batch_retrieve_input_file_id.py | 0 ...trieve_registers_missing_output_file_id.py | 0 ..._retrieve_returns_unified_input_file_id.py | 0 ..._batch_update_db_managed_output_file_id.py | 0 .../test_deleted_file_returns_403_not_404.py | 0 .../proxy/test_enterprise_routes.py | 0 .../proxy/test_file_deletion_blocking.py | 0 .../proxy/test_managed_files_access_check.py | 0 .../proxy/test_managed_files_hook.py | 0 tests/unit/gateway/__init__.py | 0 .../gateway}/test_launch.py | 0 tests/unit/litellm_proxy_extras/__init__.py | 0 .../test_litellm_proxy_extras_logging.py | 0 .../test_litellm_proxy_extras_utils.py | 6 +- 53 files changed, 152 insertions(+), 23 deletions(-) create mode 100755 .circleci/scripts/unit_selection.sh rename tests/{test_litellm/enterprise/proxy => unit/caching}/__init__.py (100%) rename tests/{local_testing => unit/caching}/test_cache_preset_key.py (100%) rename tests/{local_testing => unit/caching}/test_caching_handler.py (100%) rename tests/{local_testing => unit/caching}/test_responses_stream_cache_keys.py (100%) rename tests/{local_testing => unit/caching}/test_unit_test_caching.py (100%) rename tests/{ => unit}/enterprise/conftest.py (100%) create mode 100644 tests/unit/enterprise/enterprise_callbacks/send_emails/__init__.py rename tests/{test_litellm => unit}/enterprise/enterprise_callbacks/send_emails/test_base_email.py (100%) rename tests/{test_litellm => unit}/enterprise/enterprise_callbacks/send_emails/test_endpoints.py (100%) rename tests/{test_litellm => unit}/enterprise/enterprise_callbacks/send_emails/test_resend_email.py (100%) rename tests/{test_litellm => unit}/enterprise/enterprise_callbacks/send_emails/test_sendgrid_email.py (100%) rename tests/{enterprise/litellm_enterprise => unit/enterprise}/enterprise_callbacks/test_prometheus_logging_callbacks.py (100%) create mode 100644 tests/unit/enterprise/integrations/__init__.py rename tests/{enterprise/litellm_enterprise => unit/enterprise}/integrations/test_custom_guardrail.py (100%) rename tests/{enterprise/litellm_enterprise => unit/enterprise}/integrations/test_prometheus.py (100%) rename tests/{enterprise/litellm_enterprise => unit/enterprise}/integrations/test_prometheus_unit_tests.py (100%) create mode 100644 tests/unit/enterprise/proxy/__init__.py create mode 100644 tests/unit/enterprise/proxy/auth/__init__.py rename tests/{enterprise/litellm_enterprise => unit/enterprise}/proxy/auth/test_route_checks.py (100%) rename tests/{enterprise/litellm_enterprise => unit/enterprise}/proxy/auth/test_user_api_key_auth.py (100%) create mode 100644 tests/unit/enterprise/proxy/guardrails/__init__.py rename tests/{enterprise/litellm_enterprise => unit/enterprise}/proxy/guardrails/conftest.py (100%) rename tests/{enterprise/litellm_enterprise => unit/enterprise}/proxy/guardrails/test_apply_guardrail_endpoint.py (100%) rename tests/{enterprise/litellm_enterprise => unit/enterprise}/proxy/guardrails/test_bedrock_apply_guardrail.py (100%) create mode 100644 tests/unit/enterprise/proxy/hooks/__init__.py rename tests/{enterprise/litellm_enterprise => unit/enterprise}/proxy/hooks/test_managed_files.py (100%) create mode 100644 tests/unit/enterprise/proxy/management_endpoints/__init__.py rename tests/{enterprise/litellm_enterprise => unit/enterprise}/proxy/management_endpoints/test_internal_user_endpoints.py (100%) rename tests/{enterprise/litellm_enterprise => unit/enterprise}/proxy/management_endpoints/test_project_endpoints_prisma.py (100%) rename tests/{test_litellm => unit}/enterprise/proxy/test_afile_retrieve_returns_unified_id.py (100%) rename tests/{enterprise/litellm_enterprise => unit/enterprise}/proxy/test_audit_logging_endpoints.py (100%) rename tests/{test_litellm => unit}/enterprise/proxy/test_batch_retrieve_input_file_id.py (100%) rename tests/{test_litellm => unit}/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py (100%) rename tests/{test_litellm => unit}/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py (100%) rename tests/{test_litellm => unit}/enterprise/proxy/test_batch_update_db_managed_output_file_id.py (100%) rename tests/{test_litellm => unit}/enterprise/proxy/test_deleted_file_returns_403_not_404.py (100%) rename tests/{test_litellm => unit}/enterprise/proxy/test_enterprise_routes.py (100%) rename tests/{test_litellm => unit}/enterprise/proxy/test_file_deletion_blocking.py (100%) rename tests/{test_litellm => unit}/enterprise/proxy/test_managed_files_access_check.py (100%) rename tests/{test_litellm => unit}/enterprise/proxy/test_managed_files_hook.py (100%) create mode 100644 tests/unit/gateway/__init__.py rename tests/{test_gateway => unit/gateway}/test_launch.py (100%) create mode 100644 tests/unit/litellm_proxy_extras/__init__.py rename tests/{litellm-proxy-extras => unit/litellm_proxy_extras}/test_litellm_proxy_extras_logging.py (100%) rename tests/{litellm-proxy-extras => unit/litellm_proxy_extras}/test_litellm_proxy_extras_utils.py (99%) diff --git a/.circleci/scripts/unit_selection.sh b/.circleci/scripts/unit_selection.sh new file mode 100755 index 00000000000..8d2b8a42691 --- /dev/null +++ b/.circleci/scripts/unit_selection.sh @@ -0,0 +1,63 @@ +#!/usr/bin/env bash +set -euo pipefail + +flag="${1:?usage: unit_selection.sh }" + +legacy_flags=( + caching-local + enterprise-package + enterprise-routing + proxy-extras + proxy-infra +) + +legacy_paths() { + case "$1" in + caching-local) echo tests/unit/caching ;; + enterprise-package) + echo tests/unit/enterprise/integrations + echo tests/unit/enterprise/proxy/auth + echo tests/unit/enterprise/proxy/guardrails + echo tests/unit/enterprise/proxy/hooks + echo tests/unit/enterprise/proxy/management_endpoints + echo tests/unit/enterprise/proxy/test_audit_logging_endpoints.py + echo tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py ;; + enterprise-routing) + echo tests/unit/enterprise/enterprise_callbacks/send_emails + echo tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py + echo tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py + echo tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py + echo tests/unit/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py + echo tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py + echo tests/unit/enterprise/proxy/test_deleted_file_returns_403_not_404.py + echo tests/unit/enterprise/proxy/test_enterprise_routes.py + echo tests/unit/enterprise/proxy/test_file_deletion_blocking.py + echo tests/unit/enterprise/proxy/test_managed_files_access_check.py + echo tests/unit/enterprise/proxy/test_managed_files_hook.py ;; + proxy-extras) echo tests/unit/litellm_proxy_extras ;; + proxy-infra) echo tests/unit/gateway ;; + *) echo "unit_selection.sh: unknown flag $1" >&2; exit 1 ;; + esac +} + +expand() { + while read -r path; do + if [ -d "$path" ]; then + find "$path" -name 'test_*.py' + elif [ -f "$path" ]; then + echo "$path" + else + echo "unit_selection.sh: $path does not exist" >&2 + exit 1 + fi + done +} + +if [ "$flag" = unit ]; then + comm -23 \ + <(find tests/unit -name 'test_*.py' | sort) \ + <(for legacy in "${legacy_flags[@]}"; do legacy_paths "$legacy"; done | expand | sort) + exit 0 +fi + +legacy_paths "$flag" | expand | sort diff --git a/.circleci/tests.yml b/.circleci/tests.yml index 1afb935453f..38fe44bb625 100644 --- a/.circleci/tests.yml +++ b/.circleci/tests.yml @@ -171,6 +171,12 @@ jobs: shards: type: integer default: 6 + workers: + type: integer + default: 4 + dist: + type: string + default: loadscope base_ref: type: string default: "" @@ -199,17 +205,19 @@ jobs: no_output_timeout: 20m command: | mkdir -p test-results/<< parameters.flag >> - selection="$(find tests/unit -name 'test_*.py' | sort)" || { echo "test selection failed for << parameters.flag >>"; exit 1; } - [ -n "${selection}" ] || { echo "test selection produced no files for << parameters.flag >>"; exit 1; } + selection="$(bash .circleci/scripts/unit_selection.sh << parameters.flag >>)" || { echo "unit_selection.sh failed for << parameters.flag >>"; exit 1; } + [ -n "${selection}" ] || { echo "unit_selection.sh produced no files for << parameters.flag >>"; exit 1; } shard="$(printf '%s\n' "${selection}" | circleci tests split --split-by=timings --timings-type=filename)" || { echo "circleci tests split failed for << parameters.flag >>"; exit 1; } [ -n "${shard}" ] || { echo "shard ${CIRCLE_NODE_INDEX} received no << parameters.flag >> files; nothing to run"; exit 0; } mapfile -t files < <(printf '%s\n' "${shard}") + xdist_args=() + if [ "<< parameters.workers >>" -gt 0 ]; then xdist_args=(-n << parameters.workers >> --dist=<< parameters.dist >>); fi rerun_args=(-p no:rerunfailures) if [ "<< parameters.reruns >>" -gt 0 ]; then rerun_args=(--reruns << parameters.reruns >> --reruns-delay 1 --rerun-except "from pytest-timeout"); fi test_env=(PATH="$PATH" HOME="$HOME" CI=true COVERAGE_CORE="$COVERAGE_CORE" LITELLM_LOCAL_MODEL_COST_MAP="$LITELLM_LOCAL_MODEL_COST_MAP") set +e env -i "${test_env[@]}" \ - uv run --no-sync pytest "${files[@]}" "${rerun_args[@]}" -p no:pytest-retry --timeout=90 -n 4 --dist=loadscope --tb=short --durations=20 -o junit_family=xunit1 --junitxml=test-results/<< parameters.flag >>/junit.xml --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml:coverage.xml --cov-config=pyproject.toml + uv run --no-sync pytest "${files[@]}" "${rerun_args[@]}" -p no:pytest-retry --timeout=90 "${xdist_args[@]}" --tb=short --durations=20 -o junit_family=xunit1 --junitxml=test-results/<< parameters.flag >>/junit.xml --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml:coverage.xml --cov-config=pyproject.toml status=$? set -e if [ "$status" -eq 5 ]; then echo "pytest collected no tests from the shard; passing"; exit 0; fi @@ -293,6 +301,25 @@ workflows: - unit: base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> + - unit: + name: unit-<< matrix.flag >> + shards: 1 + workers: 2 + reruns: 2 + matrix: + parameters: + flag: [caching-local, proxy-extras, enterprise-routing] + base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> + pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> + - unit: + name: unit-<< matrix.flag >> + shards: 1 + reruns: 2 + matrix: + parameters: + flag: [enterprise-package, proxy-infra] + base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> + pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> - documentation - integration: name: integration-<< matrix.suite >> diff --git a/.github/scripts/assert_ci_coverage.py b/.github/scripts/assert_ci_coverage.py index 2e008fe7ade..d8246225a3b 100644 --- a/.github/scripts/assert_ci_coverage.py +++ b/.github/scripts/assert_ci_coverage.py @@ -120,6 +120,13 @@ def _invoked_test_tokens(scalars: Iterable[Scalar]) -> frozenset[str]: ) +def _unit_selection_tokens(repo_root: pathlib.Path = REPO_ROOT) -> frozenset[str]: + script: Final = repo_root / ".circleci/scripts/unit_selection.sh" + if not script.is_file(): + return frozenset() + return frozenset(match.group(0).rstrip("/") for match in TEST_TOKEN_RE.finditer(_uncommented(script.read_text()))) + + def _built_dockerfile_tokens(scalars: Iterable[Scalar]) -> frozenset[str]: return frozenset( match.group(0) @@ -611,7 +618,10 @@ def main() -> int: scalars = _all_scalars() integration_paths, ownership_findings = _integration_ownership() - test_findings = _uncovered_tests(allowlist, _invoked_test_tokens(scalars) | integration_paths) + ownership_findings + test_findings = ( + _uncovered_tests(allowlist, _invoked_test_tokens(scalars) | _unit_selection_tokens() | integration_paths) + + ownership_findings + ) dockerfile_findings = _uncovered_dockerfiles(allowlist, _built_dockerfile_tokens(scalars)) stale_findings = _stale_allowlist_paths(allowlist, test_files=_test_files(), dockerfiles=_dockerfiles()) diff --git a/.github/workflows/_test-unit-base.yml b/.github/workflows/_test-unit-base.yml index 8faddd11df8..ef1dc53b4a6 100644 --- a/.github/workflows/_test-unit-base.yml +++ b/.github/workflows/_test-unit-base.yml @@ -13,6 +13,15 @@ on: have its path existence-checked like any other token. required: true type: string + fork-flag: + description: >- + Codecov flag of the `.circleci/tests.yml` job that now owns part of + this shard. CircleCI does not run on pull requests from forks, so on + those events this shard also runs the files + `.circleci/scripts/unit_selection.sh` lists for the flag. + required: false + type: string + default: "" workers: description: "Number of pytest-xdist workers" required: false @@ -92,6 +101,7 @@ jobs: pull-requests: read outputs: decision: ${{ steps.changes.outputs.decision }} + has-coverage: ${{ steps.tests.outputs.has-coverage }} steps: - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 @@ -160,10 +170,13 @@ jobs: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma - name: Run tests + id: tests if: steps.changes.outputs.decision != 'skip' timeout-minutes: ${{ inputs.timeout-minutes }} env: TEST_PATH: ${{ inputs.test-path }} + FORK_FLAG: ${{ inputs.fork-flag }} + IS_FORK: ${{ github.event_name == 'pull_request' && github.event.pull_request.head.repo.full_name != github.repository }} MAX_FAILURES: ${{ inputs.max-failures }} WORKERS: ${{ inputs.workers }} RERUNS: ${{ inputs.reruns }} @@ -171,9 +184,18 @@ jobs: DIST: ${{ inputs.dist }} COVERAGE_CORE: sysmon run: | + echo "has-coverage=false" >> "$GITHUB_OUTPUT" + selection="${TEST_PATH}" + if [ "${IS_FORK}" = "true" ] && [ -n "${FORK_FLAG}" ]; then + selection="${TEST_PATH} $(bash .circleci/scripts/unit_selection.sh "${FORK_FLAG}" | tr '\n' ' ')" + fi + if [ -z "${selection// /}" ]; then + echo "shard selection is empty on this event (CircleCI flag ${FORK_FLAG:-none} owns it); nothing to run" + exit 0 + fi pytest_args=() existing_paths=0 - for token in ${TEST_PATH:?}; do + for token in ${selection}; do case "${token}" in -*) pytest_args+=("${token}") ;; *) @@ -187,7 +209,7 @@ jobs: esac done if [ "${existing_paths}" -eq 0 ]; then - echo "No path in TEST_PATH exists (${TEST_PATH}); nothing to run" + echo "No path in the selection exists (${selection}); nothing to run" exit 0 fi xdist_args=() @@ -209,8 +231,11 @@ jobs: --cov-config=pyproject.toml status=$? set -e + if [ -f coverage.xml ]; then + echo "has-coverage=true" >> "$GITHUB_OUTPUT" + fi if [ "$status" -eq 5 ]; then - echo "pytest collected no tests from ${TEST_PATH}; passing" + echo "pytest collected no tests from ${selection}; passing" exit 0 fi exit "$status" @@ -226,7 +251,7 @@ jobs: upload-coverage: name: Upload coverage to Codecov needs: run - if: always() && needs.run.outputs.decision != 'skip' + if: always() && needs.run.outputs.decision != 'skip' && needs.run.outputs.has-coverage == 'true' runs-on: ubuntu-latest permissions: contents: read diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index 686bbc89467..94de6040038 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -35,6 +35,10 @@ concurrency: # already a matrix and carries a shard-coverage guard that reads that file by # name. Folding it in here is a follow-up, together with generalising that guard # into assert_ci_coverage.py. +# +# `fork-flag` names the `.circleci/tests.yml` job that now runs part of the +# shard under the same Codecov flag. CircleCI does not build pull requests from +# forks, so the shard still runs those files there and skips them elsewhere. jobs: unit: name: ${{ matrix.shard }} @@ -65,10 +69,10 @@ jobs: - shard: enterprise-routing artifact-name: enterprise-routing test-path: >- - tests/test_litellm/enterprise tests/test_litellm/google_genai tests/test_litellm/router_utils tests/test_litellm/router_strategy + fork-flag: enterprise-routing workers: 2 reruns: 2 timeout-minutes: 20 @@ -200,7 +204,7 @@ jobs: tests/test_litellm/proxy/types_utils tests/test_litellm/proxy/logging_endpoints tests/test_litellm/proxy/test_*.py - tests/test_gateway + fork-flag: proxy-infra workers: 4 reruns: 2 timeout-minutes: 20 @@ -208,11 +212,8 @@ jobs: - shard: caching-local artifact-name: caching-local - test-path: >- - tests/local_testing/test_cache_preset_key.py - tests/local_testing/test_caching_handler.py - tests/local_testing/test_responses_stream_cache_keys.py - tests/local_testing/test_unit_test_caching.py + test-path: "" + fork-flag: caching-local workers: 2 reruns: 2 timeout-minutes: 20 @@ -220,7 +221,8 @@ jobs: - shard: proxy-extras artifact-name: proxy-extras - test-path: "tests/litellm-proxy-extras" + test-path: "" + fork-flag: proxy-extras workers: 2 reruns: 2 timeout-minutes: 20 @@ -228,7 +230,8 @@ jobs: - shard: enterprise-package artifact-name: enterprise-package - test-path: "tests/enterprise" + test-path: "" + fork-flag: enterprise-package workers: 4 reruns: 2 timeout-minutes: 20 @@ -247,6 +250,7 @@ jobs: uses: ./.github/workflows/_test-unit-base.yml with: test-path: ${{ matrix.test-path }} + fork-flag: ${{ matrix.fork-flag || '' }} workers: ${{ matrix.workers }} reruns: ${{ matrix.reruns }} timeout-minutes: ${{ matrix.timeout-minutes }} diff --git a/Makefile b/Makefile index ab7fab6aa99..6263b646c17 100644 --- a/Makefile +++ b/Makefile @@ -332,7 +332,7 @@ test-unit-core-utils: install-test-deps $(UV_RUN) pytest tests/test_litellm/litellm_core_utils --tb=short -vv -n 2 --durations=20 test-unit-other: install-test-deps - $(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/test_litellm/vector_stores tests/test_litellm/a2a_protocol tests/test_litellm/anthropic_interface tests/test_litellm/completion_extras tests/test_litellm/containers tests/test_litellm/enterprise tests/test_litellm/experimental_mcp_client tests/test_litellm/google_genai tests/test_litellm/images tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/test_litellm/types --tb=short -vv -n 4 --durations=20 + $(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/test_litellm/vector_stores tests/test_litellm/a2a_protocol tests/test_litellm/anthropic_interface tests/test_litellm/completion_extras tests/test_litellm/containers tests/unit/enterprise tests/test_litellm/experimental_mcp_client tests/test_litellm/google_genai tests/test_litellm/images tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/test_litellm/types --tb=short -vv -n 4 --durations=20 test-unit-root: install-test-deps $(UV_RUN) pytest tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20 diff --git a/tests/test_litellm/test_assert_ci_coverage.py b/tests/test_litellm/test_assert_ci_coverage.py index 983707db606..8524a905745 100644 --- a/tests/test_litellm/test_assert_ci_coverage.py +++ b/tests/test_litellm/test_assert_ci_coverage.py @@ -340,7 +340,7 @@ def test_a_dockerfile_directory_entry_is_stale_because_only_an_exact_path_exempt def test_a_workflow_that_names_a_file_clears_it_from_the_slice_check(): named = coverage._workflow_named_tokens() assert named, "the workflows must name some test paths or the check proves nothing" - assert any(coverage._token_covers(token, "tests/local_testing/test_caching_handler.py") for token in named) + assert any(coverage._token_covers(token, "tests/proxy_unit_tests/test_proxy_custom_logger.py") for token in named) def test_the_slice_check_credits_only_workflows_never_the_circleci_config(): diff --git a/tests/test_litellm/enterprise/proxy/__init__.py b/tests/unit/caching/__init__.py similarity index 100% rename from tests/test_litellm/enterprise/proxy/__init__.py rename to tests/unit/caching/__init__.py diff --git a/tests/local_testing/test_cache_preset_key.py b/tests/unit/caching/test_cache_preset_key.py similarity index 100% rename from tests/local_testing/test_cache_preset_key.py rename to tests/unit/caching/test_cache_preset_key.py diff --git a/tests/local_testing/test_caching_handler.py b/tests/unit/caching/test_caching_handler.py similarity index 100% rename from tests/local_testing/test_caching_handler.py rename to tests/unit/caching/test_caching_handler.py diff --git a/tests/local_testing/test_responses_stream_cache_keys.py b/tests/unit/caching/test_responses_stream_cache_keys.py similarity index 100% rename from tests/local_testing/test_responses_stream_cache_keys.py rename to tests/unit/caching/test_responses_stream_cache_keys.py diff --git a/tests/local_testing/test_unit_test_caching.py b/tests/unit/caching/test_unit_test_caching.py similarity index 100% rename from tests/local_testing/test_unit_test_caching.py rename to tests/unit/caching/test_unit_test_caching.py diff --git a/tests/unit/conftest.py b/tests/unit/conftest.py index b3bb19a8b8a..202ecb80d7b 100644 --- a/tests/unit/conftest.py +++ b/tests/unit/conftest.py @@ -11,7 +11,7 @@ import litellm # noqa: E402 # litellm reads LITELLM_LOCAL_MODEL_COST_MAP at im import litellm.router as litellm_router_module # noqa: E402 # same import-time dependency import litellm.utils as litellm_utils_module # noqa: E402 # same import-time dependency -LOOPBACK_HOSTS: Final = ["127.0.0.1", "::1"] +LOOPBACK_HOSTS: Final = ["127.0.0.1", "::1", "localhost"] AMBIENT_AZURE_CREDENTIAL_ENV_VARS: Final = ( "AZURE_AD_TOKEN", "AZURE_TENANT_ID", diff --git a/tests/enterprise/conftest.py b/tests/unit/enterprise/conftest.py similarity index 100% rename from tests/enterprise/conftest.py rename to tests/unit/enterprise/conftest.py diff --git a/tests/unit/enterprise/enterprise_callbacks/send_emails/__init__.py b/tests/unit/enterprise/enterprise_callbacks/send_emails/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py b/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py similarity index 100% rename from tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py rename to tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py diff --git a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_endpoints.py b/tests/unit/enterprise/enterprise_callbacks/send_emails/test_endpoints.py similarity index 100% rename from tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_endpoints.py rename to tests/unit/enterprise/enterprise_callbacks/send_emails/test_endpoints.py diff --git a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py b/tests/unit/enterprise/enterprise_callbacks/send_emails/test_resend_email.py similarity index 100% rename from tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py rename to tests/unit/enterprise/enterprise_callbacks/send_emails/test_resend_email.py diff --git a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_sendgrid_email.py b/tests/unit/enterprise/enterprise_callbacks/send_emails/test_sendgrid_email.py similarity index 100% rename from tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_sendgrid_email.py rename to tests/unit/enterprise/enterprise_callbacks/send_emails/test_sendgrid_email.py diff --git a/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py b/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py similarity index 100% rename from tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py rename to tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py diff --git a/tests/unit/enterprise/integrations/__init__.py b/tests/unit/enterprise/integrations/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/enterprise/litellm_enterprise/integrations/test_custom_guardrail.py b/tests/unit/enterprise/integrations/test_custom_guardrail.py similarity index 100% rename from tests/enterprise/litellm_enterprise/integrations/test_custom_guardrail.py rename to tests/unit/enterprise/integrations/test_custom_guardrail.py diff --git a/tests/enterprise/litellm_enterprise/integrations/test_prometheus.py b/tests/unit/enterprise/integrations/test_prometheus.py similarity index 100% rename from tests/enterprise/litellm_enterprise/integrations/test_prometheus.py rename to tests/unit/enterprise/integrations/test_prometheus.py diff --git a/tests/enterprise/litellm_enterprise/integrations/test_prometheus_unit_tests.py b/tests/unit/enterprise/integrations/test_prometheus_unit_tests.py similarity index 100% rename from tests/enterprise/litellm_enterprise/integrations/test_prometheus_unit_tests.py rename to tests/unit/enterprise/integrations/test_prometheus_unit_tests.py diff --git a/tests/unit/enterprise/proxy/__init__.py b/tests/unit/enterprise/proxy/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/enterprise/proxy/auth/__init__.py b/tests/unit/enterprise/proxy/auth/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/enterprise/litellm_enterprise/proxy/auth/test_route_checks.py b/tests/unit/enterprise/proxy/auth/test_route_checks.py similarity index 100% rename from tests/enterprise/litellm_enterprise/proxy/auth/test_route_checks.py rename to tests/unit/enterprise/proxy/auth/test_route_checks.py diff --git a/tests/enterprise/litellm_enterprise/proxy/auth/test_user_api_key_auth.py b/tests/unit/enterprise/proxy/auth/test_user_api_key_auth.py similarity index 100% rename from tests/enterprise/litellm_enterprise/proxy/auth/test_user_api_key_auth.py rename to tests/unit/enterprise/proxy/auth/test_user_api_key_auth.py diff --git a/tests/unit/enterprise/proxy/guardrails/__init__.py b/tests/unit/enterprise/proxy/guardrails/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/enterprise/litellm_enterprise/proxy/guardrails/conftest.py b/tests/unit/enterprise/proxy/guardrails/conftest.py similarity index 100% rename from tests/enterprise/litellm_enterprise/proxy/guardrails/conftest.py rename to tests/unit/enterprise/proxy/guardrails/conftest.py diff --git a/tests/enterprise/litellm_enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py b/tests/unit/enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py similarity index 100% rename from tests/enterprise/litellm_enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py rename to tests/unit/enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py diff --git a/tests/enterprise/litellm_enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py b/tests/unit/enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py similarity index 100% rename from tests/enterprise/litellm_enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py rename to tests/unit/enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py diff --git a/tests/unit/enterprise/proxy/hooks/__init__.py b/tests/unit/enterprise/proxy/hooks/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py b/tests/unit/enterprise/proxy/hooks/test_managed_files.py similarity index 100% rename from tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py rename to tests/unit/enterprise/proxy/hooks/test_managed_files.py diff --git a/tests/unit/enterprise/proxy/management_endpoints/__init__.py b/tests/unit/enterprise/proxy/management_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/unit/enterprise/proxy/management_endpoints/test_internal_user_endpoints.py similarity index 100% rename from tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_internal_user_endpoints.py rename to tests/unit/enterprise/proxy/management_endpoints/test_internal_user_endpoints.py diff --git a/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py similarity index 100% rename from tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py rename to tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py diff --git a/tests/test_litellm/enterprise/proxy/test_afile_retrieve_returns_unified_id.py b/tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py similarity index 100% rename from tests/test_litellm/enterprise/proxy/test_afile_retrieve_returns_unified_id.py rename to tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py diff --git a/tests/enterprise/litellm_enterprise/proxy/test_audit_logging_endpoints.py b/tests/unit/enterprise/proxy/test_audit_logging_endpoints.py similarity index 100% rename from tests/enterprise/litellm_enterprise/proxy/test_audit_logging_endpoints.py rename to tests/unit/enterprise/proxy/test_audit_logging_endpoints.py diff --git a/tests/test_litellm/enterprise/proxy/test_batch_retrieve_input_file_id.py b/tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py similarity index 100% rename from tests/test_litellm/enterprise/proxy/test_batch_retrieve_input_file_id.py rename to tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py diff --git a/tests/test_litellm/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py b/tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py similarity index 100% rename from tests/test_litellm/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py rename to tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py diff --git a/tests/test_litellm/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py b/tests/unit/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py similarity index 100% rename from tests/test_litellm/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py rename to tests/unit/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py diff --git a/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py b/tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py similarity index 100% rename from tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py rename to tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py diff --git a/tests/test_litellm/enterprise/proxy/test_deleted_file_returns_403_not_404.py b/tests/unit/enterprise/proxy/test_deleted_file_returns_403_not_404.py similarity index 100% rename from tests/test_litellm/enterprise/proxy/test_deleted_file_returns_403_not_404.py rename to tests/unit/enterprise/proxy/test_deleted_file_returns_403_not_404.py diff --git a/tests/test_litellm/enterprise/proxy/test_enterprise_routes.py b/tests/unit/enterprise/proxy/test_enterprise_routes.py similarity index 100% rename from tests/test_litellm/enterprise/proxy/test_enterprise_routes.py rename to tests/unit/enterprise/proxy/test_enterprise_routes.py diff --git a/tests/test_litellm/enterprise/proxy/test_file_deletion_blocking.py b/tests/unit/enterprise/proxy/test_file_deletion_blocking.py similarity index 100% rename from tests/test_litellm/enterprise/proxy/test_file_deletion_blocking.py rename to tests/unit/enterprise/proxy/test_file_deletion_blocking.py diff --git a/tests/test_litellm/enterprise/proxy/test_managed_files_access_check.py b/tests/unit/enterprise/proxy/test_managed_files_access_check.py similarity index 100% rename from tests/test_litellm/enterprise/proxy/test_managed_files_access_check.py rename to tests/unit/enterprise/proxy/test_managed_files_access_check.py diff --git a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py b/tests/unit/enterprise/proxy/test_managed_files_hook.py similarity index 100% rename from tests/test_litellm/enterprise/proxy/test_managed_files_hook.py rename to tests/unit/enterprise/proxy/test_managed_files_hook.py diff --git a/tests/unit/gateway/__init__.py b/tests/unit/gateway/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_gateway/test_launch.py b/tests/unit/gateway/test_launch.py similarity index 100% rename from tests/test_gateway/test_launch.py rename to tests/unit/gateway/test_launch.py diff --git a/tests/unit/litellm_proxy_extras/__init__.py b/tests/unit/litellm_proxy_extras/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/litellm-proxy-extras/test_litellm_proxy_extras_logging.py b/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_logging.py similarity index 100% rename from tests/litellm-proxy-extras/test_litellm_proxy_extras_logging.py rename to tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_logging.py diff --git a/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py b/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py similarity index 99% rename from tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py rename to tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py index bb329264a11..755c7617701 100644 --- a/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py +++ b/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py @@ -9,7 +9,7 @@ import pytest sys.path.insert( 0, os.path.abspath( - os.path.join(os.path.dirname(__file__), "../../litellm-proxy-extras") + os.path.join(os.path.dirname(__file__), "../../../litellm-proxy-extras") ), ) @@ -23,7 +23,7 @@ from litellm_proxy_extras.utils import ( _MIGRATIONS_DIR = os.path.abspath( os.path.join( os.path.dirname(__file__), - "../../litellm-proxy-extras/litellm_proxy_extras/migrations", + "../../../litellm-proxy-extras/litellm_proxy_extras/migrations", ) ) @@ -999,7 +999,7 @@ class TestJWTKeyMappingCascade: schema_paths = glob.glob( os.path.abspath( os.path.join( - os.path.dirname(__file__), "../../**/schema.prisma" + os.path.dirname(__file__), "../../../**/schema.prisma" ) ), recursive=True, From 7b25a151bd72dbed46137a89799b298d6da4d87e Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 15:52:44 -0700 Subject: [PATCH 008/218] feat(proxy): let callbacks filter the model listing routes per caller (#43027) * feat(proxy): let callbacks filter the model listing routes per caller * fix(proxy): offer every listed name to the listing callback, agent groups and deployment lookups included * fix(proxy): hide aliases of a team model by its public name and offer /model/info lookups the listed name * fix(proxy): map a malformed model listing filter return to the proxy error contract and document legacy team names --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/integrations/custom_logger.py | 18 + litellm/proxy/proxy_server.py | 101 ++++- litellm/proxy/utils.py | 62 ++- .../proxy_server/test_routes_model_info.py | 13 +- .../proxy/test_model_list_callback_filter.py | 425 ++++++++++++++++++ 5 files changed, 596 insertions(+), 23 deletions(-) create mode 100644 tests/test_litellm/proxy/test_model_list_callback_filter.py diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index 326abd5c6a3..d4162369a35 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -421,6 +421,24 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac ): # raise exception if invalid, return a str for the user to receive - if rejected, or return a modified dictionary for passing into litellm pass + async def async_filter_listed_models( + self, + user_api_key_dict: UserAPIKeyAuth, + model_names: Sequence[str], + ) -> Sequence[str]: + """Runs on the model listing routes (`/v1/models`, `/v1/models/{id}`, `/model/info`, + `/model_group/info`) with the public model names the route would otherwise return, so a + lookup of one model may offer just that name: decide per name, never by position in the + sequence. Return the names to keep as a sequence of strings; a name left out disappears + from every listing, any alias of it offered in the same call goes with it, and + `/v1/models/{id}` answers 404 for it, exactly as for a model that does not exist. Names + outside `model_names` are ignored, so a callback can only narrow the listing, never widen + it. Under `use_team_public_model_name: false`, `/v1/models` and `/model_group/info` list a + team model by its internal routing name while `/model/info` keeps its public name, so hide + both names to hide it on every route. + """ + return model_names + async def async_post_call_response_headers_hook( self, data: dict, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index c0071aa7c81..a869150f7e8 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -152,6 +152,7 @@ from litellm.router_utils.auto_router_tuning_baseline import ( snapshot_tuning_baselines, tuning_limit_violation, ) +from litellm.router_utils.common_utils import resolve_model_group_alias from litellm.router_utils.routing_groups import parse_routing_groups from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.utils import ( @@ -11104,6 +11105,40 @@ class ProxyStartupEvent: #### API ENDPOINTS #### +async def _names_hidden_by_listing_callbacks( + user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str] +) -> frozenset[str]: + hidden: Final = await proxy_logging_obj.hidden_by_listing_callbacks(user_api_key_dict, model_names) + if not hidden or llm_router is None: + return hidden + aliases: Final = llm_router.model_group_alias + internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, general_settings) + return hidden | frozenset( + alias + for alias in aliases + if (target := resolve_model_group_alias(aliases, alias)) is not None + and internal_to_public.get(target, target) in hidden + ) + + +async def _entries_kept_by_listing_callbacks( + entries: Sequence[tuple[str, str]], user_api_key_dict: UserAPIKeyAuth +) -> tuple[tuple[str, str], ...]: + hidden: Final = await _names_hidden_by_listing_callbacks( + user_api_key_dict, tuple(response_id for response_id, _ in entries) + ) + if not hidden: + return tuple(entries) + return tuple(entry for entry in entries if entry[0] not in hidden) + + +async def _deployment_hidden_by_listing_callbacks(deployment: Deployment, user_api_key_dict: UserAPIKeyAuth) -> bool: + listed_name: Final = _translate_model_name_for_response(deployment.model_dump(exclude_none=True)).get("model_name") + if not isinstance(listed_name, str): + return False + return listed_name in await _names_hidden_by_listing_callbacks(user_api_key_dict, (listed_name,)) + + @router.get("/v1/models", dependencies=[Depends(user_api_key_auth)], tags=["model management"]) @router.get( "/models", dependencies=[Depends(user_api_key_auth)], tags=["model management"] @@ -11254,7 +11289,9 @@ async def model_list( # The internal routing key drives the metadata/fallback lookup, while the # public name is what the client sees as the model id. model_data = [] - admin_entries: Final = TeamModelNameTranslator.listing_entries(all_models, llm_router, settings) + admin_entries: Final = await _entries_kept_by_listing_callbacks( + TeamModelNameTranslator.listing_entries(all_models, llm_router, settings), user_api_key_dict + ) for response_id, lookup_id in admin_entries: model_info = create_model_info_response( model_id=lookup_id, @@ -11310,7 +11347,10 @@ async def model_list( # public name is what the client sees as the model id. model_data = [] entries: Final = alias_listing_entries( - TeamModelNameTranslator.listing_entries(all_models, llm_router, settings), caller_aliases + await _entries_kept_by_listing_callbacks( + TeamModelNameTranslator.listing_entries(all_models, llm_router, settings), user_api_key_dict + ), + caller_aliases, ) for response_id, lookup_id in entries: model_info = create_model_info_response( @@ -11404,13 +11444,24 @@ async def model_info( llm_router=llm_router, ) hidden_names: Final = blocked_names | unhealthy_names - if hidden_names: - all_models = [m for m in all_models if m not in hidden_names] + internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, settings) + callback_hidden_names: Final = await _names_hidden_by_listing_callbacks( + user_api_key_dict, + tuple( + response_id + for response_id, _ in TeamModelNameTranslator.listing_entries( + tuple(m for m in all_models if m not in hidden_names), llm_router, settings + ) + ), + ) + if hidden_names or callback_hidden_names: + all_models = [ + m for m in all_models if m not in hidden_names and internal_to_public.get(m, m) not in callback_hidden_names + ] undiscoverable_names: Final = undiscoverable_model_names( all_models, llm_router, user_api_key_dict, team_id or user_api_key_dict.team_id ) - internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, settings) aliased_model_id: Final = alias_target( model_id, caller_alias_maps( @@ -15730,7 +15781,7 @@ async def model_info_v1( if litellm_model_id is not None: # user is trying to get specific model from litellm router deployment_info: Final = llm_router.get_deployment(model_id=litellm_model_id) - if deployment_info is None: + if deployment_info is None or await _deployment_hidden_by_listing_callbacks(deployment_info, user_api_key_dict): raise HTTPException( status_code=400, detail={"error": f"Model id = {litellm_model_id} not found on litellm proxy"}, @@ -15819,10 +15870,17 @@ async def model_info_v1( general_settings=general_settings, llm_router=llm_router, ) - visible_models: Final = discoverable_rows( + servable_rows: Final = discoverable_rows( (model for model in all_models if model.get("model_name") not in hidden_names), user_api_key_dict, ) + listed_names: Final = tuple( + dict.fromkeys(name for model in servable_rows if isinstance(name := model.get("model_name"), str)) + ) + callback_hidden_names: Final = await _names_hidden_by_listing_callbacks(user_api_key_dict, listed_names) + visible_models: Final = tuple( + model for model in servable_rows if model.get("model_name") not in callback_hidden_names + ) verbose_proxy_logger.debug("all_models: %s", visible_models) return _model_info_json_response(visible_models) @@ -15871,7 +15929,7 @@ async def model_deprecations( def _get_model_group_info( - llm_router: Router, all_models_str: list[str], model_group: str | None + llm_router: Router, all_models_str: Sequence[str], model_group: str | None ) -> list[ModelGroupInfoProxy]: model_groups: Final[list[ModelGroupInfoProxy]] = [] @@ -16104,23 +16162,34 @@ async def model_group_info( undiscoverable_group_names: Final = undiscoverable_model_names( all_models_str, llm_router, user_api_key_dict, user_api_key_dict.team_id ) - model_groups: list[ModelGroupInfoProxy] = _get_model_group_info( - llm_router=llm_router, - all_models_str=[name for name in all_models_str if name not in undiscoverable_group_names], - model_group=model_group, - ) + listed_group_names: Final = tuple(name for name in all_models_str if name not in undiscoverable_group_names) # Append A2A agents to model groups from litellm.proxy.agent_endpoints.model_list_helpers import ( append_agents_to_model_group, ) - model_groups = await append_agents_to_model_group( - model_groups=model_groups, + model_groups: Final = await append_agents_to_model_group( + model_groups=_get_model_group_info( + llm_router=llm_router, all_models_str=listed_group_names, model_group=model_group + ), user_api_key_dict=user_api_key_dict, ) + internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, general_settings) + public_group_names: Final = tuple( + internal_to_public.get(group.model_group, group.model_group) for group in model_groups + ) + callback_hidden_names: Final = await _names_hidden_by_listing_callbacks( + user_api_key_dict, tuple(dict.fromkeys(public_group_names)) + ) - return {"data": model_groups} + return { + "data": [ + group + for group, public_name in zip(model_groups, public_group_names, strict=True) + if public_name not in callback_hidden_names + ] + } @router.get( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index c617047fad9..b8cc30ad8a7 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -36,6 +36,7 @@ from typing import ( Final, Generic, Literal, + NoReturn, Optional, Protocol, TypeAlias, @@ -106,7 +107,7 @@ except ImportError: raise ImportError("backoff is not installed. Please install it via 'pip install backoff'") from fastapi import HTTPException, status -from pydantic import TypeAdapter +from pydantic import TypeAdapter, ValidationError import litellm import litellm.litellm_core_utils @@ -1136,11 +1137,51 @@ class _CallbackCapabilities: # avoids the per-request ``get_custom_logger_compatible_class`` walk for # every string entry in ``litellm.callbacks``. resolved_callbacks: tuple[object, ...] = field(default_factory=tuple) + listed_models_filters: tuple[CustomLogger, ...] = field(default_factory=tuple) + + +def _overrides_hook(callback: CustomLogger, hook_name: str) -> bool: + leaf_to_base: Final = takewhile(lambda klass: klass is not CustomLogger, type(callback).__mro__) + return any(hook_name in klass.__dict__ for klass in leaf_to_base) def _overrides_moderation_hook(callback: CustomLogger) -> bool: - leaf_to_base: Final = takewhile(lambda klass: klass is not CustomLogger, type(callback).__mro__) - return any("async_moderation_hook" in klass.__dict__ for klass in leaf_to_base) + return _overrides_hook(callback, "async_moderation_hook") + + +_LISTED_MODEL_NAMES: Final = TypeAdapter(tuple[str, ...]) + + +@dataclass(frozen=True, slots=True) +class MalformedListingFilterReturn: + callback: str + tag: Literal["malformed_listing_filter_return"] = "malformed_listing_filter_return" + + +async def _names_kept_by_listing_callbacks( + callbacks: Sequence[CustomLogger], + user_api_key_dict: UserAPIKeyAuth, + model_names: tuple[str, ...], +) -> tuple[str, ...] | MalformedListingFilterReturn: + if not callbacks or not model_names: + return model_names + returned: Final = await callbacks[0].async_filter_listed_models(user_api_key_dict, model_names) + try: + kept: Final = frozenset(_LISTED_MODEL_NAMES.validate_python(returned)) + except ValidationError: + return MalformedListingFilterReturn(callback=type(callbacks[0]).__name__) + return await _names_kept_by_listing_callbacks( + callbacks[1:], user_api_key_dict, tuple(name for name in model_names if name in kept) + ) + + +def _raise_malformed_listing_filter_return(error: MalformedListingFilterReturn) -> NoReturn: + raise ProxyException( + message=f"{error.callback}.async_filter_listed_models must return a sequence of model names", + type=ProxyErrorTypes.internal_server_error, + param=None, + code=500, + ) class ProxyLogging: @@ -2808,6 +2849,9 @@ class ProxyLogging: has_moderation_override=has_moderation_override, iterator_overrides=tuple(iterator_overrides), resolved_callbacks=tuple(resolved_callbacks), + listed_models_filters=tuple( + callback for callback in resolved_callbacks if _overrides_hook(callback, "async_filter_listed_models") + ), ) # Limit cache to handle test churn without leaking; production # callback lists are stable so this rarely grows past 1 entry. @@ -3715,6 +3759,18 @@ class ProxyLogging: verbose_proxy_logger.exception("Error in post_call_response_headers_hook: %s", str(e)) return merged_headers + async def hidden_by_listing_callbacks( + self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str] + ) -> frozenset[str]: + filters: Final = ProxyLogging._callback_capabilities().listed_models_filters + if not filters: + return frozenset() + candidates: Final = tuple(model_names) + kept: Final = await _names_kept_by_listing_callbacks(filters, user_api_key_dict, candidates) + if isinstance(kept, MalformedListingFilterReturn): + _raise_malformed_listing_filter_return(kept) + return frozenset(candidates).difference(kept) + @staticmethod def _build_litellm_call_info(data: dict, response: object) -> dict[str, object]: """ diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py b/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py index 636dc0f4d77..5175d92084c 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py @@ -322,7 +322,9 @@ def test_get_proxy_model_info_shows_litellm_params_pricing_and_names_it_as_an_ov def test_get_proxy_model_info_names_config_model_info_pricing_as_an_override(monkeypatch, local_model_cost_map): """Pricing declared under ``model_info`` in config.yaml overrides the cost map too.""" info = _enriched_model_info( - monkeypatch, {"model": "openai/gpt-5.6"}, {"id": "dep-config", "db_model": False, "output_cost_per_token": 7e-06} + monkeypatch, + {"model": "openai/gpt-5.6"}, + {"id": "dep-config", "db_model": False, "output_cost_per_token": 7e-06}, ) assert info["pricing_overrides"] == ("output_cost_per_token",) assert info["output_cost_per_token"] == 7e-06 @@ -399,7 +401,9 @@ def test_model_info_reports_null_cost_for_unpriced_deployment_and_zero_for_decla def enriched_cost(model_name: str) -> tuple: deployment = router.get_model_list(model_name=model_name)[0] - info = proxy_server._enrich_model_info_with_litellm_data({**deployment, "model_info": dict(deployment["model_info"])})["model_info"] + info = proxy_server._enrich_model_info_with_litellm_data( + {**deployment, "model_info": dict(deployment["model_info"])} + )["model_info"] return info.get("input_cost_per_token"), info.get("output_cost_per_token") assert enriched_cost("vllm-unpriced") == (None, None) @@ -643,7 +647,6 @@ def model_group_info_router(monkeypatch): monkeypatch.setattr(proxy_server, "user_model", None) monkeypatch.setattr(proxy_server, "general_settings", {}) monkeypatch.setattr(proxy_server, "prisma_client", None) - monkeypatch.setattr(proxy_server, "proxy_logging_obj", None) monkeypatch.setattr(proxy_server, "user_api_key_cache", None) monkeypatch.setattr(proxy_server, "_get_model_group_info", model_group_info) @@ -671,7 +674,9 @@ def test_model_group_info_proxy_admin_ignores_key_model_restriction( @pytest.mark.parametrize("admin_role", ["proxy_admin", "proxy_admin_viewer"]) -def test_model_group_info_proxy_admin_expands_wildcard_deployments(client, auth_as, model_group_info_router, admin_role): +def test_model_group_info_proxy_admin_expands_wildcard_deployments( + client, auth_as, model_group_info_router, admin_role +): from litellm.proxy._types import LitellmUserRoles from litellm.proxy.auth.model_checks import get_known_models_from_wildcard diff --git a/tests/test_litellm/proxy/test_model_list_callback_filter.py b/tests/test_litellm/proxy/test_model_list_callback_filter.py new file mode 100644 index 00000000000..00fbfee24ed --- /dev/null +++ b/tests/test_litellm/proxy/test_model_list_callback_filter.py @@ -0,0 +1,425 @@ +""" +Tests for `CustomLogger.async_filter_listed_models` on the model listing endpoints: +GET /v1/models (`model_list`, OpenAI and Anthropic shapes), GET /v1/models/{id} +(`model_info`), GET /v1/model/info (`model_info_v1`) and GET /model_group/info +(`model_group_info`). A registered callback that overrides the hook decides per +caller which of the names the route would list are kept; the rest disappear and +`/v1/models/{id}` answers 404 for them. +""" + +import json +from collections.abc import Sequence + +import pytest +from fastapi import HTTPException +from starlette.requests import Request + +import litellm +from litellm import Router +from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy import proxy_server +from litellm.proxy._types import LitellmUserRoles, ProxyException, UserAPIKeyAuth +from litellm.proxy.utils import ProxyLogging + + +class _Gate(CustomLogger): + def __init__(self, hidden: frozenset[str] = frozenset(), extra: tuple[str, ...] = ()) -> None: + super().__init__() + self.hidden = hidden + self.extra = extra + self.seen: list[tuple[str, ...]] = [] + + async def async_filter_listed_models( + self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str] + ) -> Sequence[str]: + self.seen.append(tuple(model_names)) + return [*(name for name in model_names if name not in self.hidden), *self.extra] + + +class _InferenceOnlyGate(CustomLogger): + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): + if data.get("model") == "restricted-model": + raise HTTPException(status_code=403, detail="not entitled to this model") + return data + + +class _RaisingGate(CustomLogger): + async def async_filter_listed_models( + self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str] + ) -> Sequence[str]: + raise HTTPException(status_code=503, detail="entitlement service down") + + +class _ReversingGate(CustomLogger): + async def async_filter_listed_models( + self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str] + ) -> Sequence[str]: + return list(reversed(model_names)) + + +class _StringReturningGate(CustomLogger): + async def async_filter_listed_models(self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str]) -> str: + return "open-model" + + +def _deployment(model_name: str, model: str = "openai/gpt-4o", **model_info): + return { + "model_name": model_name, + "litellm_params": {"model": model, "api_key": "sk-fake"}, + "model_info": {"id": f"{model_name}-id", **model_info}, + } + + +def _install_router(monkeypatch, *deployments, **router_kwargs) -> Router: + router = Router(model_list=list(deployments), **router_kwargs) + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(proxy_server, "llm_model_list", router.model_list) + monkeypatch.setattr(proxy_server, "prisma_client", None) + monkeypatch.setattr(proxy_server, "general_settings", {}) + monkeypatch.setattr(proxy_server, "user_model", None) + return router + + +def _register(monkeypatch, *callbacks: CustomLogger) -> None: + monkeypatch.setattr(litellm, "callbacks", list(callbacks)) + ProxyLogging._callback_capabilities_cache.clear() + + +@pytest.fixture +def two_model_router(monkeypatch) -> Router: + return _install_router(monkeypatch, _deployment("open-model"), _deployment("restricted-model")) + + +@pytest.fixture +def team_router(monkeypatch) -> Router: + return _install_router( + monkeypatch, + _deployment("gpt-4"), + _deployment("model_name_team1_abc", team_id="team1", team_public_model_name="team-gpt"), + _deployment("model_name_team1_def", team_id="team1", team_public_model_name="team-chat"), + ) + + +@pytest.fixture +def team_admin_privileges(monkeypatch) -> None: + from litellm.proxy.management_endpoints import common_utils + + async def _is_team_admin(**kwargs) -> bool: + return True + + monkeypatch.setattr(common_utils, "_user_has_admin_privileges", _is_team_admin) + + +def _non_admin(**kwargs) -> UserAPIKeyAuth: + return UserAPIKeyAuth(api_key="sk-test", user_role=LitellmUserRoles.INTERNAL_USER, **kwargs) + + +def _admin() -> UserAPIKeyAuth: + return UserAPIKeyAuth(api_key="sk-test", user_id="u", user_role=LitellmUserRoles.PROXY_ADMIN, team_models=[]) + + +def _team_member() -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key="sk-test", + user_id="u", + user_role=LitellmUserRoles.INTERNAL_USER, + team_id="team1", + team_models=["model_name_team1_abc", "model_name_team1_def"], + models=["model_name_team1_abc", "model_name_team1_def"], + ) + + +def _anthropic_request() -> Request: + return Request( + scope={ + "type": "http", + "method": "GET", + "path": "/v1/models", + "query_string": b"", + "headers": [(b"anthropic-version", b"2023-06-01")], + } + ) + + +async def _v1_models(user_api_key_dict: UserAPIKeyAuth, **kwargs) -> list[str]: + response = await proxy_server.model_list(user_api_key_dict=user_api_key_dict, **kwargs) + return [m["id"] for m in response["data"]] + + +async def _v1_model_info_names(user_api_key_dict: UserAPIKeyAuth) -> list[str]: + response = await proxy_server.model_info_v1(user_api_key_dict=user_api_key_dict) + return [row["model_name"] for row in json.loads(response.body)["data"]] + + +async def _model_groups(user_api_key_dict: UserAPIKeyAuth) -> list[str]: + response = await proxy_server.model_group_info(user_api_key_dict=user_api_key_dict) + return [group.model_group for group in response["data"]] + + +async def _model_by_id_status(model_id: str, user_api_key_dict: UserAPIKeyAuth) -> int: + try: + response = await proxy_server.model_info(model_id=model_id, user_api_key_dict=user_api_key_dict) + except HTTPException as error: + return error.status_code + assert response["id"] == model_id + return 200 + + +@pytest.mark.asyncio +async def test_v1_models_lists_only_the_names_the_callback_keeps(two_model_router, monkeypatch): + _register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"}))) + + assert await _v1_models(_non_admin()) == ["open-model"] + assert await _v1_models(_admin()) == ["open-model"] + assert await _v1_models(_non_admin(), request=_anthropic_request()) == ["open-model"] + + +@pytest.mark.asyncio +async def test_v1_models_scope_expand_applies_the_callback(two_model_router, team_admin_privileges, monkeypatch): + _register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"}))) + + assert await _v1_models(_non_admin(), scope="expand") == ["open-model"] + assert await _v1_models(_admin(), scope="expand") == ["open-model"] + + +@pytest.mark.asyncio +async def test_v1_models_by_id_answers_404_for_a_name_the_callback_leaves_out(two_model_router, monkeypatch): + _register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"}))) + + assert await _model_by_id_status("restricted-model", _non_admin()) == 404 + assert await _model_by_id_status("open-model", _non_admin()) == 200 + + +@pytest.mark.asyncio +async def test_v1_model_info_lists_only_the_rows_the_callback_keeps(two_model_router, monkeypatch): + _register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"}))) + + assert await _v1_model_info_names(_non_admin()) == ["open-model"] + assert await _v1_model_info_names(_admin()) == ["open-model"] + + +@pytest.mark.asyncio +async def test_model_group_info_lists_only_the_groups_the_callback_keeps(two_model_router, monkeypatch): + _register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"}))) + + assert await _model_groups(_non_admin()) == ["open-model"] + assert await _model_groups(_admin()) == ["open-model"] + + +@pytest.mark.asyncio +async def test_a_callback_without_the_hook_changes_no_listing(two_model_router, monkeypatch): + _register(monkeypatch, _InferenceOnlyGate()) + + assert await _v1_models(_non_admin()) == ["open-model", "restricted-model"] + assert await _model_by_id_status("restricted-model", _non_admin()) == 200 + assert await _v1_model_info_names(_non_admin()) == ["open-model", "restricted-model"] + assert await _model_groups(_non_admin()) == ["open-model", "restricted-model"] + + +@pytest.mark.asyncio +async def test_callback_cannot_add_a_name_it_was_not_offered(two_model_router, monkeypatch): + _register(monkeypatch, _Gate(extra=("ghost-model",))) + + assert await _v1_models(_non_admin()) == ["open-model", "restricted-model"] + assert await _model_by_id_status("ghost-model", _non_admin()) == 404 + + +@pytest.mark.asyncio +async def test_callbacks_narrow_in_registration_order(monkeypatch): + _install_router(monkeypatch, _deployment("a"), _deployment("b"), _deployment("c")) + first: _Gate = _Gate(hidden=frozenset({"a"})) + second: _Gate = _Gate(hidden=frozenset({"b"})) + _register(monkeypatch, first, second) + + assert await _v1_models(_non_admin()) == ["c"] + assert first.seen == [("a", "b", "c")] + assert second.seen == [("b", "c")] + + +@pytest.mark.asyncio +async def test_callback_sees_and_filters_team_models_by_their_public_name(team_router, monkeypatch): + gate: _Gate = _Gate(hidden=frozenset({"team-gpt"})) + _register(monkeypatch, gate) + + assert await _v1_models(_team_member()) == ["team-chat"] + assert await _model_by_id_status("team-gpt", _team_member()) == 404 + assert await _model_by_id_status("team-chat", _team_member()) == 200 + assert all("team-gpt" in seen and "model_name_team1_abc" not in seen for seen in gate.seen) + + +@pytest.mark.asyncio +async def test_callback_sees_public_team_names_on_every_listing_route(team_router, monkeypatch): + gate: _Gate = _Gate(hidden=frozenset({"team-gpt"})) + _register(monkeypatch, gate) + + assert await _v1_models(_team_member()) == ["team-chat"] + assert await _v1_model_info_names(_team_member()) == ["team-chat"] + assert await _model_groups(_team_member()) == ["model_name_team1_def"] + assert await _model_by_id_status("team-gpt", _team_member()) == 404 + assert len(gate.seen) == 4 + assert all(sorted(seen) == ["team-chat", "team-gpt"] for seen in gate.seen) + + +@pytest.mark.asyncio +async def test_router_alias_follows_its_hidden_target(monkeypatch): + _install_router( + monkeypatch, + _deployment("open-model"), + _deployment("restricted-model"), + model_group_alias={"mini": "restricted-model", "wide": "open-model"}, + ) + _register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"}))) + + assert sorted(await _v1_models(_non_admin())) == ["open-model", "wide"] + assert sorted(await _v1_model_info_names(_non_admin())) == ["open-model", "wide"] + assert sorted(await _model_groups(_non_admin())) == ["open-model", "wide"] + + _register(monkeypatch, _Gate(hidden=frozenset({"mini"}))) + + assert sorted(await _v1_models(_non_admin())) == ["open-model", "restricted-model", "wide"] + + +@pytest.mark.asyncio +async def test_router_alias_of_a_team_model_follows_its_hidden_public_name(monkeypatch): + _install_router( + monkeypatch, + _deployment("gpt-4"), + _deployment("model_name_team1_abc", team_id="team1", team_public_model_name="team-gpt"), + model_group_alias={"team-alias": "model_name_team1_abc"}, + ) + caller: UserAPIKeyAuth = _non_admin( + user_id="u", + team_id="team1", + team_models=["model_name_team1_abc", "team-alias"], + models=["model_name_team1_abc", "team-alias"], + ) + _register(monkeypatch, _Gate(hidden=frozenset())) + assert sorted(await _v1_models(caller)) == ["team-alias", "team-gpt"] + + _register(monkeypatch, _Gate(hidden=frozenset({"team-gpt"}))) + assert await _v1_models(caller) == [] + assert await _model_groups(caller) == [] + + +@pytest.mark.asyncio +async def test_v1_model_info_offers_only_the_rows_the_caller_would_see(monkeypatch): + _install_router(monkeypatch, _deployment("open-model"), _deployment("hidden-model", discoverable=False)) + gate: _Gate = _Gate() + _register(monkeypatch, gate) + + assert await _v1_model_info_names(_non_admin()) == ["open-model"] + assert await _v1_model_info_names(_admin()) == ["open-model", "hidden-model"] + assert gate.seen == [("open-model",), ("open-model", "hidden-model")] + + +@pytest.mark.asyncio +async def test_listing_keeps_its_order_whatever_order_the_callback_returns(monkeypatch): + _install_router(monkeypatch, _deployment("a"), _deployment("b"), _deployment("c")) + _register(monkeypatch, _ReversingGate()) + + assert await _v1_models(_non_admin()) == ["a", "b", "c"] + assert await _v1_model_info_names(_non_admin()) == ["a", "b", "c"] + + +@pytest.mark.asyncio +async def test_a_callback_returning_a_string_is_an_error_not_an_empty_listing(two_model_router, monkeypatch): + _register(monkeypatch, _StringReturningGate()) + + with pytest.raises(ProxyException, match=r"_StringReturningGate\.async_filter_listed_models") as raised: + await _v1_models(_non_admin()) + assert raised.value.code == "500" + + +@pytest.mark.asyncio +async def test_alias_of_a_hidden_model_is_not_listed(two_model_router, monkeypatch): + _register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"}))) + caller = _non_admin(aliases={"mini": "restricted-model", "wide": "open-model"}) + + assert await _v1_models(caller) == ["open-model", "wide"] + assert await _model_by_id_status("mini", caller) == 404 + assert await _model_by_id_status("wide", caller) == 200 + + +@pytest.mark.asyncio +async def test_callback_error_reaches_the_caller(two_model_router, monkeypatch): + _register(monkeypatch, _RaisingGate()) + + with pytest.raises(HTTPException) as raised: + await _v1_models(_non_admin()) + assert raised.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_hidden_model_still_routes_for_direct_requests(two_model_router, monkeypatch): + _register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"}))) + assert "restricted-model" not in await _v1_models(_non_admin()) + + deployment = two_model_router.get_available_deployment( + model="restricted-model", messages=[{"role": "user", "content": "hi"}] + ) + assert deployment["model_name"] == "restricted-model" + + +@pytest.mark.asyncio +async def test_model_group_info_offers_a2a_agent_groups_to_the_callback(two_model_router, monkeypatch): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry + from litellm.types.agents import AgentResponse + + monkeypatch.setattr( + global_agent_registry, + "agent_list", + [AgentResponse(agent_id="agent-1", agent_name="helper", agent_card_params={})], + ) + caller = _non_admin(object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="p1", agents=["agent-1"])) + gate: _Gate = _Gate() + _register(monkeypatch, gate) + + assert await _model_groups(caller) == ["open-model", "restricted-model", "a2a/helper"] + assert gate.seen == [("open-model", "restricted-model", "a2a/helper")] + + _register(monkeypatch, _Gate(hidden=frozenset({"a2a/helper", "restricted-model"}))) + + assert await _model_groups(caller) == ["open-model"] + + +async def _v1_model_info_by_deployment_id(deployment_id: str, user_api_key_dict: UserAPIKeyAuth) -> int | list[str]: + try: + response = await proxy_server.model_info_v1(user_api_key_dict=user_api_key_dict, litellm_model_id=deployment_id) + except HTTPException as error: + return error.status_code + return [row["model_name"] for row in json.loads(response.body)["data"]] + + +@pytest.mark.asyncio +async def test_v1_model_info_by_deployment_id_answers_like_an_unknown_id_for_a_hidden_model( + two_model_router, monkeypatch +): + _register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"}))) + + assert await _v1_model_info_by_deployment_id("restricted-model-id", _non_admin()) == 400 + assert await _v1_model_info_by_deployment_id("no-such-id", _non_admin()) == 400 + assert await _v1_model_info_by_deployment_id("open-model-id", _non_admin()) == ["open-model"] + + +@pytest.mark.asyncio +async def test_v1_model_info_by_deployment_id_offers_the_public_team_name(team_router, monkeypatch): + gate: _Gate = _Gate(hidden=frozenset({"team-gpt"})) + _register(monkeypatch, gate) + + assert await _v1_model_info_by_deployment_id("model_name_team1_abc-id", _team_member()) == 400 + assert await _v1_model_info_by_deployment_id("model_name_team1_def-id", _team_member()) == ["team-chat"] + assert gate.seen == [("team-gpt",), ("team-chat",)] + + +@pytest.mark.asyncio +async def test_v1_model_info_by_deployment_id_offers_the_name_its_listing_shows_in_legacy_mode( + team_router, monkeypatch +): + monkeypatch.setattr(proxy_server, "general_settings", {"use_team_public_model_name": False}) + gate: _Gate = _Gate(hidden=frozenset({"team-gpt"})) + _register(monkeypatch, gate) + + assert await _v1_model_info_names(_team_member()) == ["team-chat"] + assert await _v1_model_info_by_deployment_id("model_name_team1_abc-id", _team_member()) == 400 + assert gate.seen[-1] == ("team-gpt",) From 248f0eb159c6c2788a28d4671104ffc4ce24904d Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 22:59:11 +0000 Subject: [PATCH 009/218] ci: move tests/proxy_unit_tests to tests/unit/proxy and run the proxy-db shards from litellm-tests (#42903) * ci: fix the litellm-tests unit job with sysmon coverage, an env allowlist and coverage upload on failure * test: replace key-dependent proxy, enterprise and mcp unit tests with synthetic values and integration and e2e coverage * test: drop key reads at the legacy proxy, enterprise and mcp paths and wire the gemini pass-through split * ci: move caching, proxy-extras, gateway and enterprise tests into tests/unit and run them from litellm-tests under their legacy flags * ci: move caching, proxy-extras, gateway and enterprise tests into tests/unit and run them from litellm-tests under their legacy flags * ci: move tests/proxy_unit_tests to tests/unit/proxy and run the proxy-db shards from litellm-tests * ci: fail the unit shard when circleci tests split errors * test: drop restating comments from the gemini pass-through split * build: point the local proxy unit targets at the nested tests/unit/proxy tree * ci: exit the unit shard cleanly when circleci tests split assigns it no files --------- Co-authored-by: yuneng --- .circleci/scripts/classify_changes.sh | 2 +- .circleci/scripts/unit_selection.sh | 72 +++++++++ .circleci/tests.yml | 30 +++- .github/scripts/assert_ci_coverage.py | 1 - .github/workflows/test-unit-proxy-db.yml | 97 ++++------- .github/workflows/test-unit.yml | 8 +- Makefile | 10 +- litellm/llms/litellm_proxy/skills/README.md | 2 +- .../user_api_key_auth_code_coverage.py | 4 +- .../image_endpoints/test_azure_routes.py | 3 +- .../test_litellm/test_circleci_path_filter.py | 2 +- .../proxy/__init__.py} | 0 tests/unit/proxy/auth/__init__.py | 0 .../proxy/auth}/test_auth_checks.py | 0 .../test_default_end_user_budget_simple.py | 0 .../proxy/auth}/test_jwt.py | 0 .../auth}/test_models_fallback_endpoint.py | 0 .../auth}/test_multipart_bypass_repro.py | 0 .../proxy/auth}/test_proxy_routes.py | 0 .../proxy/auth}/test_user_api_key_auth.py | 0 tests/unit/proxy/common_utils/__init__.py | 0 .../common_utils}/test_check_batch_cost.py | 0 .../test_check_responses_cost.py | 0 .../test_proxy_encrypt_decrypt.py | 0 .../common_utils}/test_realtime_cache.py | 0 tests/unit/proxy/conftest.py | 150 ++++++++++++++++++ tests/unit/proxy/db/__init__.py | 0 .../proxy/db/db_transaction_queue/__init__.py | 0 .../test_e2e_pod_lock_manager.py | 0 .../proxy/db}/test_update_daily_tag_spend.py | 0 .../proxy/example_config_yaml/__init__.py | 0 .../example_config_yaml/aliases_config.yaml | 0 .../example_config_yaml/azure_config.yaml | 0 .../example_config_yaml/cache_no_params.yaml | 0 .../cache_with_params.yaml | 0 .../config_with_env_vars.yaml | 0 .../config_with_include.yaml | 0 .../config_with_missing_include.yaml | 0 .../config_with_multiple_includes.yaml | 0 .../example_config_yaml/included_models.yaml | 0 .../example_config_yaml/langfuse_config.yaml | 0 .../example_config_yaml/load_balancer.yaml | 0 .../example_config_yaml/models_file_1.yaml | 0 .../example_config_yaml/models_file_2.yaml | 0 .../opentelemetry_config.yaml | 0 .../example_config_yaml/simple_config.yaml | 0 tests/unit/proxy/google_endpoints/__init__.py | 0 .../test_gemini_agents_endpoints.py | 0 .../test_google_endpoint_routing.py | 0 .../test_google_gemini_proxy_request.py | 0 tests/unit/proxy/hooks/__init__.py | 0 .../proxy/hooks}/test_banned_keyword_list.py | 0 ...test_unit_test_max_model_budget_limiter.py | 0 .../proxy/management_endpoints/__init__.py | 0 .../test_jwt_key_mapping.py | 2 +- .../test_key_generate_prisma.py | 0 .../unit/proxy/management_helpers/__init__.py | 0 .../test_audit_logs_proxy.py | 0 tests/unit/proxy/middleware/__init__.py | 0 .../test_request_size_limit_middleware.py | 0 tests/unit/proxy/public_endpoints/__init__.py | 0 .../test_blog_posts_endpoint.py | 0 tests/unit/proxy/response_polling/__init__.py | 0 .../test_response_polling_handler.py | 0 tests/unit/proxy/spend_tracking/__init__.py | 0 .../test_search_api_logging.py | 0 .../proxy}/test_aproxy_startup.py | 0 tests/unit/proxy/test_configs/__init__.py | 0 .../proxy}/test_configs/custom_auth.py | 0 ...st_cloudflare_azure_with_cache_config.yaml | 0 .../proxy}/test_configs/test_config.yaml | 0 .../test_configs/test_config_custom_auth.yaml | 0 .../test_configs/test_config_no_auth.yaml | 0 .../test_configs/test_guardrails_config.yaml | 0 .../proxy}/test_custom_callback_input.py | 0 .../proxy}/test_custom_logger_s3_gcs.py | 0 .../proxy}/test_custom_tokenizer_bug.py | 0 .../proxy}/test_db_schema_changes.py | 0 .../test_deprecated_key_grace_period.py | 0 .../proxy}/test_get_favicon.py | 0 .../proxy}/test_get_image.py | 0 .../test_prisma_client_backoff_retry.py | 0 .../proxy}/test_prompt_test_endpoint.py | 0 .../proxy}/test_proxy_config_unit_test.py | 2 +- .../proxy}/test_proxy_custom_auth.py | 0 .../proxy}/test_proxy_reject_logging.py | 0 .../proxy}/test_proxy_server.py | 4 +- .../proxy}/test_proxy_setting_guardrails.py | 0 .../proxy}/test_proxy_token_counter.py | 0 .../proxy}/test_proxy_utils.py | 0 .../proxy}/test_reducto_ocr_route.py | 0 .../test_response_polling_pre_call_checks.py | 0 .../proxy}/test_server_root_path.py | 0 .../proxy}/test_ui_path_detection.py | 0 .../proxy}/test_unit_test_proxy_hooks.py | 0 .../proxy}/test_update_spend.py | 0 .../test_zero_cost_model_budget_bypass.py | 0 .../proxy}/vertex_key.json | 0 .../skills}/test_skills_db.py | 2 +- tests/unit/skills/test_skills_main.py | 2 +- 100 files changed, 305 insertions(+), 88 deletions(-) rename tests/{proxy_unit_tests/test_key_generate_dynamodb.py => unit/proxy/__init__.py} (100%) create mode 100644 tests/unit/proxy/auth/__init__.py rename tests/{proxy_unit_tests => unit/proxy/auth}/test_auth_checks.py (100%) rename tests/{proxy_unit_tests => unit/proxy/auth}/test_default_end_user_budget_simple.py (100%) rename tests/{proxy_unit_tests => unit/proxy/auth}/test_jwt.py (100%) rename tests/{proxy_unit_tests => unit/proxy/auth}/test_models_fallback_endpoint.py (100%) rename tests/{proxy_unit_tests => unit/proxy/auth}/test_multipart_bypass_repro.py (100%) rename tests/{proxy_unit_tests => unit/proxy/auth}/test_proxy_routes.py (100%) rename tests/{proxy_unit_tests => unit/proxy/auth}/test_user_api_key_auth.py (100%) create mode 100644 tests/unit/proxy/common_utils/__init__.py rename tests/{proxy_unit_tests => unit/proxy/common_utils}/test_check_batch_cost.py (100%) rename tests/{proxy_unit_tests => unit/proxy/common_utils}/test_check_responses_cost.py (100%) rename tests/{proxy_unit_tests => unit/proxy/common_utils}/test_proxy_encrypt_decrypt.py (100%) rename tests/{proxy_unit_tests => unit/proxy/common_utils}/test_realtime_cache.py (100%) create mode 100644 tests/unit/proxy/conftest.py create mode 100644 tests/unit/proxy/db/__init__.py create mode 100644 tests/unit/proxy/db/db_transaction_queue/__init__.py rename tests/{proxy_unit_tests => unit/proxy/db/db_transaction_queue}/test_e2e_pod_lock_manager.py (100%) rename tests/{proxy_unit_tests => unit/proxy/db}/test_update_daily_tag_spend.py (100%) create mode 100644 tests/unit/proxy/example_config_yaml/__init__.py rename tests/{proxy_unit_tests => unit/proxy}/example_config_yaml/aliases_config.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/example_config_yaml/azure_config.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/example_config_yaml/cache_no_params.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/example_config_yaml/cache_with_params.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/example_config_yaml/config_with_env_vars.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/example_config_yaml/config_with_include.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/example_config_yaml/config_with_missing_include.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/example_config_yaml/config_with_multiple_includes.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/example_config_yaml/included_models.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/example_config_yaml/langfuse_config.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/example_config_yaml/load_balancer.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/example_config_yaml/models_file_1.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/example_config_yaml/models_file_2.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/example_config_yaml/opentelemetry_config.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/example_config_yaml/simple_config.yaml (100%) create mode 100644 tests/unit/proxy/google_endpoints/__init__.py rename tests/{proxy_unit_tests => unit/proxy/google_endpoints}/test_gemini_agents_endpoints.py (100%) rename tests/{proxy_unit_tests => unit/proxy/google_endpoints}/test_google_endpoint_routing.py (100%) rename tests/{proxy_unit_tests => unit/proxy/google_endpoints}/test_google_gemini_proxy_request.py (100%) create mode 100644 tests/unit/proxy/hooks/__init__.py rename tests/{proxy_unit_tests => unit/proxy/hooks}/test_banned_keyword_list.py (100%) rename tests/{proxy_unit_tests => unit/proxy/hooks}/test_unit_test_max_model_budget_limiter.py (100%) create mode 100644 tests/unit/proxy/management_endpoints/__init__.py rename tests/{proxy_unit_tests => unit/proxy/management_endpoints}/test_jwt_key_mapping.py (99%) rename tests/{proxy_unit_tests => unit/proxy/management_endpoints}/test_key_generate_prisma.py (100%) create mode 100644 tests/unit/proxy/management_helpers/__init__.py rename tests/{proxy_unit_tests => unit/proxy/management_helpers}/test_audit_logs_proxy.py (100%) create mode 100644 tests/unit/proxy/middleware/__init__.py rename tests/{proxy_unit_tests => unit/proxy/middleware}/test_request_size_limit_middleware.py (100%) create mode 100644 tests/unit/proxy/public_endpoints/__init__.py rename tests/{proxy_unit_tests => unit/proxy/public_endpoints}/test_blog_posts_endpoint.py (100%) create mode 100644 tests/unit/proxy/response_polling/__init__.py rename tests/{proxy_unit_tests => unit/proxy/response_polling}/test_response_polling_handler.py (100%) create mode 100644 tests/unit/proxy/spend_tracking/__init__.py rename tests/{proxy_unit_tests => unit/proxy/spend_tracking}/test_search_api_logging.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_aproxy_startup.py (100%) create mode 100644 tests/unit/proxy/test_configs/__init__.py rename tests/{proxy_unit_tests => unit/proxy}/test_configs/custom_auth.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_configs/test_cloudflare_azure_with_cache_config.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_configs/test_config.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_configs/test_config_custom_auth.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_configs/test_config_no_auth.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_configs/test_guardrails_config.yaml (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_custom_callback_input.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_custom_logger_s3_gcs.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_custom_tokenizer_bug.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_db_schema_changes.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_deprecated_key_grace_period.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_get_favicon.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_get_image.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_prisma_client_backoff_retry.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_prompt_test_endpoint.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_proxy_config_unit_test.py (99%) rename tests/{proxy_unit_tests => unit/proxy}/test_proxy_custom_auth.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_proxy_reject_logging.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_proxy_server.py (99%) rename tests/{proxy_unit_tests => unit/proxy}/test_proxy_setting_guardrails.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_proxy_token_counter.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_proxy_utils.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_reducto_ocr_route.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_response_polling_pre_call_checks.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_server_root_path.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_ui_path_detection.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_unit_test_proxy_hooks.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_update_spend.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/test_zero_cost_model_budget_bypass.py (100%) rename tests/{proxy_unit_tests => unit/proxy}/vertex_key.json (100%) rename tests/{proxy_unit_tests => unit/skills}/test_skills_db.py (98%) diff --git a/.circleci/scripts/classify_changes.sh b/.circleci/scripts/classify_changes.sh index 01bc8290199..ad265a5e39f 100755 --- a/.circleci/scripts/classify_changes.sh +++ b/.circleci/scripts/classify_changes.sh @@ -31,7 +31,7 @@ while IFS= read -r file || [ -n "$file" ]; do case "$file" in model_prices_and_context_window.json | litellm/model_prices_and_context_window_backup.json | model_prices_and_context_window.schema.json) has_cost_map=true ;; - tests/test_litellm/* | tests/proxy_unit_tests/*) : ;; + tests/test_litellm/* | tests/proxy_unit_tests/* | tests/unit/proxy/*) : ;; *) outside_cost_map_set=true ;; esac done diff --git a/.circleci/scripts/unit_selection.sh b/.circleci/scripts/unit_selection.sh index 8d2b8a42691..2c60c5b1334 100755 --- a/.circleci/scripts/unit_selection.sh +++ b/.circleci/scripts/unit_selection.sh @@ -7,6 +7,18 @@ legacy_flags=( caching-local enterprise-package enterprise-routing + proxy-db-auth-checks + proxy-db-budgets + proxy-db-custom-logging + proxy-db-db-and-spend + proxy-db-endpoints-and-responses + proxy-db-guardrails-hooks + proxy-db-jwt-and-keys + proxy-db-key-generation + proxy-db-logging-misc + proxy-db-proxy-runtime + proxy-db-proxy-server-core + proxy-db-proxy-utils proxy-extras proxy-infra ) @@ -34,6 +46,66 @@ legacy_paths() { echo tests/unit/enterprise/proxy/test_file_deletion_blocking.py echo tests/unit/enterprise/proxy/test_managed_files_access_check.py echo tests/unit/enterprise/proxy/test_managed_files_hook.py ;; + proxy-db-auth-checks) + echo tests/unit/proxy/auth/test_auth_checks.py + echo tests/unit/proxy/auth/test_user_api_key_auth.py + echo tests/unit/proxy/test_deprecated_key_grace_period.py ;; + proxy-db-budgets) + echo tests/unit/proxy/auth/test_default_end_user_budget_simple.py + echo tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py + echo tests/unit/proxy/test_zero_cost_model_budget_bypass.py ;; + proxy-db-custom-logging) + echo tests/unit/proxy/test_custom_callback_input.py + echo tests/unit/proxy/test_custom_logger_s3_gcs.py ;; + proxy-db-db-and-spend) + echo tests/unit/proxy/common_utils/test_proxy_encrypt_decrypt.py + echo tests/unit/proxy/db/db_transaction_queue/test_e2e_pod_lock_manager.py + echo tests/unit/proxy/db/test_update_daily_tag_spend.py + echo tests/unit/proxy/test_db_schema_changes.py + echo tests/unit/proxy/test_prisma_client_backoff_retry.py + echo tests/unit/proxy/test_update_spend.py + echo tests/unit/skills/test_skills_db.py ;; + proxy-db-endpoints-and-responses) + echo tests/unit/proxy/auth/test_models_fallback_endpoint.py + echo tests/unit/proxy/common_utils/test_check_batch_cost.py + echo tests/unit/proxy/common_utils/test_check_responses_cost.py + echo tests/unit/proxy/common_utils/test_realtime_cache.py + echo tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py + echo tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py + echo tests/unit/proxy/google_endpoints/test_google_gemini_proxy_request.py + echo tests/unit/proxy/public_endpoints/test_blog_posts_endpoint.py + echo tests/unit/proxy/response_polling/test_response_polling_handler.py + echo tests/unit/proxy/test_custom_tokenizer_bug.py + echo tests/unit/proxy/test_get_favicon.py + echo tests/unit/proxy/test_get_image.py + echo tests/unit/proxy/test_prompt_test_endpoint.py + echo tests/unit/proxy/test_reducto_ocr_route.py + echo tests/unit/proxy/test_response_polling_pre_call_checks.py + echo tests/unit/proxy/test_ui_path_detection.py ;; + proxy-db-guardrails-hooks) + echo tests/unit/proxy/hooks/test_banned_keyword_list.py + echo tests/unit/proxy/test_proxy_setting_guardrails.py + echo tests/unit/proxy/test_unit_test_proxy_hooks.py ;; + proxy-db-jwt-and-keys) + echo tests/unit/proxy/auth/test_jwt.py + echo tests/unit/proxy/management_endpoints/test_jwt_key_mapping.py + echo tests/unit/proxy/test_proxy_custom_auth.py ;; + proxy-db-key-generation) echo tests/unit/proxy/management_endpoints/test_key_generate_prisma.py ;; + proxy-db-logging-misc) + echo tests/unit/proxy/management_helpers/test_audit_logs_proxy.py + echo tests/unit/proxy/spend_tracking/test_search_api_logging.py + echo tests/unit/proxy/test_proxy_reject_logging.py ;; + proxy-db-proxy-runtime) + echo tests/unit/proxy/auth/test_multipart_bypass_repro.py + echo tests/unit/proxy/auth/test_proxy_routes.py + echo tests/unit/proxy/middleware/test_request_size_limit_middleware.py + echo tests/unit/proxy/test_proxy_config_unit_test.py + echo tests/unit/proxy/test_proxy_token_counter.py + echo tests/unit/proxy/test_server_root_path.py ;; + proxy-db-proxy-server-core) + echo tests/unit/proxy/test_aproxy_startup.py + echo tests/unit/proxy/test_proxy_server.py ;; + proxy-db-proxy-utils) echo tests/unit/proxy/test_proxy_utils.py ;; proxy-extras) echo tests/unit/litellm_proxy_extras ;; proxy-infra) echo tests/unit/gateway ;; *) echo "unit_selection.sh: unknown flag $1" >&2; exit 1 ;; diff --git a/.circleci/tests.yml b/.circleci/tests.yml index 38fe44bb625..08e735637b7 100644 --- a/.circleci/tests.yml +++ b/.circleci/tests.yml @@ -317,7 +317,35 @@ workflows: reruns: 2 matrix: parameters: - flag: [enterprise-package, proxy-infra] + flag: + - enterprise-package + - proxy-infra + - proxy-db-auth-checks + - proxy-db-jwt-and-keys + - proxy-db-proxy-server-core + - proxy-db-proxy-runtime + - proxy-db-custom-logging + - proxy-db-logging-misc + - proxy-db-db-and-spend + - proxy-db-guardrails-hooks + - proxy-db-budgets + - proxy-db-endpoints-and-responses + base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> + pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> + - unit: + name: unit-proxy-db-proxy-utils + flag: proxy-db-proxy-utils + shards: 1 + reruns: 2 + dist: worksteal + base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> + pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> + - unit: + name: unit-proxy-db-key-generation + flag: proxy-db-key-generation + shards: 1 + workers: 0 + reruns: 2 base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> - documentation diff --git a/.github/scripts/assert_ci_coverage.py b/.github/scripts/assert_ci_coverage.py index d8246225a3b..01a01b1034b 100644 --- a/.github/scripts/assert_ci_coverage.py +++ b/.github/scripts/assert_ci_coverage.py @@ -34,7 +34,6 @@ GLOB_CHARS = frozenset("*?") # tests has to be named by some shard or it runs nowhere. A child listed here is # itself decomposed one level deeper and is checked through its own entry. SHARDED_ROOTS: tuple[str, ...] = ( - "tests/proxy_unit_tests", "tests/test_litellm", "tests/test_litellm/proxy", ) diff --git a/.github/workflows/test-unit-proxy-db.yml b/.github/workflows/test-unit-proxy-db.yml index 73015ac6e02..86b385d91a7 100644 --- a/.github/workflows/test-unit-proxy-db.yml +++ b/.github/workflows/test-unit-proxy-db.yml @@ -20,6 +20,12 @@ concurrency: # rather than alphabetical letter ranges. Adding a new test file means adding it # to whichever group it belongs to, not reshuffling slices. # +# `.circleci/tests.yml` runs each group's files on same-repo events under the +# `proxy-db-` Codecov flag; `.circleci/scripts/unit_selection.sh` holds +# the file lists. CircleCI does not build pull requests from forks, so `fork-flag` +# makes the shard run that list there. `test-path` keeps the files that still +# reach real providers and never left tests/proxy_unit_tests. +# # Design targets: # * Every shard runs in <= 7 minutes of wall-clock on the default runner. # Most of a shard's time is pytest plugin load + xdist worker imports + @@ -58,7 +64,7 @@ jobs: proxy-db: needs: assert-shard-coverage # Display only the semantic shard name in the checks UI instead of GHA's - # default "proxy-db (key-generation, tests/proxy_unit_tests/โ€ฆ, 0, loadscope, 20)" + # default "proxy-db (key-generation, tests/unit/proxy/โ€ฆ, 0, loadscope, 20)" # which includes every matrix field and gets truncated past the test-path. name: ${{ matrix.test-group }} permissions: @@ -71,132 +77,93 @@ jobs: include: # Must run serially โ€” event-loop conflict with the logging worker. - test-group: key-generation - test-path: "tests/proxy_unit_tests/test_key_generate_prisma.py" + test-path: "" + fork-flag: proxy-db-key-generation workers: 0 dist: loadscope timeout: 20 # ---- auth: split into 2 shards ---- - test-group: auth-checks - test-path: >- - tests/proxy_unit_tests/test_auth_checks.py - tests/proxy_unit_tests/test_user_api_key_auth.py - tests/proxy_unit_tests/test_deprecated_key_grace_period.py + test-path: "" + fork-flag: proxy-db-auth-checks workers: 4 dist: loadscope timeout: 15 - test-group: jwt-and-keys - test-path: >- - tests/proxy_unit_tests/test_jwt.py - tests/proxy_unit_tests/test_jwt_key_mapping.py - tests/proxy_unit_tests/test_proxy_custom_auth.py - tests/proxy_unit_tests/test_key_generate_dynamodb.py + test-path: "" + fork-flag: proxy-db-jwt-and-keys workers: 4 dist: loadscope timeout: 15 # ---- test_proxy_utils.py, single shard, worksteal distribution ---- - test-group: proxy-utils - test-path: "tests/proxy_unit_tests/test_proxy_utils.py" + test-path: "" + fork-flag: proxy-db-proxy-utils workers: 4 dist: worksteal timeout: 15 # ---- proxy server: split into 2 shards ---- - test-group: proxy-server-core - test-path: >- - tests/proxy_unit_tests/test_proxy_server.py - tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py - tests/proxy_unit_tests/test_aproxy_startup.py + test-path: "tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py" + fork-flag: proxy-db-proxy-server-core workers: 4 dist: loadscope timeout: 15 - test-group: proxy-runtime - test-path: >- - tests/proxy_unit_tests/test_proxy_config_unit_test.py - tests/proxy_unit_tests/test_proxy_routes.py - tests/proxy_unit_tests/test_server_root_path.py - tests/proxy_unit_tests/test_proxy_token_counter.py - tests/proxy_unit_tests/test_request_size_limit_middleware.py - tests/proxy_unit_tests/test_multipart_bypass_repro.py + test-path: "" + fork-flag: proxy-db-proxy-runtime workers: 4 dist: loadscope timeout: 15 # ---- logging: split into 2 shards ---- - test-group: custom-logging - test-path: >- - tests/proxy_unit_tests/test_custom_callback_input.py - tests/proxy_unit_tests/test_custom_logger_s3_gcs.py - tests/proxy_unit_tests/test_proxy_custom_logger.py + test-path: "tests/proxy_unit_tests/test_proxy_custom_logger.py" + fork-flag: proxy-db-custom-logging workers: 4 dist: loadscope timeout: 15 - test-group: logging-misc - test-path: >- - tests/proxy_unit_tests/test_proxy_reject_logging.py - tests/proxy_unit_tests/test_audit_logs_proxy.py - tests/proxy_unit_tests/test_search_api_logging.py + test-path: "" + fork-flag: proxy-db-logging-misc workers: 4 dist: loadscope timeout: 15 - test-group: db-and-spend - test-path: >- - tests/proxy_unit_tests/test_prisma_client_backoff_retry.py - tests/proxy_unit_tests/test_db_schema_changes.py - tests/proxy_unit_tests/test_e2e_pod_lock_manager.py - tests/proxy_unit_tests/test_skills_db.py - tests/proxy_unit_tests/test_update_daily_tag_spend.py - tests/proxy_unit_tests/test_update_spend.py - tests/proxy_unit_tests/test_proxy_encrypt_decrypt.py + test-path: "" + fork-flag: proxy-db-db-and-spend workers: 4 dist: loadscope timeout: 15 # ---- guardrails + budget + hooks: split into 2 ---- - test-group: guardrails-hooks - test-path: >- - tests/proxy_unit_tests/test_proxy_setting_guardrails.py - tests/proxy_unit_tests/test_banned_keyword_list.py - tests/proxy_unit_tests/test_unit_test_proxy_hooks.py + test-path: "" + fork-flag: proxy-db-guardrails-hooks workers: 4 dist: loadscope timeout: 15 - test-group: budgets - test-path: >- - tests/proxy_unit_tests/test_default_end_user_budget_simple.py - tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py - tests/proxy_unit_tests/test_zero_cost_model_budget_bypass.py + test-path: "" + fork-flag: proxy-db-budgets workers: 4 dist: loadscope timeout: 15 - test-group: endpoints-and-responses - test-path: >- - tests/proxy_unit_tests/test_blog_posts_endpoint.py - tests/proxy_unit_tests/test_models_fallback_endpoint.py - tests/proxy_unit_tests/test_google_endpoint_routing.py - tests/proxy_unit_tests/test_google_gemini_proxy_request.py - tests/proxy_unit_tests/test_gemini_agents_endpoints.py - tests/proxy_unit_tests/test_get_favicon.py - tests/proxy_unit_tests/test_get_image.py - tests/proxy_unit_tests/test_reducto_ocr_route.py - tests/proxy_unit_tests/test_ui_path_detection.py - tests/proxy_unit_tests/test_prompt_test_endpoint.py - tests/proxy_unit_tests/test_check_batch_cost.py - tests/proxy_unit_tests/test_check_responses_cost.py - tests/proxy_unit_tests/test_response_polling_handler.py - tests/proxy_unit_tests/test_response_polling_pre_call_checks.py - tests/proxy_unit_tests/test_realtime_cache.py - tests/proxy_unit_tests/test_proxy_exception_mapping.py - tests/proxy_unit_tests/test_custom_tokenizer_bug.py + test-path: "tests/proxy_unit_tests/test_proxy_exception_mapping.py" + fork-flag: proxy-db-endpoints-and-responses workers: 4 dist: loadscope timeout: 15 uses: ./.github/workflows/_test-unit-base.yml with: test-path: ${{ matrix.test-path }} + fork-flag: ${{ matrix.fork-flag }} workers: ${{ matrix.workers }} reruns: 2 timeout-minutes: ${{ matrix.timeout }} diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index 94de6040038..bf2e1602be8 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -31,10 +31,10 @@ concurrency: # number, so a partially-specified entry would fail the call rather than fall # back to the default. # -# tests/proxy_unit_tests keeps its own caller (test-unit-proxy-db.yml): it is -# already a matrix and carries a shard-coverage guard that reads that file by -# name. Folding it in here is a follow-up, together with generalising that guard -# into assert_ci_coverage.py. +# tests/unit/proxy keeps its own caller (test-unit-proxy-db.yml): it is already +# a matrix and carries a shard-coverage guard that reads that file by name. +# Folding it in here is a follow-up, together with generalising that guard into +# assert_ci_coverage.py. # # `fork-flag` names the `.circleci/tests.yml` job that now runs part of the # shard under the same Codecov flag. CircleCI does not build pull requests from diff --git a/Makefile b/Makefile index 6263b646c17..28daf589a23 100644 --- a/Makefile +++ b/Makefile @@ -51,8 +51,8 @@ help: @echo " make test-unit-core-utils - Run core utils tests (~32 files)" @echo " make test-unit-other - Run other tests (caching, responses, etc., ~69 files)" @echo " make test-unit-root - Run root-level tests (~34 files)" - @echo " make test-proxy-unit-a - Run proxy_unit_tests (a-o, ~20 files)" - @echo " make test-proxy-unit-b - Run proxy_unit_tests (p-z, ~28 files)" + @echo " make test-proxy-unit-a - Run tests/unit/proxy (a-o)" + @echo " make test-proxy-unit-b - Run tests/unit/proxy (p-z)" @echo " make test-integration - Run integration tests" @echo " make test-unit-helm - Run helm unit tests" @echo " make test-rust-extension - Build the Rust extension and run its public Python tests" @@ -337,12 +337,12 @@ test-unit-other: install-test-deps test-unit-root: install-test-deps $(UV_RUN) pytest tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20 -# Proxy unit tests (tests/proxy_unit_tests split alphabetically) +# Proxy unit tests (tests/unit/proxy split alphabetically) test-proxy-unit-a: install-test-deps - $(UV_RUN) pytest tests/proxy_unit_tests/test_[a-o]*.py --tb=short -vv -n 2 --durations=20 + $(UV_RUN) pytest tests/unit/proxy --ignore-glob='tests/unit/proxy/test_[p-z]*.py' --tb=short -vv -n 2 --durations=20 test-proxy-unit-b: install-test-deps - $(UV_RUN) pytest tests/proxy_unit_tests/test_[p-z]*.py --tb=short -vv -n 2 --durations=20 + $(UV_RUN) pytest tests/unit/proxy/test_[p-z]*.py tests/unit/skills --tb=short -vv -n 2 --durations=20 test-integration: install-test-deps $(UV_RUN) pytest tests/ -k "not test_litellm" diff --git a/litellm/llms/litellm_proxy/skills/README.md b/litellm/llms/litellm_proxy/skills/README.md index a896aa1166e..ccbd394cddc 100644 --- a/litellm/llms/litellm_proxy/skills/README.md +++ b/litellm/llms/litellm_proxy/skills/README.md @@ -369,7 +369,7 @@ model LiteLLM_SkillsTable { Run the tests: ```bash -pytest tests/proxy_unit_tests/test_skills_db.py -v +pytest tests/unit/skills/test_skills_db.py -v ``` Tests cover: diff --git a/tests/code_coverage_tests/user_api_key_auth_code_coverage.py b/tests/code_coverage_tests/user_api_key_auth_code_coverage.py index a9c2f8ef15f..2f221a7ebe7 100644 --- a/tests/code_coverage_tests/user_api_key_auth_code_coverage.py +++ b/tests/code_coverage_tests/user_api_key_auth_code_coverage.py @@ -31,11 +31,11 @@ def get_function_names_from_file(file_path): def get_all_functions_called_in_tests(base_dir): """ Returns a set of function names that are called in test functions - inside 'local_testing' and 'proxy_unit_tests' directories, + inside 'local_testing' and 'unit/proxy' directories, specifically in files containing the word 'router'. """ called_functions = set() - test_dirs = ["local_testing", "proxy_unit_tests"] + test_dirs = ["local_testing", "unit/proxy"] for test_dir in test_dirs: dir_path = os.path.join(base_dir, test_dir) diff --git a/tests/test_litellm/proxy/image_endpoints/test_azure_routes.py b/tests/test_litellm/proxy/image_endpoints/test_azure_routes.py index 91fff717d25..46fe9a6f893 100644 --- a/tests/test_litellm/proxy/image_endpoints/test_azure_routes.py +++ b/tests/test_litellm/proxy/image_endpoints/test_azure_routes.py @@ -53,7 +53,8 @@ def client_no_auth(): config_fp = ( repo_root / "tests" - / "proxy_unit_tests" + / "unit" + / "proxy" / "test_configs" / "test_config_no_auth.yaml" ) diff --git a/tests/test_litellm/test_circleci_path_filter.py b/tests/test_litellm/test_circleci_path_filter.py index 84e2327057d..dcce7f57113 100644 --- a/tests/test_litellm/test_circleci_path_filter.py +++ b/tests/test_litellm/test_circleci_path_filter.py @@ -107,7 +107,7 @@ CI = [".github/workflows/test-litellm-ui-unit.yml"] ), ( "cost-map-only", - ["model_prices_and_context_window.json", "tests/proxy_unit_tests/test_y.py"], + ["model_prices_and_context_window.json", "tests/unit/proxy/test_y.py"], "run", ), ( diff --git a/tests/proxy_unit_tests/test_key_generate_dynamodb.py b/tests/unit/proxy/__init__.py similarity index 100% rename from tests/proxy_unit_tests/test_key_generate_dynamodb.py rename to tests/unit/proxy/__init__.py diff --git a/tests/unit/proxy/auth/__init__.py b/tests/unit/proxy/auth/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/proxy_unit_tests/test_auth_checks.py b/tests/unit/proxy/auth/test_auth_checks.py similarity index 100% rename from tests/proxy_unit_tests/test_auth_checks.py rename to tests/unit/proxy/auth/test_auth_checks.py diff --git a/tests/proxy_unit_tests/test_default_end_user_budget_simple.py b/tests/unit/proxy/auth/test_default_end_user_budget_simple.py similarity index 100% rename from tests/proxy_unit_tests/test_default_end_user_budget_simple.py rename to tests/unit/proxy/auth/test_default_end_user_budget_simple.py diff --git a/tests/proxy_unit_tests/test_jwt.py b/tests/unit/proxy/auth/test_jwt.py similarity index 100% rename from tests/proxy_unit_tests/test_jwt.py rename to tests/unit/proxy/auth/test_jwt.py diff --git a/tests/proxy_unit_tests/test_models_fallback_endpoint.py b/tests/unit/proxy/auth/test_models_fallback_endpoint.py similarity index 100% rename from tests/proxy_unit_tests/test_models_fallback_endpoint.py rename to tests/unit/proxy/auth/test_models_fallback_endpoint.py diff --git a/tests/proxy_unit_tests/test_multipart_bypass_repro.py b/tests/unit/proxy/auth/test_multipart_bypass_repro.py similarity index 100% rename from tests/proxy_unit_tests/test_multipart_bypass_repro.py rename to tests/unit/proxy/auth/test_multipart_bypass_repro.py diff --git a/tests/proxy_unit_tests/test_proxy_routes.py b/tests/unit/proxy/auth/test_proxy_routes.py similarity index 100% rename from tests/proxy_unit_tests/test_proxy_routes.py rename to tests/unit/proxy/auth/test_proxy_routes.py diff --git a/tests/proxy_unit_tests/test_user_api_key_auth.py b/tests/unit/proxy/auth/test_user_api_key_auth.py similarity index 100% rename from tests/proxy_unit_tests/test_user_api_key_auth.py rename to tests/unit/proxy/auth/test_user_api_key_auth.py diff --git a/tests/unit/proxy/common_utils/__init__.py b/tests/unit/proxy/common_utils/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/proxy_unit_tests/test_check_batch_cost.py b/tests/unit/proxy/common_utils/test_check_batch_cost.py similarity index 100% rename from tests/proxy_unit_tests/test_check_batch_cost.py rename to tests/unit/proxy/common_utils/test_check_batch_cost.py diff --git a/tests/proxy_unit_tests/test_check_responses_cost.py b/tests/unit/proxy/common_utils/test_check_responses_cost.py similarity index 100% rename from tests/proxy_unit_tests/test_check_responses_cost.py rename to tests/unit/proxy/common_utils/test_check_responses_cost.py diff --git a/tests/proxy_unit_tests/test_proxy_encrypt_decrypt.py b/tests/unit/proxy/common_utils/test_proxy_encrypt_decrypt.py similarity index 100% rename from tests/proxy_unit_tests/test_proxy_encrypt_decrypt.py rename to tests/unit/proxy/common_utils/test_proxy_encrypt_decrypt.py diff --git a/tests/proxy_unit_tests/test_realtime_cache.py b/tests/unit/proxy/common_utils/test_realtime_cache.py similarity index 100% rename from tests/proxy_unit_tests/test_realtime_cache.py rename to tests/unit/proxy/common_utils/test_realtime_cache.py diff --git a/tests/unit/proxy/conftest.py b/tests/unit/proxy/conftest.py new file mode 100644 index 00000000000..148751c33f2 --- /dev/null +++ b/tests/unit/proxy/conftest.py @@ -0,0 +1,150 @@ +# conftest.py + +import asyncio +import copy +import inspect +import warnings + +import pytest + + +import litellm +import litellm.proxy.proxy_server + + +# Top-level assignments of these types are the ones importlib.reload(litellm) +# would have effectively reset. We snapshot them at conftest import time and +# deep-copy the snapshot back before every test. +_SNAPSHOT_TYPES = (list, dict, set, tuple, str, int, float, bool, bytes) + + +def _snapshot_mutable_state(module): + """Capture a per-module snapshot of primitive and collection attributes.""" + snapshot = {} + for attr in list(vars(module)): + if attr.startswith("_"): + continue + try: + value = getattr(module, attr) + except Exception as exc: + warnings.warn( + f"conftest: could not read {module.__name__}.{attr} during snapshot: {exc}", + stacklevel=2, + ) + continue + if value is None or isinstance(value, _SNAPSHOT_TYPES): + try: + snapshot[attr] = copy.deepcopy(value) + except Exception as exc: + warnings.warn( + f"conftest: could not snapshot {module.__name__}.{attr}: {exc}", + stacklevel=2, + ) + return snapshot + + +def _restore_mutable_state(module, snapshot): + for attr, default in snapshot.items(): + try: + setattr(module, attr, copy.deepcopy(default)) + except Exception as exc: + warnings.warn( + f"conftest: could not restore {module.__name__}.{attr}: {exc}", + stacklevel=2, + ) + + +def _collect_flushable_caches(): + """Return (module, attr) pairs whose values expose flush_cache().""" + targets = [] + for module in (litellm, litellm.proxy.proxy_server): + for attr in list(vars(module)): + if attr.startswith("_"): + continue + try: + value = getattr(module, attr) + except Exception: + continue + # Only instances โ€” a class reference has an unbound flush_cache + # that can't be called without a self argument. + if inspect.isclass(value) or inspect.ismodule(value): + continue + if callable(getattr(value, "flush_cache", None)): + targets.append((module, attr)) + return targets + + +def _flush_caches(targets): + for module, attr in targets: + try: + value = getattr(module, attr) + except Exception: + continue + flush = getattr(value, "flush_cache", None) + if callable(flush): + try: + flush() + except Exception as exc: + warnings.warn( + f"conftest: flush_cache failed on {module.__name__}.{attr}: {exc}", + stacklevel=2, + ) + + +# Snapshot once at conftest import โ€” these are the "clean" module states. +_LITELLM_STATE = _snapshot_mutable_state(litellm) +_PROXY_SERVER_STATE = _snapshot_mutable_state(litellm.proxy.proxy_server) +_FLUSHABLE_CACHES = _collect_flushable_caches() + + +@pytest.fixture(scope="function", autouse=True) +def setup_and_teardown(): + """Reset mutable module state on litellm and proxy_server before each test. + + Replaces a previous importlib.reload(litellm) approach that cost ~17s + per test (re-executing the full litellm __init__ import chain). + + What IS reset: + - Top-level module attributes of type list / dict / set / tuple + / str / int / float / bool / bytes, and None-valued attributes. + These cover callback lists, general_settings, master_key, + premium_user, prisma_client, etc. โ€” anything the old reload() reset + by re-executing the module body. + - Any module-level object instance that exposes flush_cache() (the + DualCache and LLMClientCache family), which handles cache state + that can't round-trip through deepcopy because of internal locks. + + What is NOT reset: + - Class instances without flush_cache() (e.g. ProxyLogging, + JWTHandler, FastAPI routers, loggers). If a test mutates such an + instance in-place (setattr on the instance, appending to one of + its internal lists, etc.), the mutation will leak into later tests. + Use pytest's monkeypatch.setattr() or a local fixture for those + cases โ€” don't rely on this autouse fixture to undo them. + """ + _restore_mutable_state(litellm, _LITELLM_STATE) + _restore_mutable_state(litellm.proxy.proxy_server, _PROXY_SERVER_STATE) + _flush_caches(_FLUSHABLE_CACHES) + + loop = asyncio.get_event_loop_policy().new_event_loop() + asyncio.set_event_loop(loop) + try: + yield + finally: + loop.close() + asyncio.set_event_loop(None) + + +def pytest_collection_modifyitems(config, items): + # Separate tests in 'test_amazing_proxy_custom_logger.py' and other tests + custom_logger_tests = [ + item for item in items if "custom_logger" in item.parent.name + ] + other_tests = [item for item in items if "custom_logger" not in item.parent.name] + + # Sort tests based on their names + custom_logger_tests.sort(key=lambda x: x.name) + other_tests.sort(key=lambda x: x.name) + + # Reorder the items list + items[:] = custom_logger_tests + other_tests diff --git a/tests/unit/proxy/db/__init__.py b/tests/unit/proxy/db/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/db/db_transaction_queue/__init__.py b/tests/unit/proxy/db/db_transaction_queue/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/proxy_unit_tests/test_e2e_pod_lock_manager.py b/tests/unit/proxy/db/db_transaction_queue/test_e2e_pod_lock_manager.py similarity index 100% rename from tests/proxy_unit_tests/test_e2e_pod_lock_manager.py rename to tests/unit/proxy/db/db_transaction_queue/test_e2e_pod_lock_manager.py diff --git a/tests/proxy_unit_tests/test_update_daily_tag_spend.py b/tests/unit/proxy/db/test_update_daily_tag_spend.py similarity index 100% rename from tests/proxy_unit_tests/test_update_daily_tag_spend.py rename to tests/unit/proxy/db/test_update_daily_tag_spend.py diff --git a/tests/unit/proxy/example_config_yaml/__init__.py b/tests/unit/proxy/example_config_yaml/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/proxy_unit_tests/example_config_yaml/aliases_config.yaml b/tests/unit/proxy/example_config_yaml/aliases_config.yaml similarity index 100% rename from tests/proxy_unit_tests/example_config_yaml/aliases_config.yaml rename to tests/unit/proxy/example_config_yaml/aliases_config.yaml diff --git a/tests/proxy_unit_tests/example_config_yaml/azure_config.yaml b/tests/unit/proxy/example_config_yaml/azure_config.yaml similarity index 100% rename from tests/proxy_unit_tests/example_config_yaml/azure_config.yaml rename to tests/unit/proxy/example_config_yaml/azure_config.yaml diff --git a/tests/proxy_unit_tests/example_config_yaml/cache_no_params.yaml b/tests/unit/proxy/example_config_yaml/cache_no_params.yaml similarity index 100% rename from tests/proxy_unit_tests/example_config_yaml/cache_no_params.yaml rename to tests/unit/proxy/example_config_yaml/cache_no_params.yaml diff --git a/tests/proxy_unit_tests/example_config_yaml/cache_with_params.yaml b/tests/unit/proxy/example_config_yaml/cache_with_params.yaml similarity index 100% rename from tests/proxy_unit_tests/example_config_yaml/cache_with_params.yaml rename to tests/unit/proxy/example_config_yaml/cache_with_params.yaml diff --git a/tests/proxy_unit_tests/example_config_yaml/config_with_env_vars.yaml b/tests/unit/proxy/example_config_yaml/config_with_env_vars.yaml similarity index 100% rename from tests/proxy_unit_tests/example_config_yaml/config_with_env_vars.yaml rename to tests/unit/proxy/example_config_yaml/config_with_env_vars.yaml diff --git a/tests/proxy_unit_tests/example_config_yaml/config_with_include.yaml b/tests/unit/proxy/example_config_yaml/config_with_include.yaml similarity index 100% rename from tests/proxy_unit_tests/example_config_yaml/config_with_include.yaml rename to tests/unit/proxy/example_config_yaml/config_with_include.yaml diff --git a/tests/proxy_unit_tests/example_config_yaml/config_with_missing_include.yaml b/tests/unit/proxy/example_config_yaml/config_with_missing_include.yaml similarity index 100% rename from tests/proxy_unit_tests/example_config_yaml/config_with_missing_include.yaml rename to tests/unit/proxy/example_config_yaml/config_with_missing_include.yaml diff --git a/tests/proxy_unit_tests/example_config_yaml/config_with_multiple_includes.yaml b/tests/unit/proxy/example_config_yaml/config_with_multiple_includes.yaml similarity index 100% rename from tests/proxy_unit_tests/example_config_yaml/config_with_multiple_includes.yaml rename to tests/unit/proxy/example_config_yaml/config_with_multiple_includes.yaml diff --git a/tests/proxy_unit_tests/example_config_yaml/included_models.yaml b/tests/unit/proxy/example_config_yaml/included_models.yaml similarity index 100% rename from tests/proxy_unit_tests/example_config_yaml/included_models.yaml rename to tests/unit/proxy/example_config_yaml/included_models.yaml diff --git a/tests/proxy_unit_tests/example_config_yaml/langfuse_config.yaml b/tests/unit/proxy/example_config_yaml/langfuse_config.yaml similarity index 100% rename from tests/proxy_unit_tests/example_config_yaml/langfuse_config.yaml rename to tests/unit/proxy/example_config_yaml/langfuse_config.yaml diff --git a/tests/proxy_unit_tests/example_config_yaml/load_balancer.yaml b/tests/unit/proxy/example_config_yaml/load_balancer.yaml similarity index 100% rename from tests/proxy_unit_tests/example_config_yaml/load_balancer.yaml rename to tests/unit/proxy/example_config_yaml/load_balancer.yaml diff --git a/tests/proxy_unit_tests/example_config_yaml/models_file_1.yaml b/tests/unit/proxy/example_config_yaml/models_file_1.yaml similarity index 100% rename from tests/proxy_unit_tests/example_config_yaml/models_file_1.yaml rename to tests/unit/proxy/example_config_yaml/models_file_1.yaml diff --git a/tests/proxy_unit_tests/example_config_yaml/models_file_2.yaml b/tests/unit/proxy/example_config_yaml/models_file_2.yaml similarity index 100% rename from tests/proxy_unit_tests/example_config_yaml/models_file_2.yaml rename to tests/unit/proxy/example_config_yaml/models_file_2.yaml diff --git a/tests/proxy_unit_tests/example_config_yaml/opentelemetry_config.yaml b/tests/unit/proxy/example_config_yaml/opentelemetry_config.yaml similarity index 100% rename from tests/proxy_unit_tests/example_config_yaml/opentelemetry_config.yaml rename to tests/unit/proxy/example_config_yaml/opentelemetry_config.yaml diff --git a/tests/proxy_unit_tests/example_config_yaml/simple_config.yaml b/tests/unit/proxy/example_config_yaml/simple_config.yaml similarity index 100% rename from tests/proxy_unit_tests/example_config_yaml/simple_config.yaml rename to tests/unit/proxy/example_config_yaml/simple_config.yaml diff --git a/tests/unit/proxy/google_endpoints/__init__.py b/tests/unit/proxy/google_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/proxy_unit_tests/test_gemini_agents_endpoints.py b/tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py similarity index 100% rename from tests/proxy_unit_tests/test_gemini_agents_endpoints.py rename to tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py diff --git a/tests/proxy_unit_tests/test_google_endpoint_routing.py b/tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py similarity index 100% rename from tests/proxy_unit_tests/test_google_endpoint_routing.py rename to tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py diff --git a/tests/proxy_unit_tests/test_google_gemini_proxy_request.py b/tests/unit/proxy/google_endpoints/test_google_gemini_proxy_request.py similarity index 100% rename from tests/proxy_unit_tests/test_google_gemini_proxy_request.py rename to tests/unit/proxy/google_endpoints/test_google_gemini_proxy_request.py diff --git a/tests/unit/proxy/hooks/__init__.py b/tests/unit/proxy/hooks/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/proxy_unit_tests/test_banned_keyword_list.py b/tests/unit/proxy/hooks/test_banned_keyword_list.py similarity index 100% rename from tests/proxy_unit_tests/test_banned_keyword_list.py rename to tests/unit/proxy/hooks/test_banned_keyword_list.py diff --git a/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py b/tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py similarity index 100% rename from tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py rename to tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py diff --git a/tests/unit/proxy/management_endpoints/__init__.py b/tests/unit/proxy/management_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/proxy_unit_tests/test_jwt_key_mapping.py b/tests/unit/proxy/management_endpoints/test_jwt_key_mapping.py similarity index 99% rename from tests/proxy_unit_tests/test_jwt_key_mapping.py rename to tests/unit/proxy/management_endpoints/test_jwt_key_mapping.py index e95ed42013b..50b7a5c03fd 100644 --- a/tests/proxy_unit_tests/test_jwt_key_mapping.py +++ b/tests/unit/proxy/management_endpoints/test_jwt_key_mapping.py @@ -5,7 +5,7 @@ from datetime import datetime, timezone from unittest.mock import AsyncMock, MagicMock, patch # Add project root to sys.path -sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../.."))) +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../.."))) from litellm.proxy.auth.user_api_key_auth import ( _resolve_jwt_to_virtual_key, diff --git a/tests/proxy_unit_tests/test_key_generate_prisma.py b/tests/unit/proxy/management_endpoints/test_key_generate_prisma.py similarity index 100% rename from tests/proxy_unit_tests/test_key_generate_prisma.py rename to tests/unit/proxy/management_endpoints/test_key_generate_prisma.py diff --git a/tests/unit/proxy/management_helpers/__init__.py b/tests/unit/proxy/management_helpers/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/proxy_unit_tests/test_audit_logs_proxy.py b/tests/unit/proxy/management_helpers/test_audit_logs_proxy.py similarity index 100% rename from tests/proxy_unit_tests/test_audit_logs_proxy.py rename to tests/unit/proxy/management_helpers/test_audit_logs_proxy.py diff --git a/tests/unit/proxy/middleware/__init__.py b/tests/unit/proxy/middleware/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/proxy_unit_tests/test_request_size_limit_middleware.py b/tests/unit/proxy/middleware/test_request_size_limit_middleware.py similarity index 100% rename from tests/proxy_unit_tests/test_request_size_limit_middleware.py rename to tests/unit/proxy/middleware/test_request_size_limit_middleware.py diff --git a/tests/unit/proxy/public_endpoints/__init__.py b/tests/unit/proxy/public_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/proxy_unit_tests/test_blog_posts_endpoint.py b/tests/unit/proxy/public_endpoints/test_blog_posts_endpoint.py similarity index 100% rename from tests/proxy_unit_tests/test_blog_posts_endpoint.py rename to tests/unit/proxy/public_endpoints/test_blog_posts_endpoint.py diff --git a/tests/unit/proxy/response_polling/__init__.py b/tests/unit/proxy/response_polling/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/proxy_unit_tests/test_response_polling_handler.py b/tests/unit/proxy/response_polling/test_response_polling_handler.py similarity index 100% rename from tests/proxy_unit_tests/test_response_polling_handler.py rename to tests/unit/proxy/response_polling/test_response_polling_handler.py diff --git a/tests/unit/proxy/spend_tracking/__init__.py b/tests/unit/proxy/spend_tracking/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/proxy_unit_tests/test_search_api_logging.py b/tests/unit/proxy/spend_tracking/test_search_api_logging.py similarity index 100% rename from tests/proxy_unit_tests/test_search_api_logging.py rename to tests/unit/proxy/spend_tracking/test_search_api_logging.py diff --git a/tests/proxy_unit_tests/test_aproxy_startup.py b/tests/unit/proxy/test_aproxy_startup.py similarity index 100% rename from tests/proxy_unit_tests/test_aproxy_startup.py rename to tests/unit/proxy/test_aproxy_startup.py diff --git a/tests/unit/proxy/test_configs/__init__.py b/tests/unit/proxy/test_configs/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/proxy_unit_tests/test_configs/custom_auth.py b/tests/unit/proxy/test_configs/custom_auth.py similarity index 100% rename from tests/proxy_unit_tests/test_configs/custom_auth.py rename to tests/unit/proxy/test_configs/custom_auth.py diff --git a/tests/proxy_unit_tests/test_configs/test_cloudflare_azure_with_cache_config.yaml b/tests/unit/proxy/test_configs/test_cloudflare_azure_with_cache_config.yaml similarity index 100% rename from tests/proxy_unit_tests/test_configs/test_cloudflare_azure_with_cache_config.yaml rename to tests/unit/proxy/test_configs/test_cloudflare_azure_with_cache_config.yaml diff --git a/tests/proxy_unit_tests/test_configs/test_config.yaml b/tests/unit/proxy/test_configs/test_config.yaml similarity index 100% rename from tests/proxy_unit_tests/test_configs/test_config.yaml rename to tests/unit/proxy/test_configs/test_config.yaml diff --git a/tests/proxy_unit_tests/test_configs/test_config_custom_auth.yaml b/tests/unit/proxy/test_configs/test_config_custom_auth.yaml similarity index 100% rename from tests/proxy_unit_tests/test_configs/test_config_custom_auth.yaml rename to tests/unit/proxy/test_configs/test_config_custom_auth.yaml diff --git a/tests/proxy_unit_tests/test_configs/test_config_no_auth.yaml b/tests/unit/proxy/test_configs/test_config_no_auth.yaml similarity index 100% rename from tests/proxy_unit_tests/test_configs/test_config_no_auth.yaml rename to tests/unit/proxy/test_configs/test_config_no_auth.yaml diff --git a/tests/proxy_unit_tests/test_configs/test_guardrails_config.yaml b/tests/unit/proxy/test_configs/test_guardrails_config.yaml similarity index 100% rename from tests/proxy_unit_tests/test_configs/test_guardrails_config.yaml rename to tests/unit/proxy/test_configs/test_guardrails_config.yaml diff --git a/tests/proxy_unit_tests/test_custom_callback_input.py b/tests/unit/proxy/test_custom_callback_input.py similarity index 100% rename from tests/proxy_unit_tests/test_custom_callback_input.py rename to tests/unit/proxy/test_custom_callback_input.py diff --git a/tests/proxy_unit_tests/test_custom_logger_s3_gcs.py b/tests/unit/proxy/test_custom_logger_s3_gcs.py similarity index 100% rename from tests/proxy_unit_tests/test_custom_logger_s3_gcs.py rename to tests/unit/proxy/test_custom_logger_s3_gcs.py diff --git a/tests/proxy_unit_tests/test_custom_tokenizer_bug.py b/tests/unit/proxy/test_custom_tokenizer_bug.py similarity index 100% rename from tests/proxy_unit_tests/test_custom_tokenizer_bug.py rename to tests/unit/proxy/test_custom_tokenizer_bug.py diff --git a/tests/proxy_unit_tests/test_db_schema_changes.py b/tests/unit/proxy/test_db_schema_changes.py similarity index 100% rename from tests/proxy_unit_tests/test_db_schema_changes.py rename to tests/unit/proxy/test_db_schema_changes.py diff --git a/tests/proxy_unit_tests/test_deprecated_key_grace_period.py b/tests/unit/proxy/test_deprecated_key_grace_period.py similarity index 100% rename from tests/proxy_unit_tests/test_deprecated_key_grace_period.py rename to tests/unit/proxy/test_deprecated_key_grace_period.py diff --git a/tests/proxy_unit_tests/test_get_favicon.py b/tests/unit/proxy/test_get_favicon.py similarity index 100% rename from tests/proxy_unit_tests/test_get_favicon.py rename to tests/unit/proxy/test_get_favicon.py diff --git a/tests/proxy_unit_tests/test_get_image.py b/tests/unit/proxy/test_get_image.py similarity index 100% rename from tests/proxy_unit_tests/test_get_image.py rename to tests/unit/proxy/test_get_image.py diff --git a/tests/proxy_unit_tests/test_prisma_client_backoff_retry.py b/tests/unit/proxy/test_prisma_client_backoff_retry.py similarity index 100% rename from tests/proxy_unit_tests/test_prisma_client_backoff_retry.py rename to tests/unit/proxy/test_prisma_client_backoff_retry.py diff --git a/tests/proxy_unit_tests/test_prompt_test_endpoint.py b/tests/unit/proxy/test_prompt_test_endpoint.py similarity index 100% rename from tests/proxy_unit_tests/test_prompt_test_endpoint.py rename to tests/unit/proxy/test_prompt_test_endpoint.py diff --git a/tests/proxy_unit_tests/test_proxy_config_unit_test.py b/tests/unit/proxy/test_proxy_config_unit_test.py similarity index 99% rename from tests/proxy_unit_tests/test_proxy_config_unit_test.py rename to tests/unit/proxy/test_proxy_config_unit_test.py index 5f236806685..2181c932586 100644 --- a/tests/proxy_unit_tests/test_proxy_config_unit_test.py +++ b/tests/unit/proxy/test_proxy_config_unit_test.py @@ -31,7 +31,7 @@ async def test_basic_reading_configs_from_files(): example_config_yaml_path = os.path.join(current_path, "example_config_yaml") # get all the files from example_config_yaml - files = os.listdir(example_config_yaml_path) + files = [f for f in os.listdir(example_config_yaml_path) if f.endswith((".yaml", ".yml"))] print(files) for file in files: diff --git a/tests/proxy_unit_tests/test_proxy_custom_auth.py b/tests/unit/proxy/test_proxy_custom_auth.py similarity index 100% rename from tests/proxy_unit_tests/test_proxy_custom_auth.py rename to tests/unit/proxy/test_proxy_custom_auth.py diff --git a/tests/proxy_unit_tests/test_proxy_reject_logging.py b/tests/unit/proxy/test_proxy_reject_logging.py similarity index 100% rename from tests/proxy_unit_tests/test_proxy_reject_logging.py rename to tests/unit/proxy/test_proxy_reject_logging.py diff --git a/tests/proxy_unit_tests/test_proxy_server.py b/tests/unit/proxy/test_proxy_server.py similarity index 99% rename from tests/proxy_unit_tests/test_proxy_server.py rename to tests/unit/proxy/test_proxy_server.py index 5be27b3ad72..eae80f311d8 100644 --- a/tests/proxy_unit_tests/test_proxy_server.py +++ b/tests/unit/proxy/test_proxy_server.py @@ -477,7 +477,7 @@ async def test_team_disable_guardrails(mock_acompletion, client_no_auth): assert e.code == str(403) -from test_custom_callback_input import CompletionCustomHandler +from tests.unit.proxy.test_custom_callback_input import CompletionCustomHandler @mock_patch_acompletion() @@ -1114,7 +1114,7 @@ from litellm.proxy._types import ( ) from litellm.proxy.management_endpoints.internal_user_endpoints import new_user from litellm.proxy.management_endpoints.team_endpoints import team_member_add -from test_key_generate_prisma import prisma_client +from tests.unit.proxy.management_endpoints.test_key_generate_prisma import prisma_client @pytest.fixture diff --git a/tests/proxy_unit_tests/test_proxy_setting_guardrails.py b/tests/unit/proxy/test_proxy_setting_guardrails.py similarity index 100% rename from tests/proxy_unit_tests/test_proxy_setting_guardrails.py rename to tests/unit/proxy/test_proxy_setting_guardrails.py diff --git a/tests/proxy_unit_tests/test_proxy_token_counter.py b/tests/unit/proxy/test_proxy_token_counter.py similarity index 100% rename from tests/proxy_unit_tests/test_proxy_token_counter.py rename to tests/unit/proxy/test_proxy_token_counter.py diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/unit/proxy/test_proxy_utils.py similarity index 100% rename from tests/proxy_unit_tests/test_proxy_utils.py rename to tests/unit/proxy/test_proxy_utils.py diff --git a/tests/proxy_unit_tests/test_reducto_ocr_route.py b/tests/unit/proxy/test_reducto_ocr_route.py similarity index 100% rename from tests/proxy_unit_tests/test_reducto_ocr_route.py rename to tests/unit/proxy/test_reducto_ocr_route.py diff --git a/tests/proxy_unit_tests/test_response_polling_pre_call_checks.py b/tests/unit/proxy/test_response_polling_pre_call_checks.py similarity index 100% rename from tests/proxy_unit_tests/test_response_polling_pre_call_checks.py rename to tests/unit/proxy/test_response_polling_pre_call_checks.py diff --git a/tests/proxy_unit_tests/test_server_root_path.py b/tests/unit/proxy/test_server_root_path.py similarity index 100% rename from tests/proxy_unit_tests/test_server_root_path.py rename to tests/unit/proxy/test_server_root_path.py diff --git a/tests/proxy_unit_tests/test_ui_path_detection.py b/tests/unit/proxy/test_ui_path_detection.py similarity index 100% rename from tests/proxy_unit_tests/test_ui_path_detection.py rename to tests/unit/proxy/test_ui_path_detection.py diff --git a/tests/proxy_unit_tests/test_unit_test_proxy_hooks.py b/tests/unit/proxy/test_unit_test_proxy_hooks.py similarity index 100% rename from tests/proxy_unit_tests/test_unit_test_proxy_hooks.py rename to tests/unit/proxy/test_unit_test_proxy_hooks.py diff --git a/tests/proxy_unit_tests/test_update_spend.py b/tests/unit/proxy/test_update_spend.py similarity index 100% rename from tests/proxy_unit_tests/test_update_spend.py rename to tests/unit/proxy/test_update_spend.py diff --git a/tests/proxy_unit_tests/test_zero_cost_model_budget_bypass.py b/tests/unit/proxy/test_zero_cost_model_budget_bypass.py similarity index 100% rename from tests/proxy_unit_tests/test_zero_cost_model_budget_bypass.py rename to tests/unit/proxy/test_zero_cost_model_budget_bypass.py diff --git a/tests/proxy_unit_tests/vertex_key.json b/tests/unit/proxy/vertex_key.json similarity index 100% rename from tests/proxy_unit_tests/vertex_key.json rename to tests/unit/proxy/vertex_key.json diff --git a/tests/proxy_unit_tests/test_skills_db.py b/tests/unit/skills/test_skills_db.py similarity index 98% rename from tests/proxy_unit_tests/test_skills_db.py rename to tests/unit/skills/test_skills_db.py index 8eb07a5ad48..20ffed6fec1 100644 --- a/tests/proxy_unit_tests/test_skills_db.py +++ b/tests/unit/skills/test_skills_db.py @@ -42,7 +42,7 @@ def create_skill_zip(skill_name: str): The zip file is automatically cleaned up after use. """ - test_dir = Path(__file__).parent.parent / "llm_translation" / "test_skills_data" + test_dir = Path(__file__).parents[2] / "llm_translation" / "test_skills_data" skill_dir = test_dir / skill_name # Create a zip file containing the skill directory diff --git a/tests/unit/skills/test_skills_main.py b/tests/unit/skills/test_skills_main.py index e1c66c8d9ea..71d65d45a08 100644 --- a/tests/unit/skills/test_skills_main.py +++ b/tests/unit/skills/test_skills_main.py @@ -30,7 +30,7 @@ def test_create_skill_forwards_description_and_instructions_from_top_level_kwarg def test_create_skill_forwards_description_and_instructions_from_extra_body(monkeypatch) -> None: - """The SDK convention (see tests/proxy_unit_tests/test_skills_db.py) nests them under + """The SDK convention (see tests/unit/skills/test_skills_db.py) nests them under extra_body instead of passing them as top-level kwargs; both paths must reach the DB.""" handler = MagicMock() monkeypatch.setattr(skills_main, "_get_litellm_skills_handler", lambda: handler) From ba776469916bcdb35ab238dec92f33ad428d724d Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 23:07:48 +0000 Subject: [PATCH 010/218] ci: move provider-independent MCP tests into tests/unit and run mcp-integration from litellm-tests (#42904) * ci: fix the litellm-tests unit job with sysmon coverage, an env allowlist and coverage upload on failure * test: replace key-dependent proxy, enterprise and mcp unit tests with synthetic values and integration and e2e coverage * test: drop key reads at the legacy proxy, enterprise and mcp paths and wire the gemini pass-through split * ci: move caching, proxy-extras, gateway and enterprise tests into tests/unit and run them from litellm-tests under their legacy flags * ci: move caching, proxy-extras, gateway and enterprise tests into tests/unit and run them from litellm-tests under their legacy flags * ci: move tests/proxy_unit_tests to tests/unit/proxy and run the proxy-db shards from litellm-tests * ci: move provider-independent MCP tests into tests/unit and run mcp-integration from litellm-tests * ci: fail the unit shard when circleci tests split errors * test: drop restating comments from the gemini pass-through split * build: point the local proxy unit targets at the nested tests/unit/proxy tree * ci: exit the unit shard cleanly when circleci tests split assigns it no files --------- Co-authored-by: yuneng --- .circleci/scripts/unit_selection.sh | 5 ++ .circleci/tests.yml | 21 +++++ .github/workflows/test-unit.yml | 1 + tests/unit/proxy/_experimental/__init__.py | 0 .../_experimental/mcp_server/__init__.py | 0 .../_experimental/mcp_server/conftest.py | 78 +++++++++++++++++++ .../test_mcp_auth_header_extraction.py | 0 .../mcp_server}/test_mcp_auth_priority.py | 0 .../mcp_server}/test_mcp_chat_completions.py | 0 .../mcp_server}/test_mcp_client_unit.py | 0 .../mcp_server}/test_mcp_logging.py | 0 .../mcp_server}/test_mcp_server.py | 0 .../mcp_server}/test_oauth2_mcp_config.yaml | 0 .../mcp_server}/test_openapi_spec_path_url.py | 0 .../mcp_server}/test_per_user_oauth_cache.py | 0 tests/unit/responses/__init__.py | 0 tests/unit/responses/mcp/__init__.py | 0 .../mcp}/test_aresponses_api_with_mcp.py | 0 18 files changed, 105 insertions(+) create mode 100644 tests/unit/proxy/_experimental/__init__.py create mode 100644 tests/unit/proxy/_experimental/mcp_server/__init__.py create mode 100644 tests/unit/proxy/_experimental/mcp_server/conftest.py rename tests/{mcp_tests => unit/proxy/_experimental/mcp_server}/test_mcp_auth_header_extraction.py (100%) rename tests/{mcp_tests => unit/proxy/_experimental/mcp_server}/test_mcp_auth_priority.py (100%) rename tests/{mcp_tests => unit/proxy/_experimental/mcp_server}/test_mcp_chat_completions.py (100%) rename tests/{mcp_tests => unit/proxy/_experimental/mcp_server}/test_mcp_client_unit.py (100%) rename tests/{mcp_tests => unit/proxy/_experimental/mcp_server}/test_mcp_logging.py (100%) rename tests/{mcp_tests => unit/proxy/_experimental/mcp_server}/test_mcp_server.py (100%) rename tests/{mcp_tests => unit/proxy/_experimental/mcp_server}/test_oauth2_mcp_config.yaml (100%) rename tests/{mcp_tests => unit/proxy/_experimental/mcp_server}/test_openapi_spec_path_url.py (100%) rename tests/{mcp_tests => unit/proxy/_experimental/mcp_server}/test_per_user_oauth_cache.py (100%) create mode 100644 tests/unit/responses/__init__.py create mode 100644 tests/unit/responses/mcp/__init__.py rename tests/{mcp_tests => unit/responses/mcp}/test_aresponses_api_with_mcp.py (100%) diff --git a/.circleci/scripts/unit_selection.sh b/.circleci/scripts/unit_selection.sh index 2c60c5b1334..f2ee7550df3 100755 --- a/.circleci/scripts/unit_selection.sh +++ b/.circleci/scripts/unit_selection.sh @@ -7,6 +7,7 @@ legacy_flags=( caching-local enterprise-package enterprise-routing + mcp-integration proxy-db-auth-checks proxy-db-budgets proxy-db-custom-logging @@ -46,6 +47,10 @@ legacy_paths() { echo tests/unit/enterprise/proxy/test_file_deletion_blocking.py echo tests/unit/enterprise/proxy/test_managed_files_access_check.py echo tests/unit/enterprise/proxy/test_managed_files_hook.py ;; + mcp-integration) + echo tests/unit/proxy/_experimental/mcp_server + echo tests/unit/responses/mcp + echo tests/mcp_tests/test_proxy_mcp_e2e.py ;; proxy-db-auth-checks) echo tests/unit/proxy/auth/test_auth_checks.py echo tests/unit/proxy/auth/test_user_api_key_auth.py diff --git a/.circleci/tests.yml b/.circleci/tests.yml index 08e735637b7..264d7695a94 100644 --- a/.circleci/tests.yml +++ b/.circleci/tests.yml @@ -183,6 +183,9 @@ jobs: pull_request_url: type: string default: "" + legacy_mcp_peer: + type: boolean + default: false reruns: type: integer default: 0 @@ -200,6 +203,15 @@ jobs: base_ref: << parameters.base_ref >> pull_request_url: << parameters.pull_request_url >> - setup_test_deps + - when: + condition: << parameters.legacy_mcp_peer >> + steps: + - run: + name: Install MCP SDK1 peer + command: | + uv venv --python 3.12 .venv-mcp-peer + uv pip install --python .venv-mcp-peer 'mcp==1.28.1' 'langchain-mcp-adapters==0.2.1' + echo "export MCP_TEST_PEER_PYTHON=$PWD/.venv-mcp-peer/bin/python" >> "$BASH_ENV" - run: name: "Run << parameters.flag >> shard" no_output_timeout: 20m @@ -215,6 +227,7 @@ jobs: rerun_args=(-p no:rerunfailures) if [ "<< parameters.reruns >>" -gt 0 ]; then rerun_args=(--reruns << parameters.reruns >> --reruns-delay 1 --rerun-except "from pytest-timeout"); fi test_env=(PATH="$PATH" HOME="$HOME" CI=true COVERAGE_CORE="$COVERAGE_CORE" LITELLM_LOCAL_MODEL_COST_MAP="$LITELLM_LOCAL_MODEL_COST_MAP") + if [ -n "${MCP_TEST_PEER_PYTHON:-}" ]; then test_env+=(MCP_TEST_PEER_PYTHON="$MCP_TEST_PEER_PYTHON"); fi set +e env -i "${test_env[@]}" \ uv run --no-sync pytest "${files[@]}" "${rerun_args[@]}" -p no:pytest-retry --timeout=90 "${xdist_args[@]}" --tb=short --durations=20 -o junit_family=xunit1 --junitxml=test-results/<< parameters.flag >>/junit.xml --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml:coverage.xml --cov-config=pyproject.toml @@ -311,6 +324,14 @@ workflows: flag: [caching-local, proxy-extras, enterprise-routing] base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> + - unit: + name: unit-mcp-integration + flag: mcp-integration + shards: 1 + workers: 2 + legacy_mcp_peer: true + base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> + pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> - unit: name: unit-<< matrix.flag >> shards: 1 diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index bf2e1602be8..126a6e26e6f 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -53,6 +53,7 @@ jobs: - shard: mcp-integration artifact-name: mcp-integration test-path: "tests/mcp_tests tests/test_litellm/experimental_mcp_client" + fork-flag: mcp-integration workers: 2 reruns: 0 timeout-minutes: 20 diff --git a/tests/unit/proxy/_experimental/__init__.py b/tests/unit/proxy/_experimental/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/_experimental/mcp_server/__init__.py b/tests/unit/proxy/_experimental/mcp_server/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/_experimental/mcp_server/conftest.py b/tests/unit/proxy/_experimental/mcp_server/conftest.py new file mode 100644 index 00000000000..d8b91e07467 --- /dev/null +++ b/tests/unit/proxy/_experimental/mcp_server/conftest.py @@ -0,0 +1,78 @@ +import asyncio +import importlib + +import pytest + +import litellm +from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER + + +@pytest.fixture(scope="session") +def event_loop(): + try: + loop = asyncio.get_running_loop() + except RuntimeError: + loop = asyncio.new_event_loop() + yield loop + loop.close() + + +@pytest.fixture(scope="function", autouse=True) +def setup_and_teardown(): + """ + This fixture reloads litellm before every function. To speed up testing by removing callbacks being chained. + """ + importlib.reload(litellm) + import asyncio + + loop = asyncio.get_event_loop_policy().new_event_loop() + asyncio.set_event_loop(loop) + yield + + # Teardown code (executes after the yield point) + # LoggingWorker carries still-queued coroutines onto the next test's loop, where they'd log into that test's callbacks + asyncio.run(GLOBAL_LOGGING_WORKER.clear_queue()) + loop.close() # Close the loop created earlier + asyncio.set_event_loop(None) # Remove the reference to the loop + + +@pytest.fixture(scope="function", autouse=True) +async def drain_logging_worker(): + """ + The logging queue is bound to the running loop, so anything left queued when a test's loop + goes away is carried onto the next test's loop and fires against its callbacks. + """ + from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER + + yield + + try: + await asyncio.wait_for(GLOBAL_LOGGING_WORKER.clear_queue(), timeout=10) + except asyncio.TimeoutError: + pass + + +def pytest_collection_modifyitems(config, items): + # Separate tests in 'test_amazing_proxy_custom_logger.py' and other tests + custom_logger_tests = [ + item for item in items if "custom_logger" in item.parent.name + ] + other_tests = [item for item in items if "custom_logger" not in item.parent.name] + + # Sort tests based on their names + custom_logger_tests.sort(key=lambda x: x.name) + other_tests.sort(key=lambda x: x.name) + + # Reorder the items list + items[:] = custom_logger_tests + other_tests + + +@pytest.fixture +def config_only_mcp_manager_factory(): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + + class ConfigOnlyManager(MCPServerManager): + def initialize_tool_name_to_mcp_server_name_mapping(self): + return None + + return ConfigOnlyManager diff --git a/tests/mcp_tests/test_mcp_auth_header_extraction.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_auth_header_extraction.py similarity index 100% rename from tests/mcp_tests/test_mcp_auth_header_extraction.py rename to tests/unit/proxy/_experimental/mcp_server/test_mcp_auth_header_extraction.py diff --git a/tests/mcp_tests/test_mcp_auth_priority.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_auth_priority.py similarity index 100% rename from tests/mcp_tests/test_mcp_auth_priority.py rename to tests/unit/proxy/_experimental/mcp_server/test_mcp_auth_priority.py diff --git a/tests/mcp_tests/test_mcp_chat_completions.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_chat_completions.py similarity index 100% rename from tests/mcp_tests/test_mcp_chat_completions.py rename to tests/unit/proxy/_experimental/mcp_server/test_mcp_chat_completions.py diff --git a/tests/mcp_tests/test_mcp_client_unit.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_client_unit.py similarity index 100% rename from tests/mcp_tests/test_mcp_client_unit.py rename to tests/unit/proxy/_experimental/mcp_server/test_mcp_client_unit.py diff --git a/tests/mcp_tests/test_mcp_logging.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_logging.py similarity index 100% rename from tests/mcp_tests/test_mcp_logging.py rename to tests/unit/proxy/_experimental/mcp_server/test_mcp_logging.py diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py similarity index 100% rename from tests/mcp_tests/test_mcp_server.py rename to tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py diff --git a/tests/mcp_tests/test_oauth2_mcp_config.yaml b/tests/unit/proxy/_experimental/mcp_server/test_oauth2_mcp_config.yaml similarity index 100% rename from tests/mcp_tests/test_oauth2_mcp_config.yaml rename to tests/unit/proxy/_experimental/mcp_server/test_oauth2_mcp_config.yaml diff --git a/tests/mcp_tests/test_openapi_spec_path_url.py b/tests/unit/proxy/_experimental/mcp_server/test_openapi_spec_path_url.py similarity index 100% rename from tests/mcp_tests/test_openapi_spec_path_url.py rename to tests/unit/proxy/_experimental/mcp_server/test_openapi_spec_path_url.py diff --git a/tests/mcp_tests/test_per_user_oauth_cache.py b/tests/unit/proxy/_experimental/mcp_server/test_per_user_oauth_cache.py similarity index 100% rename from tests/mcp_tests/test_per_user_oauth_cache.py rename to tests/unit/proxy/_experimental/mcp_server/test_per_user_oauth_cache.py diff --git a/tests/unit/responses/__init__.py b/tests/unit/responses/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/responses/mcp/__init__.py b/tests/unit/responses/mcp/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/mcp_tests/test_aresponses_api_with_mcp.py b/tests/unit/responses/mcp/test_aresponses_api_with_mcp.py similarity index 100% rename from tests/mcp_tests/test_aresponses_api_with_mcp.py rename to tests/unit/responses/mcp/test_aresponses_api_with_mcp.py From 3fa688223d97bb46345c8ed8032b583e03e29218 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 16:15:58 -0700 Subject: [PATCH 011/218] fix(vertex_ai): translate /v1/responses batch rows through the Responses-to-Chat bridge (#43042) * fix(vertex_ai): translate /v1/responses batch rows through the Responses-to-Chat bridge Vertex batch uploads treated every non-embeddings JSONL row as a chat completions body, so a /v1/responses row lost its input and reached GCS as a blank text part. Route detection now recognizes /v1/responses rows and bridges them to chat through the same Responses-to-Chat bridge the real-time path uses. That bridge call moves out of the Bedrock files transformation into a shared helper both providers call, forwarding the record's fields as sent, like real time, instead of validating them against the SDK TypedDicts whose required keys clients omit. * chore(batches): type the Vertex responses test helper and drop the quoted input cast * fix(batches): translate developer messages to system on Vertex and Bedrock batch rows like real time --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/llms/base_llm/files/batch_records.py | 54 +++++++ litellm/llms/bedrock/files/transformation.py | 53 +----- .../llms/vertex_ai/files/transformation.py | 32 +++- .../test_bedrock_files_transformation.py | 43 +++++ .../test_vertex_ai_files_transformation.py | 153 ++++++++++++++++++ 5 files changed, 280 insertions(+), 55 deletions(-) create mode 100644 litellm/llms/base_llm/files/batch_records.py diff --git a/litellm/llms/base_llm/files/batch_records.py b/litellm/llms/base_llm/files/batch_records.py new file mode 100644 index 00000000000..6bb98456e69 --- /dev/null +++ b/litellm/llms/base_llm/files/batch_records.py @@ -0,0 +1,54 @@ +from collections.abc import Iterable, Mapping +from functools import cache +from types import MappingProxyType +from typing import Final, cast, get_type_hints + +from litellm.types.llms.openai import ResponseInputParam, ResponsesAPIOptionalRequestParams + + +def _frozen_mapping(items: Iterable[tuple[str, object]]) -> Mapping[str, object]: + return MappingProxyType(dict(items)) + + +@cache +def _responses_request_keys() -> frozenset[str]: + return frozenset(get_type_hints(ResponsesAPIOptionalRequestParams)) + + +def responses_batch_body_to_chat_body( + openai_request_body: Mapping[str, object], + custom_llm_provider: str | None = None, +) -> dict[str, object]: # mutable-ok: provider transforms take the bridged chat body as a plain dict + """ + Rewrite the body of an OpenAI `/v1/responses` batch record as a Chat Completions body. + + Batch providers translate chat bodies into their own request shape, so a Responses + record goes through the same Responses-to-Chat bridge the real-time path uses for + providers without a native Responses API: `input`, `instructions`, `max_output_tokens` + and the tool params translate identically in batch and real time. Like real time, the + record's fields are forwarded as sent instead of validated against the SDK TypedDicts, + whose required keys (a function tool's `strict`, an image part's `detail`) clients omit. + """ + from litellm.responses.litellm_completion_transformation.transformation import ( + LiteLLMCompletionResponsesConfig, + ) + + responses_input: Final = openai_request_body.get("input") + if responses_input is None: + raise ValueError( + "Batch record for /v1/responses is missing required `input` field: " + f"model={openai_request_body.get('model', '')}" + ) + model: Final = openai_request_body.get("model") + chat_input: Final = cast(str | ResponseInputParam, responses_input) # cast-ok: forwarded as sent + responses_request: Final = cast( # cast-ok: client-supplied fields forwarded verbatim, as real time does + ResponsesAPIOptionalRequestParams, + _frozen_mapping((key, value) for key, value in openai_request_body.items() if key in _responses_request_keys()), + ) + return LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request( # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # transformer declares a bare dict return + model=model if isinstance(model, str) else "", + input=chat_input, + responses_api_request=responses_request, + custom_llm_provider=custom_llm_provider, + metadata=openai_request_body.get("metadata"), + ) diff --git a/litellm/llms/bedrock/files/transformation.py b/litellm/llms/bedrock/files/transformation.py index 43faa7d79ea..a79f4de1e3d 100644 --- a/litellm/llms/bedrock/files/transformation.py +++ b/litellm/llms/bedrock/files/transformation.py @@ -8,7 +8,6 @@ from collections.abc import Iterable, Mapping, MutableMapping, Sequence from contextlib import suppress from dataclasses import dataclass from datetime import datetime -from functools import cache from itertools import chain from types import MappingProxyType from typing import Any, Final, Literal, TypeAlias, TypedDict @@ -17,7 +16,7 @@ from urllib.parse import quote, unquote, urlencode import httpx from httpx import Headers, Response from openai.types.file_deleted import FileDeleted -from pydantic import BaseModel, ConfigDict, Field, TypeAdapter +from pydantic import BaseModel, ConfigDict, Field from typing_extensions import ReadOnly from litellm._logging import verbose_logger @@ -41,7 +40,9 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( extract_file_data, text_completion_prompt_to_messages, ) +from litellm.llms.base_llm.base_utils import map_developer_role_to_system_role from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.base_llm.files.batch_records import responses_batch_body_to_chat_body from litellm.llms.base_llm.files.transformation import ( BaseFilesConfig, LiteLLMLoggingObj, @@ -56,8 +57,6 @@ from litellm.types.llms.openai import ( OpenAICreateFileRequestOptionalParams, OpenAIFileObject, PathLike, - ResponseInputParam, - ResponsesAPIOptionalRequestParams, ) from litellm.types.utils import ExtractedFileData, LlmProviders, SpecialEnums from litellm.utils import get_llm_provider @@ -130,22 +129,6 @@ class _S3UploadResponse(TypedDict, total=False): ContentLength: ReadOnly[int] -# JSONL batch records are untyped json, so the `/v1/responses` fields are -# validated into their concrete Responses API types before being handed to the -# Responses-to-Chat bridge. Both adapters drop keys the Responses API doesn't -# define, which is what the bridge would ignore anyway. Built on first use -# rather than at import: `ResponseInputParam` is a deep union and only batch -# files carrying `/v1/responses` records need it. -@cache -def _responses_input_adapter() -> TypeAdapter[str | ResponseInputParam]: - return TypeAdapter(str | ResponseInputParam) - - -@cache -def _responses_request_adapter() -> TypeAdapter[ResponsesAPIOptionalRequestParams]: - return TypeAdapter(ResponsesAPIOptionalRequestParams) - - class _BedrockS3RequestParams(AwsAuthParams): """Typed view of the credential/region params the S3 GetObject path reads.""" @@ -859,33 +842,9 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): Delegates to the same Responses-to-Chat bridge the real-time path uses for providers without a native Responses API (which is every Bedrock model), so `input`, `instructions`, `max_output_tokens` and the tool - params translate identically in batch and real time. The bridge always - emits a `tools` key; an empty one is dropped rather than shipped as an - empty array inside `modelInput`. + params translate identically in batch and real time. """ - from litellm.responses.litellm_completion_transformation.transformation import ( - LiteLLMCompletionResponsesConfig, - ) - - responses_input: Final = openai_request_body.get("input") - if responses_input is None: - raise ValueError( - "Batch record for /v1/responses is missing required `input` field: " - f"model={openai_request_body.get('model', '')}" - ) - chat_body: Final[Mapping[str, object]] = ( - LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request( - model=openai_request_body.get("model", ""), - input=_responses_input_adapter().validate_python(responses_input), - responses_api_request=_responses_request_adapter().validate_python( - _frozen_mapping( - (key, value) for key, value in openai_request_body.items() if key not in ("model", "input") - ) - ), - metadata=openai_request_body.get("metadata"), - ) - ) - return _frozen_mapping((key, value) for key, value in chat_body.items() if key != "tools" or value) + return responses_batch_body_to_chat_body(openai_request_body) @staticmethod def _transform_batch_body_to_chat_body( @@ -922,7 +881,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): """ from litellm.types.utils import LlmProviders - messages: Final = openai_request_body.get("messages", []) + messages: Final = map_developer_role_to_system_role(openai_request_body.get("messages", [])) optional_params: Final = {k: v for k, v in openai_request_body.items() if k not in ["model", "messages"]} # --- Anthropic: use existing AmazonAnthropicClaudeConfig --- diff --git a/litellm/llms/vertex_ai/files/transformation.py b/litellm/llms/vertex_ai/files/transformation.py index 789b36ef3d0..dbb41b57348 100644 --- a/litellm/llms/vertex_ai/files/transformation.py +++ b/litellm/llms/vertex_ai/files/transformation.py @@ -35,7 +35,9 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( extract_file_data, extract_file_metadata, ) +from litellm.llms.base_llm.base_utils import map_developer_role_to_system_role from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.base_llm.files.batch_records import responses_batch_body_to_chat_body from litellm.llms.base_llm.files.transformation import ( BaseFilesConfig, BaseFileUploadStream, @@ -529,21 +531,30 @@ def is_passthrough_batch_upload(create_file_data: Mapping[str, object], litellm_ return create_file_data.get("purpose") == "batch" and litellm_params.get("passthrough") is True -def _is_embeddings_batch_entry(openai_entry: Mapping[str, object]) -> bool: +def _batch_entry_route_path(openai_entry: Mapping[str, object]) -> str: """ - Whether an OpenAI batch JSONL line targets the embeddings endpoint. + The route an OpenAI batch JSONL line targets, without query string or trailing slash. OpenAI puts the target route on each line's `url` (e.g. `/v1/embeddings`); Vertex has no equivalent per-line field, so the route decides which Vertex request shape the line has to be translated into. """ - url = openai_entry.get("url") + url: Final = openai_entry.get("url") if not isinstance(url, str): - return False - path = url.split("?")[0].rstrip("/") + return "" + return url.split("?")[0].rstrip("/") + + +def _is_embeddings_batch_entry(openai_entry: Mapping[str, object]) -> bool: + path: Final = _batch_entry_route_path(openai_entry) return path == "embeddings" or path.endswith("/embeddings") +def _is_responses_batch_entry(openai_entry: Mapping[str, object]) -> bool: + path: Final = _batch_entry_route_path(openai_entry) + return path == "responses" or path.endswith("/responses") + + def _openai_embedding_input_elements( embedding_input: GeminiEmbeddingInput, ) -> tuple[str | list[str], ...]: @@ -665,10 +676,15 @@ def _openai_batch_jsonl_entry_to_vertex_rows( return _openai_batch_jsonl_entry_to_vertex_embeddings_rows(openai_entry) openai_request_body: Final = openai_entry.get("body") or {} + chat_request_body: Final = ( + responses_batch_body_to_chat_body(openai_request_body, custom_llm_provider="vertex_ai") + if _is_responses_batch_entry(openai_entry) + else openai_request_body + ) vertex_request_body: Final = _transform_request_body( - messages=openai_request_body.get("messages", []), - model=openai_request_body.get("model", ""), - optional_params=map_openai_to_vertex_params(openai_request_body), + messages=map_developer_role_to_system_role(chat_request_body.get("messages", [])), + model=chat_request_body.get("model", ""), + optional_params=map_openai_to_vertex_params(chat_request_body), custom_llm_provider="vertex_ai", litellm_params={}, cached_content=None, diff --git a/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py b/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py index d0921e68424..7a7159a1624 100644 --- a/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py +++ b/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py @@ -1672,6 +1672,49 @@ class TestBedrockBatchNonChatEndpointRecords: assert "input" not in model_input assert "max_output_tokens" not in model_input + def test_anthropic_responses_record_accepts_a_function_tool_without_strict(self): + """Clients omit the SDK's required `strict`; the record is forwarded like real time, not validated.""" + parameters = {"type": "object", "properties": {"city": {"type": "string"}}} + model_input = self._transform( + { + "custom_id": "4a", + "method": "POST", + "url": "/v1/responses", + "body": { + "model": self.ANTHROPIC_MODEL, + "input": "Weather in Paris?", + "tools": [{"type": "function", "name": "get_weather", "parameters": parameters}], + }, + } + ) + + assert model_input["messages"][0]["content"] == [{"type": "text", "text": "Weather in Paris?"}] + tool = model_input["tools"][0] + function = tool.get("function", tool) + assert (function["name"], function.get("parameters", function.get("input_schema"))) == ("get_weather", parameters) + + @pytest.mark.parametrize( + ("url", "body"), + [ + ( + "/v1/responses", + {"input": [{"role": "developer", "content": "be terse"}, {"role": "user", "content": "ping"}]}, + ), + ( + "/v1/chat/completions", + {"messages": [{"role": "developer", "content": "be terse"}, {"role": "user", "content": "ping"}]}, + ), + ], + ids=["responses", "chat"], + ) + def test_anthropic_developer_role_becomes_the_system_prompt_like_real_time(self, url, body): + model_input = self._transform( + {"custom_id": "4c", "method": "POST", "url": url, "body": {"model": self.ANTHROPIC_MODEL, **body}} + ) + + assert model_input["system"] == [{"type": "text", "text": "be terse"}] + assert [message["role"] for message in model_input["messages"]] == ["user"] + def test_responses_record_keeps_metadata(self): """`metadata` reaches the bridge, which reads it as its own kwarg.""" model_input = self._transform( diff --git a/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_transformation.py b/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_transformation.py index 48464e79876..7434eae72a4 100644 --- a/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_transformation.py +++ b/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_transformation.py @@ -5,6 +5,7 @@ Includes tests for Vertex AI batch output transformation to OpenAI format. import json import urllib.parse +from collections.abc import Mapping from types import MappingProxyType from urllib.parse import parse_qs, urlparse @@ -1447,6 +1448,158 @@ class TestVertexEmbeddingsBatchInputTranslation: assert "content" in embeddings_row["request"] +def _responses_entry( + body: Mapping[str, object] | None = None, + custom_id: str = "resp-1", + url: str = "/v1/responses", +) -> dict[str, object]: + return { + "custom_id": custom_id, + "method": "POST", + "url": url, + "body": body + if body is not None + else {"model": "gemini-2.5-flash", "input": "What was the top headline in world news yesterday?"}, + } + + +class TestVertexResponsesBatchInputTranslation: + """ + /v1/responses batch lines carry `input`, not `messages`, so they go through the + Responses-to-Chat bridge before the Gemini translation instead of uploading as an + empty text part. + """ + + def test_string_input_becomes_the_user_prompt(self): + (row,) = _wrap_entries([_responses_entry()]) + + assert row["request"]["contents"] == [ + {"role": "user", "parts": [{"text": "What was the top headline in world news yesterday?"}]} + ] + assert row["request"]["labels"]["litellm_custom_id"] == "resp-1" + + def test_instructions_and_input_items_map_like_real_time(self): + (row,) = _wrap_entries( + [ + _responses_entry( + body={ + "model": "gemini-2.5-flash", + "instructions": "be terse", + "input": [ + {"role": "user", "content": "what is 2+2?"}, + {"role": "assistant", "content": "4"}, + {"role": "user", "content": "and 3+3?"}, + ], + "max_output_tokens": 32, + "temperature": 0.2, + } + ) + ] + ) + + request = row["request"] + assert request["system_instruction"] == {"parts": [{"text": "be terse"}]} + assert [content["role"] for content in request["contents"]] == ["user", "model", "user"] + assert request["contents"][-1]["parts"] == [{"text": "and 3+3?"}] + assert request["generationConfig"]["max_output_tokens"] == 32 + assert request["generationConfig"]["temperature"] == 0.2 + + def test_web_search_tool_keeps_the_prompt(self): + (row,) = _wrap_entries( + [ + _responses_entry( + body={ + "model": "gemini-2.5-flash", + "input": "What was the top headline in world news yesterday?", + "tools": [{"type": "web_search"}], + } + ) + ] + ) + + assert row["request"]["contents"] == [ + {"role": "user", "parts": [{"text": "What was the top headline in world news yesterday?"}]} + ] + assert row["request"]["tools"] + + def test_sdk_optional_keys_are_not_required_like_real_time(self): + (row,) = _wrap_entries( + [ + _responses_entry( + body={ + "model": "gemini-2.5-flash", + "input": [ + { + "role": "user", + "content": [ + {"type": "input_text", "text": "Weather in the pictured city?"}, + {"type": "input_image", "image_url": "https://example.com/paris.png"}, + ], + } + ], + "tools": [ + { + "type": "function", + "name": "get_weather", + "parameters": {"type": "object", "properties": {"city": {"type": "string"}}}, + } + ], + } + ) + ] + ) + + request = row["request"] + assert request["contents"][0]["parts"] == [ + {"text": "Weather in the pictured city?"}, + {"file_data": {"mime_type": "image/png", "file_uri": "https://example.com/paris.png"}}, + ] + assert request["tools"][0]["function_declarations"][0]["name"] == "get_weather" + + @pytest.mark.parametrize( + "url", + ["/v1/responses", "/v1/responses/", "/v1/responses?beta=1", "responses", "https://api.openai.com/v1/responses"], + ) + def test_route_spellings_are_all_responses(self, url): + (row,) = _wrap_entries([_responses_entry(url=url)]) + + assert row["request"]["contents"][0]["parts"] == [ + {"text": "What was the top headline in world news yesterday?"} + ] + + def test_missing_input_fails_the_upload(self): + with pytest.raises(ValueError, match="missing required `input` field"): + _wrap_entries([_responses_entry(body={"model": "gemini-2.5-flash"})]) + + @pytest.mark.parametrize( + "entry", + [ + _responses_entry( + body={ + "model": "gemini-2.5-flash", + "input": [{"role": "developer", "content": "be terse"}, {"role": "user", "content": "ping"}], + } + ), + { + "custom_id": "chat-1", + "method": "POST", + "url": "/v1/chat/completions", + "body": { + "model": "gemini-2.5-flash", + "messages": [{"role": "developer", "content": "be terse"}, {"role": "user", "content": "ping"}], + }, + }, + ], + ids=["responses", "chat"], + ) + def test_developer_role_becomes_the_system_instruction_like_real_time(self, entry): + (row,) = _wrap_entries([entry]) + + request = row["request"] + assert request["system_instruction"] == {"parts": [{"text": "be terse"}]} + assert request["contents"] == [{"role": "user", "parts": [{"text": "ping"}]}] + + class TestVertexEmbeddingsBatchOutputTranslation: """Vertex Gemini Embedding batch output rows must come back as OpenAI batch rows.""" From 1dc3b62dbc160c5f97bd8420829dcdd629e1e0c7 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 16:26:44 -0700 Subject: [PATCH 012/218] fix(cost-map): add vertex priority prices for gemini-3-pro-image-preview and batch price for gemini-embedding-001 (#43069) Co-authored-by: kerry Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/model_prices_and_context_window_backup.json | 13 +++++++++++++ model_prices_and_context_window.json | 13 +++++++++++++ 2 files changed, 26 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 9bec83f6b08..dc979cbd651 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -26312,10 +26312,14 @@ "gemini-3-pro-image-preview": { "input_cost_per_image": 0.0011, "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_priority": 3.6e-07, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07, "cache_read_input_token_cost_batches": 1e-07, "input_cost_per_token": 2e-06, + "input_cost_per_token_priority": 3.6e-06, "input_cost_per_token_above_200k_tokens": 4e-06, + "input_cost_per_token_above_200k_tokens_priority": 7.2e-06, "input_cost_per_token_batches": 1e-06, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 65536, @@ -26325,7 +26329,9 @@ "output_cost_per_image": 0.134, "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, + "output_cost_per_token_priority": 2.16e-05, "output_cost_per_token_above_200k_tokens": 1.8e-05, + "output_cost_per_token_above_200k_tokens_priority": 3.24e-05, "output_cost_per_token_batches": 6e-06, "source": "https://ai.google.dev/gemini-api/docs/pricing", "supported_endpoints": [ @@ -27848,6 +27854,7 @@ "gemini-embedding-001": { "deprecation_date": "2028-05-20", "input_cost_per_token": 1.5e-07, + "input_cost_per_token_batches": 1.2e-07, "litellm_provider": "vertex_ai-embedding-models", "max_input_tokens": 2048, "max_tokens": 2048, @@ -49492,10 +49499,14 @@ "vertex_ai/gemini-3-pro-image-preview": { "input_cost_per_image": 0.0011, "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_priority": 3.6e-07, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07, "cache_read_input_token_cost_batches": 1e-07, "input_cost_per_token": 2e-06, + "input_cost_per_token_priority": 3.6e-06, "input_cost_per_token_above_200k_tokens": 4e-06, + "input_cost_per_token_above_200k_tokens_priority": 7.2e-06, "input_cost_per_token_batches": 1e-06, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 65536, @@ -49505,7 +49516,9 @@ "output_cost_per_image": 0.134, "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, + "output_cost_per_token_priority": 2.16e-05, "output_cost_per_token_above_200k_tokens": 1.8e-05, + "output_cost_per_token_above_200k_tokens_priority": 3.24e-05, "output_cost_per_token_batches": 6e-06, "supports_reasoning": false, "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 9bec83f6b08..dc979cbd651 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -26312,10 +26312,14 @@ "gemini-3-pro-image-preview": { "input_cost_per_image": 0.0011, "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_priority": 3.6e-07, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07, "cache_read_input_token_cost_batches": 1e-07, "input_cost_per_token": 2e-06, + "input_cost_per_token_priority": 3.6e-06, "input_cost_per_token_above_200k_tokens": 4e-06, + "input_cost_per_token_above_200k_tokens_priority": 7.2e-06, "input_cost_per_token_batches": 1e-06, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 65536, @@ -26325,7 +26329,9 @@ "output_cost_per_image": 0.134, "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, + "output_cost_per_token_priority": 2.16e-05, "output_cost_per_token_above_200k_tokens": 1.8e-05, + "output_cost_per_token_above_200k_tokens_priority": 3.24e-05, "output_cost_per_token_batches": 6e-06, "source": "https://ai.google.dev/gemini-api/docs/pricing", "supported_endpoints": [ @@ -27848,6 +27854,7 @@ "gemini-embedding-001": { "deprecation_date": "2028-05-20", "input_cost_per_token": 1.5e-07, + "input_cost_per_token_batches": 1.2e-07, "litellm_provider": "vertex_ai-embedding-models", "max_input_tokens": 2048, "max_tokens": 2048, @@ -49492,10 +49499,14 @@ "vertex_ai/gemini-3-pro-image-preview": { "input_cost_per_image": 0.0011, "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_priority": 3.6e-07, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07, "cache_read_input_token_cost_batches": 1e-07, "input_cost_per_token": 2e-06, + "input_cost_per_token_priority": 3.6e-06, "input_cost_per_token_above_200k_tokens": 4e-06, + "input_cost_per_token_above_200k_tokens_priority": 7.2e-06, "input_cost_per_token_batches": 1e-06, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 65536, @@ -49505,7 +49516,9 @@ "output_cost_per_image": 0.134, "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, + "output_cost_per_token_priority": 2.16e-05, "output_cost_per_token_above_200k_tokens": 1.8e-05, + "output_cost_per_token_above_200k_tokens_priority": 3.24e-05, "output_cost_per_token_batches": 6e-06, "supports_reasoning": false, "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" From 4584958574cd22b6f7f9b65aa6c92c58cdf940b6 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 18:26:50 -0500 Subject: [PATCH 013/218] feat(agents): add optional per-agent kill switch webhook (#42841) --- .../migration.sql | 2 + .../litellm_proxy_extras/schema.prisma | 1 + litellm/constants.py | 2 + litellm/proxy/_lazy_openapi_snapshot.json | 293 ++++++++++++++++++ litellm/proxy/_types.py | 4 +- .../proxy/agent_endpoints/agent_registry.py | 50 ++- litellm/proxy/agent_endpoints/endpoints.py | 90 +++++- litellm/proxy/agent_endpoints/kill_switch.py | 239 ++++++++++++++ litellm/proxy/schema.prisma | 1 + litellm/types/agents.py | 76 ++++- litellm/types/llms/custom_http.py | 1 + schema.prisma | 1 + .../agent_endpoints/test_agent_registry.py | 162 +++++++++- .../proxy/agent_endpoints/test_endpoints.py | 223 ++++++++++++- .../proxy/agent_endpoints/test_kill_switch.py | 248 +++++++++++++++ .../proxy/auth/test_route_checks.py | 1 + .../agents/_components/AgentFormKit.tsx | 5 + .../AgentKillSwitchDangerZone.test.tsx | 124 ++++++++ .../_components/AgentKillSwitchDangerZone.tsx | 147 +++++++++ .../_components/KillSwitchFormFields.tsx | 223 +++++++++++++ .../agents/_components/agent_config.ts | 23 ++ .../agents/_components/agent_form_fields.tsx | 9 + .../agents/_components/agent_info.test.tsx | 28 ++ .../agents/_components/agent_info.tsx | 11 +- .../_components/agent_type_utils.test.ts | 26 ++ .../agents/_components/agent_type_utils.ts | 2 + .../dynamic_agent_form_fields.test.ts | 65 ++++ .../_components/dynamic_agent_form_fields.tsx | 21 +- .../_components/kill_switch_config.test.ts | 118 +++++++ .../agents/_components/kill_switch_config.ts | 124 ++++++++ .../src/components/agents/types.ts | 3 + .../src/components/networking.tsx | 11 + .../AuditLogDrawer/AuditLogDrawer.tsx | 1 + .../components/view_logs/AuditLogsTable.tsx | 2 + .../view_logs/AuditLogsTableColumns.tsx | 9 +- .../src/contexts/PluginModeContext.tsx | 10 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 149 +++++++++ 37 files changed, 2485 insertions(+), 20 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260923000000_add_agent_kill_switch/migration.sql create mode 100644 litellm/proxy/agent_endpoints/kill_switch.py create mode 100644 tests/test_litellm/proxy/agent_endpoints/test_kill_switch.py create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentKillSwitchDangerZone.test.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentKillSwitchDangerZone.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/agents/_components/KillSwitchFormFields.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/agents/_components/dynamic_agent_form_fields.test.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/agents/_components/kill_switch_config.test.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/agents/_components/kill_switch_config.ts diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260923000000_add_agent_kill_switch/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260923000000_add_agent_kill_switch/migration.sql new file mode 100644 index 00000000000..dd21ed644eb --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260923000000_add_agent_kill_switch/migration.sql @@ -0,0 +1,2 @@ +-- AlterTable +ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "kill_switch" JSONB; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 85996430bc5..69c63d9ecd6 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -72,6 +72,7 @@ model LiteLLM_AgentsTable { agent_card_params Json static_headers Json? @default("{}") extra_headers String[] @default([]) + kill_switch Json? agent_access_groups String[] @default([]) access_group_ids String[] @default([]) object_permission_id String? diff --git a/litellm/constants.py b/litellm/constants.py index 7b40f432446..807694c2f8a 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -557,6 +557,8 @@ SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS: Final[float] = float( request_timeout: float = float(os.getenv("REQUEST_TIMEOUT", str(int(DEFAULT_REQUEST_TIMEOUT_SECONDS)))) request_timeout_explicitly_set: bool = "REQUEST_TIMEOUT" in os.environ DEFAULT_A2A_AGENT_TIMEOUT: Final[float] = float(os.getenv("DEFAULT_A2A_AGENT_TIMEOUT", 6000)) # 10 minutes +AGENT_KILL_SWITCH_TIMEOUT_SECONDS: Final = 10.0 +AGENT_KILL_SWITCH_RESPONSE_BODY_MAX_CHARS: Final = 2000 # Patterns that indicate a localhost/internal URL in A2A agent cards that should be # replaced with the original base_url. This is a common misconfiguration where # developers deploy agents with development URLs in their agent cards. diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 0b43c3864ab..3d44315341b 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -2392,6 +2392,16 @@ ], "title": "Extra Headers" }, + "kill_switch": { + "anyOf": [ + { + "$ref": "#/components/schemas/AgentKillSwitchConfig" + }, + { + "type": "null" + } + ] + }, "litellm_params": { "additionalProperties": true, "title": "Litellm Params", @@ -2561,6 +2571,221 @@ "title": "AgentKeySummary", "type": "object" }, + "AgentKillSwitchApiKeyAuth": { + "additionalProperties": false, + "properties": { + "api_key": { + "title": "Api Key", + "type": "string" + }, + "header_name": { + "default": "x-api-key", + "title": "Header Name", + "type": "string" + }, + "type": { + "const": "api_key", + "title": "Type", + "type": "string" + } + }, + "required": [ + "type", + "api_key" + ], + "title": "AgentKillSwitchApiKeyAuth", + "type": "object" + }, + "AgentKillSwitchBasicAuth": { + "additionalProperties": false, + "properties": { + "password": { + "title": "Password", + "type": "string" + }, + "type": { + "const": "basic", + "title": "Type", + "type": "string" + }, + "username": { + "title": "Username", + "type": "string" + } + }, + "required": [ + "type", + "username", + "password" + ], + "title": "AgentKillSwitchBasicAuth", + "type": "object" + }, + "AgentKillSwitchBearerAuth": { + "additionalProperties": false, + "properties": { + "token": { + "title": "Token", + "type": "string" + }, + "type": { + "const": "bearer", + "title": "Type", + "type": "string" + } + }, + "required": [ + "type", + "token" + ], + "title": "AgentKillSwitchBearerAuth", + "type": "object" + }, + "AgentKillSwitchConfig": { + "additionalProperties": false, + "description": "Webhook an admin fires to shut an agent down out of band. LiteLLM only\nmakes the call; whatever the endpoint does with it is the agent's business.", + "properties": { + "auth": { + "anyOf": [ + { + "discriminator": { + "mapping": { + "api_key": "#/components/schemas/AgentKillSwitchApiKeyAuth", + "basic": "#/components/schemas/AgentKillSwitchBasicAuth", + "bearer": "#/components/schemas/AgentKillSwitchBearerAuth" + }, + "propertyName": "type" + }, + "oneOf": [ + { + "$ref": "#/components/schemas/AgentKillSwitchBearerAuth" + }, + { + "$ref": "#/components/schemas/AgentKillSwitchApiKeyAuth" + }, + { + "$ref": "#/components/schemas/AgentKillSwitchBasicAuth" + } + ] + }, + { + "type": "null" + } + ], + "title": "Auth" + }, + "body": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "title": "Body" + }, + "headers": { + "additionalProperties": { + "type": "string" + }, + "title": "Headers", + "type": "object" + }, + "method": { + "default": "POST", + "enum": [ + "POST", + "PUT", + "PATCH", + "DELETE", + "GET" + ], + "title": "Method", + "type": "string" + }, + "query_params": { + "additionalProperties": { + "type": "string" + }, + "title": "Query Params", + "type": "object" + }, + "url": { + "title": "Url", + "type": "string" + } + }, + "required": [ + "url" + ], + "title": "AgentKillSwitchConfig", + "type": "object" + }, + "AgentKillSwitchResult": { + "properties": { + "agent_id": { + "title": "Agent Id", + "type": "string" + }, + "error": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Error" + }, + "method": { + "enum": [ + "POST", + "PUT", + "PATCH", + "DELETE", + "GET" + ], + "title": "Method", + "type": "string" + }, + "response_body": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Response Body" + }, + "status_code": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "Status Code" + }, + "url": { + "title": "Url", + "type": "string" + } + }, + "required": [ + "agent_id", + "url", + "method" + ], + "title": "AgentKillSwitchResult", + "type": "object" + }, "AgentMakePublicResponse": { "properties": { "message": { @@ -2775,6 +3000,16 @@ ], "title": "Keys" }, + "kill_switch": { + "anyOf": [ + { + "$ref": "#/components/schemas/AgentKillSwitchConfig" + }, + { + "type": "null" + } + ] + }, "litellm_params": { "anyOf": [ { @@ -3569,6 +3804,16 @@ ], "title": "Extra Headers" }, + "kill_switch": { + "anyOf": [ + { + "$ref": "#/components/schemas/AgentKillSwitchConfig" + }, + { + "type": "null" + } + ] + }, "litellm_params": { "additionalProperties": true, "title": "Litellm Params", @@ -4331,6 +4576,54 @@ ] } }, + "/v1/agents/{agent_id}/kill_switch": { + "post": { + "description": "Fire the agent's configured kill switch webhook. Proxy admin only.\n\nLiteLLM only makes the configured HTTP call and reports what came back; it\ndoes not change the agent's state in LiteLLM. Returns 200 when the webhook\nanswered 2xx, 502 with the same result body otherwise. Every attempt is\nwritten to the audit log as a `kill_switch_fired` row against the agent.\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000/kill_switch\" \\\n -H \"Authorization: Bearer \"\n```", + "operationId": "trigger_agent_kill_switch_v1_agents__agent_id__kill_switch_post", + "parameters": [ + { + "in": "path", + "name": "agent_id", + "required": true, + "schema": { + "title": "Agent Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/AgentKillSwitchResult" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Trigger Agent Kill Switch", + "tags": [ + "agents" + ] + } + }, "/v1/agents/{agent_id}/make_public": { "post": { "description": "Make an agent publicly discoverable\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000/make_public\" \\\n -H \"Authorization: Bearer \" \\\n -H \"Content-Type: application/json\"\n```\n\nExample Response:\n```json\n{\n \"agent_id\": \"123e4567-e89b-12d3-a456-426614174000\",\n \"agent_name\": \"my-custom-agent\",\n \"litellm_params\": {\n \"make_public\": true\n },\n \"agent_card_params\": {...},\n \"created_at\": \"2025-11-15T10:30:00Z\",\n \"updated_at\": \"2025-11-15T10:35:00Z\",\n \"created_by\": \"user123\",\n \"updated_by\": \"user123\"\n}\n```", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index b6de36f8423..12b4d4b2412 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -238,6 +238,7 @@ class LitellmTableNames(str, enum.Enum): CONFIG_TABLE_NAME = "LiteLLM_Config" SSO_CONFIG_TABLE_NAME = "LiteLLM_SSOConfig" UI_SETTINGS_TABLE_NAME = "LiteLLM_UISettings" + AGENT_TABLE_NAME = "LiteLLM_AgentsTable" class Litellm_EntityType(enum.Enum): @@ -578,6 +579,7 @@ class LiteLLMRoutes(enum.Enum): "/v1/agents/{agent_id}", "/v1/agents/make_public", "/v1/agents/{agent_id}/make_public", + "/v1/agents/{agent_id}/kill_switch", ) # Backwards-compat union โ€” virtual keys may be configured with @@ -3688,7 +3690,7 @@ from litellm.models.spend_logs import ( # noqa: E402 ) from litellm.models.tag import LiteLLM_TagTable as LiteLLM_TagTable # noqa: E402 -AUDIT_ACTIONS = Literal["created", "updated", "deleted", "blocked", "unblocked", "rotated"] +AUDIT_ACTIONS = Literal["created", "updated", "deleted", "blocked", "unblocked", "rotated", "kill_switch_fired"] class LiteLLM_AuditLogs(LiteLLMPydanticObjectBase): diff --git a/litellm/proxy/agent_endpoints/agent_registry.py b/litellm/proxy/agent_endpoints/agent_registry.py index 3d56c2b5326..3e775d7648e 100644 --- a/litellm/proxy/agent_endpoints/agent_registry.py +++ b/litellm/proxy/agent_endpoints/agent_registry.py @@ -13,13 +13,14 @@ import litellm from litellm.constants import REDACTED_BY_LITELM_STRING from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker +from litellm.proxy.agent_endpoints.kill_switch import restore_kill_switch from litellm.proxy.management_helpers.object_permission_utils import ( handle_update_object_permission_common, ) from litellm.proxy.utils import PrismaClient from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import AgentsRepository, ObjectPermissionRepository -from litellm.types.agents import AgentConfig, AgentResponse, PatchAgentRequest +from litellm.types.agents import AgentConfig, AgentKillSwitchConfig, AgentResponse, PatchAgentRequest if TYPE_CHECKING: from prisma import models as prisma_models @@ -31,6 +32,10 @@ class AgentObjectPermissionRecord(Protocol): def dict(self) -> dict[str, object]: ... +class AgentIdWhere(TypedDict): + agent_id: ReadOnly[str] + + class AgentRecordDump(TypedDict): agent_id: str agent_name: str @@ -38,6 +43,7 @@ class AgentRecordDump(TypedDict): agent_card_params: dict[str, object] static_headers: dict[str, str] | None extra_headers: list[str] | None + kill_switch: ReadOnly[AgentKillSwitchConfig | None] access_group_ids: ReadOnly[Sequence[str] | None] object_permission: dict[str, object] | None spend: float @@ -70,6 +76,9 @@ class AgentRecord(Protocol): @property def access_group_ids(self) -> Sequence[str] | None: ... + @property + def kill_switch(self) -> Mapping[str, object] | None: ... + @property def spend(self) -> float: ... @@ -211,6 +220,29 @@ def parse_agent_litellm_params(value: object) -> Mapping[str, object]: return _EMPTY_LITELLM_PARAMS +_KILL_SWITCH_ADAPTER: Final[TypeAdapter[AgentKillSwitchConfig | None]] = TypeAdapter(AgentKillSwitchConfig | None) + + +def parse_agent_kill_switch(value: object) -> AgentKillSwitchConfig | None: + if value is None: + return None + try: + if isinstance(value, str): + return _KILL_SWITCH_ADAPTER.validate_json(value) + return _KILL_SWITCH_ADAPTER.validate_python(value) + except ValidationError: + return None + + +def serialize_agent_kill_switch(incoming: object, existing: object) -> str: + """prisma-client-py drops ``None`` from update data, so a cleared kill switch is stored as the JSON literal + ``null`` (read back as ``None``), the same convention ``memory_endpoints`` uses for ``Json?`` columns.""" + restored: Final = restore_kill_switch( + _KILL_SWITCH_ADAPTER.validate_python(incoming), parse_agent_kill_switch(existing) + ) + return safe_dumps(restored.model_dump() if restored is not None else None) + + _MISSING_AGENT_PARAM: Final = object() _RESTORE_AGENT_PARAMS_MAX_DEPTH: Final = 10 @@ -293,6 +325,12 @@ def _patched_access_group_ids(agent: PatchAgentRequest) -> Mapping[str, object]: return MappingProxyType({"access_group_ids": tuple(dict.fromkeys(agent.get("access_group_ids") or ()))}) +def _patched_kill_switch(agent: PatchAgentRequest, existing: object) -> Mapping[str, object]: + if "kill_switch" not in agent: + return MappingProxyType({}) + return MappingProxyType({"kill_switch": serialize_agent_kill_switch(agent.get("kill_switch"), existing)}) + + def _restore_redacted_litellm_params( incoming: Mapping[str, object], existing: Mapping[str, object], @@ -531,6 +569,7 @@ class AgentRegistry: "agent_name": agent_name, "litellm_params": litellm_params, "agent_card_params": agent_card_params, + "kill_switch": serialize_agent_kill_switch(agent.get("kill_switch"), None), "created_by": created_by, "updated_by": created_by, "created_at": datetime.now(timezone.utc), @@ -613,7 +652,10 @@ class AgentRegistry: existing_agent: Final[Mapping[str, object]] = dict(existing_record) augment_agent: Final = {**existing_agent, **agent} - update_data: Final[dict[str, object]] = {**_patched_access_group_ids(agent)} + update_data: Final[dict[str, object]] = { + **_patched_access_group_ids(agent), + **_patched_kill_switch(agent, existing_agent.get("kill_switch")), + } if augment_agent.get("agent_name"): update_data["agent_name"] = augment_agent.get("agent_name") if "litellm_params" in agent: @@ -716,6 +758,9 @@ class AgentRegistry: ) extra_headers_val_u: Final = agent.get("extra_headers") or [] access_group_ids_val_u: Final = tuple(dict.fromkeys(agent.get("access_group_ids") or ())) + kill_switch_val_u: Final = serialize_agent_kill_switch( + agent.get("kill_switch"), existing_row.kill_switch if existing_row is not None else None + ) update_data: Final[dict[str, object]] = { "agent_name": agent_name, @@ -723,6 +768,7 @@ class AgentRegistry: "agent_card_params": agent_card_params, "static_headers": static_headers_val_u, "extra_headers": extra_headers_val_u, + "kill_switch": kill_switch_val_u, "access_group_ids": access_group_ids_val_u, "updated_by": updated_by, "updated_at": datetime.now(timezone.utc), diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index aa8979a73c6..28c82a715e0 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -33,6 +33,8 @@ from litellm.proxy.a2a.agent_card import ( normalize_protocol_version, ) from litellm.proxy.agent_endpoints.agent_registry import ( + AgentIdWhere, + parse_agent_kill_switch, parse_agent_litellm_params, redact_sensitive_agent_litellm_params, ) @@ -45,6 +47,15 @@ from litellm.proxy.agent_endpoints.agent_search import ( search_agents, ) from litellm.proxy.agent_endpoints.auth.agent_permission_handler import accessible_agents +from litellm.proxy.agent_endpoints.kill_switch import ( + KillSwitchAuditLogWriter, + KillSwitchHttpClient, + build_kill_switch_audit_log, + default_kill_switch_audit_log_writer, + default_kill_switch_http_client, + fire_kill_switch, + redact_kill_switch, +) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity @@ -53,6 +64,8 @@ from litellm.types.agents import ( AgentCard, AgentConfig, AgentKeySummary, + AgentKillSwitchConfig, + AgentKillSwitchResult, AgentMakePublicResponse, AgentResponse, MakeAgentsPublicRequest, @@ -160,9 +173,10 @@ def _redact_sensitive_agent_fields( ) -> list[AgentResponse]: """ Return copies of the given agents with credential-bearing litellm_params - values replaced by a fixed marker (never returned to ANY caller, - admin included) and, for non-admin callers, virtual-key and header - fields stripped entirely. The original objects are not modified. + values and kill-switch auth secrets replaced by a fixed marker (never + returned to ANY caller, admin included) and, for non-admin callers, + virtual-key, header and kill-switch fields stripped entirely. The original + objects are not modified. """ redacted: Final[list[AgentResponse]] = [] for agent in agents: @@ -171,8 +185,10 @@ def _redact_sensitive_agent_fields( copy.static_headers = None copy.extra_headers = None copy.keys = None + copy.kill_switch = None if copy.litellm_params: copy.litellm_params = _redact_agent_litellm_params_dict(copy.litellm_params) + copy.kill_switch = redact_kill_switch(copy.kill_switch) redacted.append(copy) return redacted @@ -872,6 +888,74 @@ async def delete_agent( raise HTTPException(status_code=500, detail=str(e)) +@router.post( + "/v1/agents/{agent_id}/kill_switch", + tags=["[beta] A2A Agents"], # mutable-ok: fastapi types tags as list[str | Enum] + dependencies=(Depends(user_api_key_auth),), + response_model=AgentKillSwitchResult, +) +async def trigger_agent_kill_switch( + agent_id: str, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + http_client: Annotated[KillSwitchHttpClient, Depends(default_kill_switch_http_client)], + audit_log_writer: Annotated[KillSwitchAuditLogWriter, Depends(default_kill_switch_audit_log_writer)], +): + """ + Fire the agent's configured kill switch webhook. Proxy admin only. + + LiteLLM only makes the configured HTTP call and reports what came back; it + does not change the agent's state in LiteLLM. Returns 200 when the webhook + answered 2xx, 502 with the same result body otherwise. Every attempt is + written to the audit log as a `kill_switch_fired` row against the agent. + + Example Request: + ```bash + curl -X POST "http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000/kill_switch" \\ + -H "Authorization: Bearer " + ``` + """ + from litellm.proxy.proxy_server import litellm_proxy_admin_name + + await check_feature_access_for_user(user_api_key_dict, "agents") + _check_agent_management_permission(user_api_key_dict) + + resolved: Final = await _resolve_agent_kill_switch(agent_id) + if resolved is None: + raise HTTPException(status_code=404, detail=f"Agent with ID {agent_id} not found") + resolved_agent_id, config = resolved + if config is None: + raise HTTPException(status_code=400, detail=f"Agent with ID {agent_id} has no kill_switch configured") + + result: Final = await fire_kill_switch(agent_id=resolved_agent_id, config=config, http_client=http_client) + await audit_log_writer( + build_kill_switch_audit_log( + result=result, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ) + ) + if not result.succeeded: + raise HTTPException(status_code=502, detail=result.model_dump()) + return result + + +async def _resolve_agent_kill_switch(agent_id: str) -> tuple[str, AgentKillSwitchConfig | None] | None: + """The DB row wins over this replica's in-memory registry so a trigger never fires a webhook another + replica has since changed; config.yaml agents have no row and fall back to the registry.""" + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is not None: + where: Final[AgentIdWhere] = {"agent_id": agent_id} + row: Final = await agents_table(prisma_client).find_unique(where=where) + if row is not None: + return row.agent_id, parse_agent_kill_switch(row.kill_switch) + + agent: Final = AGENT_REGISTRY.get_agent_by_id(agent_id=agent_id) + if agent is None: + return None + return agent.agent_id, agent.kill_switch + + @router.post( "/v1/agents/{agent_id}/make_public", tags=["[beta] A2A Agents"], diff --git a/litellm/proxy/agent_endpoints/kill_switch.py b/litellm/proxy/agent_endpoints/kill_switch.py new file mode 100644 index 00000000000..8b3f64e74ee --- /dev/null +++ b/litellm/proxy/agent_endpoints/kill_switch.py @@ -0,0 +1,239 @@ +from base64 import b64encode +from collections.abc import AsyncIterator, Awaitable, Callable, Mapping +from dataclasses import dataclass +from datetime import datetime, timezone +from types import MappingProxyType +from typing import Final, Protocol, TypeAlias + +import httpx +from typing_extensions import assert_never + +from litellm._logging import verbose_proxy_logger +from litellm._uuid import uuid +from litellm.constants import ( + AGENT_KILL_SWITCH_RESPONSE_BODY_MAX_CHARS, + AGENT_KILL_SWITCH_TIMEOUT_SECONDS, + REDACTED_BY_LITELM_STRING, +) +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # its params arg is a bare dict in http_handler +) +from litellm.proxy._types import LiteLLM_AuditLogs, LitellmTableNames, UserAPIKeyAuth +from litellm.proxy.management_helpers.audit_logs import create_audit_log_for_update, get_audit_log_changed_by +from litellm.types.agents import ( + AgentKillSwitchApiKeyAuth, + AgentKillSwitchAuth, + AgentKillSwitchBasicAuth, + AgentKillSwitchBearerAuth, + AgentKillSwitchConfig, + AgentKillSwitchResult, +) +from litellm.types.llms.custom_http import httpxSpecialProvider + + +def _with_auth(config: AgentKillSwitchConfig, auth: AgentKillSwitchAuth) -> AgentKillSwitchConfig: + return AgentKillSwitchConfig( + url=config.url, + method=config.method, + headers=config.headers, + query_params=config.query_params, + body=config.body, + auth=auth, + ) + + +def redact_kill_switch(config: AgentKillSwitchConfig | None) -> AgentKillSwitchConfig | None: + if config is None or config.auth is None: + return config + return _with_auth(config, _redact_auth(config.auth)) + + +def _redact_auth(auth: AgentKillSwitchAuth) -> AgentKillSwitchAuth: + match auth: + case AgentKillSwitchBearerAuth(): + return AgentKillSwitchBearerAuth(type="bearer", token=REDACTED_BY_LITELM_STRING) + case AgentKillSwitchApiKeyAuth(): + return AgentKillSwitchApiKeyAuth( + type="api_key", header_name=auth.header_name, api_key=REDACTED_BY_LITELM_STRING + ) + case AgentKillSwitchBasicAuth(): + return AgentKillSwitchBasicAuth(type="basic", username=auth.username, password=REDACTED_BY_LITELM_STRING) + case _: + assert_never(auth) + + +def restore_kill_switch( + incoming: AgentKillSwitchConfig | None, + existing: AgentKillSwitchConfig | None, +) -> AgentKillSwitchConfig | None: + """Put the stored secret back behind an auth field echoed as the redaction + marker; a marker with no stored secret of the same auth type becomes "".""" + if incoming is None or incoming.auth is None: + return incoming + existing_auth: Final = existing.auth if existing is not None else None + return _with_auth(incoming, _restore_auth(incoming.auth, existing_auth)) + + +def _restore_secret(incoming_value: str, existing_value: str | None) -> str: + if incoming_value != REDACTED_BY_LITELM_STRING: + return incoming_value + return existing_value if existing_value is not None else "" + + +def _restore_auth(incoming: AgentKillSwitchAuth, existing: AgentKillSwitchAuth | None) -> AgentKillSwitchAuth: + match incoming: + case AgentKillSwitchBearerAuth(): + stored_token: Final = existing.token if isinstance(existing, AgentKillSwitchBearerAuth) else None + return AgentKillSwitchBearerAuth(type="bearer", token=_restore_secret(incoming.token, stored_token)) + case AgentKillSwitchApiKeyAuth(): + stored_key: Final = existing.api_key if isinstance(existing, AgentKillSwitchApiKeyAuth) else None + return AgentKillSwitchApiKeyAuth( + type="api_key", + header_name=incoming.header_name, + api_key=_restore_secret(incoming.api_key, stored_key), + ) + case AgentKillSwitchBasicAuth(): + stored_password: Final = existing.password if isinstance(existing, AgentKillSwitchBasicAuth) else None + return AgentKillSwitchBasicAuth( + type="basic", + username=incoming.username, + password=_restore_secret(incoming.password, stored_password), + ) + case _: + assert_never(incoming) + + +@dataclass(frozen=True, slots=True) +class KillSwitchRequest: + method: str + url: str + headers: Mapping[str, str] + json_body: Mapping[str, object] | None + + +def _auth_headers(auth: AgentKillSwitchAuth | None) -> Mapping[str, str]: + match auth: + case None: + return MappingProxyType({}) + case AgentKillSwitchBearerAuth(): + return MappingProxyType({"Authorization": f"Bearer {auth.token}"}) + case AgentKillSwitchApiKeyAuth(): + return MappingProxyType({auth.header_name: auth.api_key}) + case AgentKillSwitchBasicAuth(): + credentials: Final = b64encode(f"{auth.username}:{auth.password}".encode()).decode() + return MappingProxyType({"Authorization": f"Basic {credentials}"}) + case _: + assert_never(auth) + + +def build_kill_switch_request(config: AgentKillSwitchConfig) -> KillSwitchRequest: + url: Final = httpx.URL(config.url).copy_merge_params(config.query_params) + return KillSwitchRequest( + method=config.method, + url=str(url), + headers=MappingProxyType({**config.headers, **_auth_headers(config.auth)}), + json_body=config.body, + ) + + +class KillSwitchHttpClient(Protocol): + def build_request( + self, + method: str, + url: str, + *, + headers: Mapping[str, str], + json: Mapping[str, object] | None, + timeout: float, + ) -> httpx.Request: ... + + async def send(self, request: httpx.Request, *, stream: bool, follow_redirects: bool) -> httpx.Response: ... + + +def default_kill_switch_http_client() -> KillSwitchHttpClient: + return get_async_httpx_client(llm_provider=httpxSpecialProvider.AgentKillSwitch).client + + +KillSwitchAuditLogWriter: TypeAlias = Callable[[LiteLLM_AuditLogs], Awaitable[None]] # mutable-ok: Callable params + + +def default_kill_switch_audit_log_writer() -> KillSwitchAuditLogWriter: + return create_audit_log_for_update + + +def build_kill_switch_audit_log( + *, + result: AgentKillSwitchResult, + user_api_key_dict: UserAPIKeyAuth, + litellm_proxy_admin_name: str | None, +) -> LiteLLM_AuditLogs: + return LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by=get_audit_log_changed_by( + litellm_changed_by=None, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ), + changed_by_api_key=user_api_key_dict.api_key, + table_name=LitellmTableNames.AGENT_TABLE_NAME, + object_id=result.agent_id, + action="kill_switch_fired", + updated_values=result.model_dump_json(exclude_none=True), + ) + + +async def fire_kill_switch( + *, + agent_id: str, + config: AgentKillSwitchConfig, + http_client: KillSwitchHttpClient, + timeout: float = AGENT_KILL_SWITCH_TIMEOUT_SECONDS, +) -> AgentKillSwitchResult: + request: Final = build_kill_switch_request(config) + reported_url: Final = str(httpx.URL(request.url).copy_with(query=None)) + verbose_proxy_logger.info("Firing kill switch for agent %s: %s %s", agent_id, request.method, reported_url) + try: + response: Final = await http_client.send( + http_client.build_request( + request.method, + request.url, + headers=request.headers, + json=request.json_body, + timeout=timeout, + ), + stream=True, + follow_redirects=False, + ) + body: Final = await _read_text_prefix(response, AGENT_KILL_SWITCH_RESPONSE_BODY_MAX_CHARS) + except httpx.HTTPError as exc: + verbose_proxy_logger.warning("Kill switch for agent %s failed: %s", agent_id, type(exc).__name__) + return AgentKillSwitchResult( + agent_id=agent_id, + url=reported_url, + method=config.method, + error=type(exc).__name__, + ) + return AgentKillSwitchResult( + agent_id=agent_id, + url=reported_url, + method=config.method, + status_code=response.status_code, + response_body=body, + ) + + +async def _read_text_prefix(response: httpx.Response, max_chars: int) -> str: + try: + return await _take_text(response.aiter_text(), max_chars) + finally: + await response.aclose() + + +async def _take_text(chunks: AsyncIterator[str], max_chars: int) -> str: + taken = "" # rebind-ok: running prefix of a stream that is abandoned once the cap is hit + async for chunk in chunks: + taken += chunk # rebind-ok: see above + if len(taken) >= max_chars: + break + return taken[:max_chars] diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 85996430bc5..69c63d9ecd6 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -72,6 +72,7 @@ model LiteLLM_AgentsTable { agent_card_params Json static_headers Json? @default("{}") extra_headers String[] @default([]) + kill_switch Json? agent_access_groups String[] @default([]) access_group_ids String[] @default([]) object_permission_id String? diff --git a/litellm/types/agents.py b/litellm/types/agents.py index 7f8d8c6af66..f7aef09fa29 100644 --- a/litellm/types/agents.py +++ b/litellm/types/agents.py @@ -1,8 +1,9 @@ from collections.abc import Mapping, Sequence from datetime import datetime -from typing import TYPE_CHECKING, Any, Final, Literal +from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, TypeAlias +from urllib.parse import urlsplit -from pydantic import BaseModel, ConfigDict, PrivateAttr, StrictInt +from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, StrictInt, field_validator from typing_extensions import ReadOnly, Required, TypedDict from litellm.types.llms.base import LiteLLMPydanticObjectBase @@ -178,6 +179,74 @@ class AgentObjectPermission(TypedDict, total=False): agents: list[str] | None +class AgentKillSwitchBearerAuth(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + + type: Literal["bearer"] + token: str + + +class AgentKillSwitchApiKeyAuth(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + + type: Literal["api_key"] + header_name: str = "x-api-key" + api_key: str + + +class AgentKillSwitchBasicAuth(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + + type: Literal["basic"] + username: str + password: str + + +AgentKillSwitchAuth: TypeAlias = Annotated[ + AgentKillSwitchBearerAuth | AgentKillSwitchApiKeyAuth | AgentKillSwitchBasicAuth, + Field(discriminator="type"), +] + +AgentKillSwitchMethod: TypeAlias = Literal["POST", "PUT", "PATCH", "DELETE", "GET"] + + +class AgentKillSwitchConfig(BaseModel): + """Webhook an admin fires to shut an agent down out of band. LiteLLM only + makes the call; whatever the endpoint does with it is the agent's business.""" + + model_config = ConfigDict(frozen=True, extra="forbid") + + url: str + method: AgentKillSwitchMethod = "POST" + headers: Mapping[str, str] = Field(default_factory=dict) + query_params: Mapping[str, str] = Field(default_factory=dict) + body: Mapping[str, object] | None = None + auth: AgentKillSwitchAuth | None = None + + @field_validator("url") + @classmethod + def _require_absolute_http_url(cls, value: str) -> str: + parts: Final = urlsplit(value) + if parts.scheme not in ("http", "https") or not parts.netloc: + raise ValueError("kill_switch.url must be an absolute http(s) URL") + return value + + +class AgentKillSwitchResult(BaseModel): + model_config = ConfigDict(frozen=True) + + agent_id: str + url: str + method: AgentKillSwitchMethod + status_code: int | None = None + response_body: str | None = None + error: str | None = None + + @property + def succeeded(self) -> bool: + return self.status_code is not None and 200 <= self.status_code < 300 + + class AgentConfig(TypedDict, total=False): agent_name: Required[str] agent_card_params: Required[AgentCard] @@ -190,6 +259,7 @@ class AgentConfig(TypedDict, total=False): static_headers: dict[str, str] | None extra_headers: list[str] | None access_group_ids: ReadOnly[Sequence[str] | None] + kill_switch: ReadOnly[AgentKillSwitchConfig | None] class PatchAgentRequest(TypedDict, total=False): @@ -204,6 +274,7 @@ class PatchAgentRequest(TypedDict, total=False): static_headers: dict[str, str] | None extra_headers: list[str] | None access_group_ids: ReadOnly[Sequence[str] | None] + kill_switch: ReadOnly[AgentKillSwitchConfig | None] AGENT_CALLER_USER_ID_HEADER: Final = "x-litellm-user-id" @@ -243,6 +314,7 @@ class AgentResponse(BaseModel): static_headers: dict[str, str] | None = None extra_headers: list[str] | None = None access_group_ids: Sequence[str] | None = None + kill_switch: AgentKillSwitchConfig | None = None keys: list[AgentKeySummary] | None = None search_score: float | None = None created_at: datetime | None = None diff --git a/litellm/types/llms/custom_http.py b/litellm/types/llms/custom_http.py index 06982a16755..fa2d1373ea1 100644 --- a/litellm/types/llms/custom_http.py +++ b/litellm/types/llms/custom_http.py @@ -26,6 +26,7 @@ class httpxSpecialProvider(str, Enum): RAG = "rag" A2AProvider = "a2a_provider" AgentHealthCheck = "agent_health_check" + AgentKillSwitch = "agent_kill_switch" A2A = "a2a" PromptManagement = "prompt_management" UI = "ui" diff --git a/schema.prisma b/schema.prisma index 85996430bc5..69c63d9ecd6 100644 --- a/schema.prisma +++ b/schema.prisma @@ -72,6 +72,7 @@ model LiteLLM_AgentsTable { agent_card_params Json static_headers Json? @default("{}") extra_headers String[] @default([]) + kill_switch Json? agent_access_groups String[] @default([]) access_group_ids String[] @default([]) object_permission_id String? diff --git a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py index b036e0dac4d..ef20e88c368 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py @@ -451,7 +451,7 @@ async def test_update_agent_in_db_raises_when_row_deleted_mid_update(): registry: Final = AgentRegistry() mock_prisma: Final = MagicMock() mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( - return_value=SimpleNamespace(litellm_params={}, object_permission_id=None) + return_value=SimpleNamespace(litellm_params={}, object_permission_id=None, kill_switch=None) ) mock_prisma.db.litellm_agentstable.update = AsyncMock(return_value=None) @@ -736,6 +736,7 @@ async def test_update_agent_in_db_preserves_secret_when_echoed_back_redacted(): "model": "bedrock/agentcore/my-agent", }, object_permission_id=None, + kill_switch=None, ) ) updated_agent = MagicMock() @@ -784,6 +785,7 @@ async def test_update_agent_in_db_preserves_secret_when_key_omitted_entirely(): return_value=SimpleNamespace( litellm_params={"aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY}, object_permission_id=None, + kill_switch=None, ) ) updated_agent = MagicMock() @@ -830,6 +832,7 @@ async def test_update_agent_in_db_preserves_secret_nested_under_a_non_sensitive_ } }, object_permission_id=None, + kill_switch=None, ) ) updated_agent = MagicMock() @@ -878,6 +881,7 @@ async def test_update_agent_in_db_clears_secret_on_explicit_empty_value(): return_value=SimpleNamespace( litellm_params={"aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY}, object_permission_id=None, + kill_switch=None, ) ) updated_agent = MagicMock() @@ -1110,7 +1114,9 @@ async def test_update_agent_in_db_always_writes_access_group_ids(body_access_gro registry: Final = AgentRegistry() mock_prisma: Final = MagicMock() mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( - return_value=SimpleNamespace(litellm_params={}, object_permission_id=None, access_group_ids=["ag-1"]) + return_value=SimpleNamespace( + litellm_params={}, object_permission_id=None, kill_switch=None, access_group_ids=["ag-1"] + ) ) mock_update = AsyncMock(return_value=_agent_row_mock(expected)) mock_prisma.db.litellm_agentstable.update = mock_update @@ -1126,3 +1132,155 @@ async def test_update_agent_in_db_always_writes_access_group_ids(body_access_gro ) assert tuple(mock_update.call_args.kwargs["data"]["access_group_ids"]) == tuple(expected) + + +_KILL_SWITCH: Final = { + "url": "https://ops.example.com/kill", + "method": "POST", + "headers": {"X-Env": "prod"}, + "query_params": {"reason": "manual"}, + "body": {"action": "stop"}, + "auth": {"type": "bearer", "token": "tok-real"}, +} + + +@pytest.mark.asyncio +async def test_add_agent_to_db_stores_kill_switch_json_and_a_json_null_when_unset(): + registry: Final = AgentRegistry() + mock_prisma: Final = MagicMock() + mock_create = AsyncMock(return_value=_agent_row_mock([])) + mock_prisma.db.litellm_agentstable.create = mock_create + + await registry.add_agent_to_db( + agent={ + "agent_name": "Test Agent", + "agent_card_params": _sample_agent_card_params(), + "kill_switch": _KILL_SWITCH, + }, + prisma_client=mock_prisma, + created_by="test-user", + ) + assert json.loads(mock_create.call_args.kwargs["data"]["kill_switch"]) == _KILL_SWITCH + + await registry.add_agent_to_db( + agent={"agent_name": "Plain Agent", "agent_card_params": _sample_agent_card_params()}, + prisma_client=mock_prisma, + created_by="test-user", + ) + assert mock_create.call_args.kwargs["data"]["kill_switch"] == json.dumps(None) + + +@pytest.mark.asyncio +async def test_add_agent_to_db_rejects_a_kill_switch_with_a_non_http_url(): + registry: Final = AgentRegistry() + mock_prisma: Final = MagicMock() + mock_prisma.db.litellm_agentstable.create = AsyncMock(return_value=_agent_row_mock([])) + + with pytest.raises(Exception, match="absolute http"): + await registry.add_agent_to_db( + agent={ + "agent_name": "Test Agent", + "agent_card_params": _sample_agent_card_params(), + "kill_switch": {**_KILL_SWITCH, "url": "ops.example.com/kill"}, + }, + prisma_client=mock_prisma, + created_by="test-user", + ) + mock_prisma.db.litellm_agentstable.create.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_patch_agent_in_db_keeps_kill_switch_when_omitted_and_clears_it_on_null(): + registry: Final = AgentRegistry() + mock_prisma: Final = MagicMock() + mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( + return_value={ + "agent_id": "agent-123", + "agent_name": "Old", + "litellm_params": {}, + "object_permission_id": None, + "kill_switch": _KILL_SWITCH, + } + ) + mock_update = AsyncMock(return_value=_agent_row_mock([])) + mock_prisma.db.litellm_agentstable.update = mock_update + + await registry.patch_agent_in_db( + agent_id="agent-123", agent={"agent_name": "New"}, prisma_client=mock_prisma, updated_by="u" + ) + assert "kill_switch" not in mock_update.call_args.kwargs["data"] + + await registry.patch_agent_in_db( + agent_id="agent-123", agent={"kill_switch": None}, prisma_client=mock_prisma, updated_by="u" + ) + assert mock_update.call_args.kwargs["data"]["kill_switch"] == json.dumps(None), ( + "prisma-client-py silently drops None, so the clear must be written as the JSON literal null" + ) + + +@pytest.mark.asyncio +async def test_patch_agent_in_db_restores_the_stored_kill_switch_secret_behind_the_marker(): + registry: Final = AgentRegistry() + mock_prisma: Final = MagicMock() + mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( + return_value={ + "agent_id": "agent-123", + "agent_name": "A", + "litellm_params": {}, + "object_permission_id": None, + "kill_switch": _KILL_SWITCH, + } + ) + mock_update = AsyncMock(return_value=_agent_row_mock([])) + mock_prisma.db.litellm_agentstable.update = mock_update + + await registry.patch_agent_in_db( + agent_id="agent-123", + agent={ + "kill_switch": { + **_KILL_SWITCH, + "url": "https://ops.example.com/v2/kill", + "auth": {"type": "bearer", "token": REDACTED_BY_LITELM_STRING}, + } + }, + prisma_client=mock_prisma, + updated_by="u", + ) + + assert json.loads(mock_update.call_args.kwargs["data"]["kill_switch"]) == { + **_KILL_SWITCH, + "url": "https://ops.example.com/v2/kill", + } + + +@pytest.mark.asyncio +async def test_update_agent_in_db_clears_kill_switch_when_omitted_and_restores_secret_when_echoed(): + registry: Final = AgentRegistry() + mock_prisma: Final = MagicMock() + mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( + return_value=SimpleNamespace(litellm_params={}, object_permission_id=None, kill_switch=json.dumps(_KILL_SWITCH)) + ) + mock_update = AsyncMock(return_value=_agent_row_mock([])) + mock_prisma.db.litellm_agentstable.update = mock_update + base: Final = {"agent_name": "Test Agent", "agent_card_params": _sample_agent_card_params(), "litellm_params": {}} + + await registry.update_agent_in_db(agent_id="agent-123", agent=base, prisma_client=mock_prisma, updated_by="u") + assert mock_update.call_args.kwargs["data"]["kill_switch"] == json.dumps(None) + + echoed: Final = {**_KILL_SWITCH, "auth": {"type": "bearer", "token": REDACTED_BY_LITELM_STRING}} + await registry.update_agent_in_db( + agent_id="agent-123", agent={**base, "kill_switch": echoed}, prisma_client=mock_prisma, updated_by="u" + ) + assert json.loads(mock_update.call_args.kwargs["data"]["kill_switch"]) == _KILL_SWITCH + + +def test_load_agents_from_config_exposes_a_typed_kill_switch(): + registry: Final = AgentRegistry() + + registry.load_agents_from_config( + [{"agent_name": "cfg-agent", "agent_card_params": _sample_agent_card_params(), "kill_switch": _KILL_SWITCH}] + ) + + (agent,) = registry.get_agent_list() + assert agent.kill_switch is not None + assert agent.kill_switch.model_dump() == _KILL_SWITCH diff --git a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py index 482294e7b92..526f24c5221 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py @@ -1,13 +1,15 @@ import json +from types import SimpleNamespace from typing import Final from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest from fastapi import FastAPI from fastapi.testclient import TestClient from litellm.constants import REDACTED_BY_LITELM_STRING -from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy._types import LiteLLM_AuditLogs, LitellmTableNames, LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.agent_endpoints import endpoints as agent_endpoints from litellm.proxy.agent_endpoints.auth.agent_permission_handler import ( RestrictedAgentAccess, @@ -1136,3 +1138,222 @@ def test_make_agent_public_rejects_an_agent_published_only_in_the_db(monkeypatch assert duplicate.status_code == 400 assert "already in public agent groups" in duplicate.json()["detail"] + + +_KILL_SWITCH: Final = { + "url": "https://ops.example.com/kill", + "method": "POST", + "headers": {"X-Env": "prod"}, + "query_params": {"reason": "manual"}, + "body": {"action": "stop"}, + "auth": {"type": "bearer", "token": "tok-real"}, +} + + +def _agent_with_kill_switch() -> AgentResponse: + return AgentResponse( + agent_id="agent-123", + agent_name="Test Agent", + agent_card_params=_sample_agent_card_params(), + litellm_params={}, + kill_switch=_KILL_SWITCH, + ) + + +class _FakeKillSwitchClient: + def __init__(self, response: httpx.Response) -> None: + self.calls: list[tuple[str, str, dict[str, str], object, float]] = [] # mutable-ok: test double records calls + self._response: Final = response + + def build_request(self, method: str, url: str, *, headers, json, timeout: float) -> httpx.Request: + self.calls.append((method, url, dict(headers), json, timeout)) + return httpx.Request(method, url, headers=dict(headers), json=json) + + async def send(self, request: httpx.Request, *, stream: bool, follow_redirects: bool) -> httpx.Response: + return self._response + + +class _AuditLogRecorder: + def __init__(self) -> None: + self.rows: list[LiteLLM_AuditLogs] = [] # mutable-ok: test double records writes + + async def __call__(self, request_data: LiteLLM_AuditLogs) -> None: + self.rows.append(request_data) + + +def _kill_switch_app( + role: LitellmUserRoles, + http_client: _FakeKillSwitchClient, + audit_log: _AuditLogRecorder | None = None, +) -> TestClient: + test_client: Final = _make_app_with_role(role) + test_client.app.dependency_overrides[agent_endpoints.default_kill_switch_http_client] = lambda: http_client + test_client.app.dependency_overrides[agent_endpoints.default_kill_switch_audit_log_writer] = ( + lambda: audit_log or _AuditLogRecorder() + ) + return test_client + + +def test_kill_switch_trigger_fires_the_configured_webhook_and_returns_the_result(monkeypatch) -> None: + registry: Final = MagicMock() + registry.get_agent_by_id = MagicMock(return_value=_agent_with_kill_switch()) + monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + fake: Final = _FakeKillSwitchClient(httpx.Response(200, text="ok")) + + resp: Final = _kill_switch_app(LitellmUserRoles.PROXY_ADMIN, fake).post( + "/v1/agents/agent-123/kill_switch", headers={"Authorization": "Bearer k"} + ) + + assert resp.status_code == 200, resp.text + assert resp.json() == { + "agent_id": "agent-123", + "url": "https://ops.example.com/kill", + "method": "POST", + "status_code": 200, + "response_body": "ok", + "error": None, + } + (method, url, headers, body, _timeout) = fake.calls[0] + assert (method, url, body) == ("POST", "https://ops.example.com/kill?reason=manual", {"action": "stop"}) + assert headers == {"X-Env": "prod", "Authorization": "Bearer tok-real"} + + +def test_kill_switch_trigger_writes_an_audit_log_row_naming_the_admin_and_the_sanitized_result(monkeypatch) -> None: + registry: Final = MagicMock() + registry.get_agent_by_id = MagicMock(return_value=_agent_with_kill_switch()) + monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + fake: Final = _FakeKillSwitchClient(httpx.Response(202, text='{"stopped": true}')) + audit: Final = _AuditLogRecorder() + test_client: Final = _kill_switch_app(LitellmUserRoles.PROXY_ADMIN, fake, audit) + test_client.app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user", user_role=LitellmUserRoles.PROXY_ADMIN, api_key="hashed-k" + ) + + resp: Final = test_client.post("/v1/agents/agent-123/kill_switch", headers={"Authorization": "Bearer k"}) + + assert resp.status_code == 200, resp.text + (row,) = audit.rows + assert (row.action, row.table_name, row.object_id) == ( + "kill_switch_fired", + LitellmTableNames.AGENT_TABLE_NAME, + "agent-123", + ) + assert (row.changed_by, row.changed_by_api_key) == ("test-user", "hashed-k") + assert row.before_value is None + assert json.loads(row.updated_values) == { + "agent_id": "agent-123", + "url": "https://ops.example.com/kill", + "method": "POST", + "status_code": 202, + "response_body": '{"stopped": true}', + } + assert "tok-real" not in row.model_dump_json() + + +def test_kill_switch_trigger_returns_502_and_still_audits_when_the_webhook_rejects(monkeypatch) -> None: + registry: Final = MagicMock() + registry.get_agent_by_id = MagicMock(return_value=_agent_with_kill_switch()) + monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + fake: Final = _FakeKillSwitchClient(httpx.Response(401, text="bad token")) + audit: Final = _AuditLogRecorder() + + resp: Final = _kill_switch_app(LitellmUserRoles.PROXY_ADMIN, fake, audit).post( + "/v1/agents/agent-123/kill_switch", headers={"Authorization": "Bearer k"} + ) + + assert resp.status_code == 502, resp.text + assert resp.json()["detail"]["status_code"] == 401 + assert resp.json()["detail"]["response_body"] == "bad token" + (row,) = audit.rows + assert row.action == "kill_switch_fired" + assert json.loads(row.updated_values)["status_code"] == 401 + + +@pytest.mark.parametrize("role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY]) +def test_kill_switch_trigger_is_refused_before_any_webhook_call_for_non_admins(monkeypatch, role) -> None: + registry: Final = MagicMock() + registry.get_agent_by_id = MagicMock(return_value=_agent_with_kill_switch()) + monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + fake: Final = _FakeKillSwitchClient(httpx.Response(200)) + audit: Final = _AuditLogRecorder() + + resp: Final = _kill_switch_app(role, fake, audit).post( + "/v1/agents/agent-123/kill_switch", headers={"Authorization": "Bearer k"} + ) + + assert resp.status_code == 403, resp.text + assert fake.calls == [] + assert audit.rows == [] + + +def test_kill_switch_trigger_404s_unknown_agent_and_400s_an_agent_without_one(monkeypatch) -> None: + registry: Final = MagicMock() + registry.get_agent_by_id = MagicMock(side_effect=[None, _sample_agent_response()]) + monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + fake: Final = _FakeKillSwitchClient(httpx.Response(200)) + audit: Final = _AuditLogRecorder() + test_client: Final = _kill_switch_app(LitellmUserRoles.PROXY_ADMIN, fake, audit) + + missing: Final = test_client.post("/v1/agents/nope/kill_switch", headers={"Authorization": "Bearer k"}) + unconfigured: Final = test_client.post("/v1/agents/agent-123/kill_switch", headers={"Authorization": "Bearer k"}) + + assert missing.status_code == 404 + assert unconfigured.status_code == 400 + assert "no kill_switch configured" in unconfigured.json()["detail"] + assert fake.calls == [] + assert audit.rows == [] + + +def test_kill_switch_trigger_fires_the_db_row_config_over_a_stale_in_memory_copy(monkeypatch) -> None: + """Another replica may have updated the agent; the row is the source of truth for what gets fired.""" + registry: Final = MagicMock() + registry.get_agent_by_id = MagicMock(return_value=_agent_with_kill_switch()) + monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry) + db_row: Final = SimpleNamespace( + agent_id="agent-123", + kill_switch={"url": "https://ops.example.com/kill-v2", "method": "DELETE", "auth": None}, + ) + prisma: Final = MagicMock() + prisma.db.litellm_agentstable.find_unique = AsyncMock(return_value=db_row) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma) + fake: Final = _FakeKillSwitchClient(httpx.Response(204)) + + resp: Final = _kill_switch_app(LitellmUserRoles.PROXY_ADMIN, fake).post( + "/v1/agents/agent-123/kill_switch", headers={"Authorization": "Bearer k"} + ) + + assert resp.status_code == 200, resp.text + (method, url, headers, body, _timeout) = fake.calls[0] + assert (method, url, headers, body) == ("DELETE", "https://ops.example.com/kill-v2", {}, None) + assert prisma.db.litellm_agentstable.find_unique.await_args.kwargs == {"where": {"agent_id": "agent-123"}} + registry.get_agent_by_id.assert_not_called() + + +def test_get_agent_redacts_kill_switch_secret_for_admins_and_hides_it_from_others(monkeypatch) -> None: + registry: Final = MagicMock() + registry.get_agent_by_id = MagicMock(return_value=_agent_with_kill_switch()) + registry.ids_for_agent = MagicMock(return_value=("agent-123",)) + monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry) + + def _get_as(role: LitellmUserRoles): + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(return_value=None) + mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + return _make_app_with_role(role).get("/v1/agents/agent-123", headers={"Authorization": "Bearer k"}) + + admin: Final = _get_as(LitellmUserRoles.PROXY_ADMIN) + assert admin.status_code == 200, admin.text + assert admin.json()["kill_switch"] == { + **_KILL_SWITCH, + "auth": {"type": "bearer", "token": REDACTED_BY_LITELM_STRING}, + } + + internal: Final = _get_as(LitellmUserRoles.INTERNAL_USER) + assert internal.status_code == 200, internal.text + assert internal.json()["kill_switch"] is None + assert "tok-real" not in internal.text diff --git a/tests/test_litellm/proxy/agent_endpoints/test_kill_switch.py b/tests/test_litellm/proxy/agent_endpoints/test_kill_switch.py new file mode 100644 index 00000000000..a6bb945713e --- /dev/null +++ b/tests/test_litellm/proxy/agent_endpoints/test_kill_switch.py @@ -0,0 +1,248 @@ +from base64 import b64encode +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Final + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.constants import REDACTED_BY_LITELM_STRING +from litellm.proxy.agent_endpoints.kill_switch import ( + build_kill_switch_request, + fire_kill_switch, + redact_kill_switch, + restore_kill_switch, +) +from litellm.types.agents import AgentKillSwitchConfig + + +@dataclass(frozen=True, slots=True) +class _SentRequest: + method: str + url: str + headers: Mapping[str, str] + json: Mapping[str, object] | None + timeout: float + + +class _RecordingClient: + def __init__(self, respond: httpx.Response | httpx.HTTPError) -> None: + self.sent: list[_SentRequest] = [] # mutable-ok: test double records calls + self.follow_redirects: list[bool] = [] # mutable-ok: test double records calls + self._respond: Final = respond + + def build_request( + self, + method: str, + url: str, + *, + headers: Mapping[str, str], + json: Mapping[str, object] | None, + timeout: float, + ) -> httpx.Request: + self.sent.append(_SentRequest(method, url, headers, json, timeout)) + return httpx.Request(method, url, headers=dict(headers), json=json) + + async def send(self, request: httpx.Request, *, stream: bool, follow_redirects: bool) -> httpx.Response: + self.follow_redirects.append(follow_redirects) + if isinstance(self._respond, httpx.HTTPError): + raise self._respond + return self._respond + + +class _CountingStream(httpx.AsyncByteStream): + def __init__(self, chunk: bytes, chunks: int) -> None: + self.pulled: int = 0 # rebind-ok: test double counts reads + self._chunk: Final = chunk + self._chunks: Final = chunks + + async def __aiter__(self): + for _ in range(self._chunks): + self.pulled += 1 # rebind-ok: test double counts reads + yield self._chunk + + +def _config(**overrides: object) -> AgentKillSwitchConfig: + return AgentKillSwitchConfig.model_validate({"url": "https://ops.example.com/agents/kill", **overrides}) + + +def test_request_carries_endpoint_method_query_params_headers_and_body() -> None: + request: Final = build_kill_switch_request( + _config( + url="https://ops.example.com/kill?env=prod", + method="PUT", + query_params={"agent": "billing-bot", "reason": "manual stop"}, + headers={"X-Trace": "abc"}, + body={"action": "stop", "hard": True}, + ) + ) + + assert request.method == "PUT" + assert str(httpx.URL(request.url)) == "https://ops.example.com/kill?env=prod&agent=billing-bot&reason=manual+stop" + assert dict(request.headers) == {"X-Trace": "abc"} + assert request.json_body == {"action": "stop", "hard": True} + + +def test_request_defaults_to_post_with_no_body_and_untouched_url() -> None: + request: Final = build_kill_switch_request(_config()) + + assert (request.method, request.url, dict(request.headers), request.json_body) == ( + "POST", + "https://ops.example.com/agents/kill", + {}, + None, + ) + + +@pytest.mark.parametrize( + ("auth", "expected_headers"), + [ + ({"type": "bearer", "token": "tok-123"}, {"Authorization": "Bearer tok-123"}), + ({"type": "api_key", "api_key": "k-456"}, {"x-api-key": "k-456"}), + ({"type": "api_key", "header_name": "X-Ops-Key", "api_key": "k-456"}, {"X-Ops-Key": "k-456"}), + ( + {"type": "basic", "username": "ops", "password": "pw:1"}, + {"Authorization": f"Basic {b64encode(b'ops:pw:1').decode()}"}, + ), + ], +) +def test_auth_becomes_the_matching_request_header(auth: Mapping[str, object], expected_headers: dict[str, str]) -> None: + request: Final = build_kill_switch_request(_config(auth=auth)) + + assert dict(request.headers) == expected_headers + + +def test_auth_header_wins_over_a_conflicting_custom_header() -> None: + request: Final = build_kill_switch_request( + _config(headers={"Authorization": "stale", "X-Env": "prod"}, auth={"type": "bearer", "token": "fresh"}) + ) + + assert dict(request.headers) == {"Authorization": "Bearer fresh", "X-Env": "prod"} + + +@pytest.mark.parametrize("url", ["ftp://ops.example.com/kill", "/relative/kill", "ops.example.com/kill", ""]) +def test_config_rejects_non_http_urls(url: str) -> None: + with pytest.raises(ValidationError, match="absolute http"): + _config(url=url) + + +def test_config_rejects_unknown_auth_type_and_unknown_fields() -> None: + with pytest.raises(ValidationError): + _config(auth={"type": "hmac", "secret": "x"}) + with pytest.raises(ValidationError): + _config(endpoint="https://typo.example.com") + + +@pytest.mark.parametrize( + ("auth", "secret_field"), + [ + ({"type": "bearer", "token": "tok-123"}, "token"), + ({"type": "api_key", "header_name": "X-K", "api_key": "k-456"}, "api_key"), + ({"type": "basic", "username": "ops", "password": "pw"}, "password"), + ], +) +def test_redact_replaces_only_the_secret_and_restore_puts_it_back(auth: dict[str, str], secret_field: str) -> None: + original: Final = _config(auth=auth) + + redacted: Final = redact_kill_switch(original) + assert redacted is not None and redacted.auth is not None + assert redacted.auth.model_dump() == {**auth, secret_field: REDACTED_BY_LITELM_STRING} + assert original.auth is not None and original.auth.model_dump() == auth, "redact must not mutate its input" + + restored: Final = restore_kill_switch(redacted, original) + assert restored == original + + +def test_restore_keeps_a_rotated_secret_and_never_stores_the_marker_itself() -> None: + rotated: Final = _config(auth={"type": "bearer", "token": "new-token"}) + stored: Final = _config(auth={"type": "bearer", "token": "old-token"}) + assert restore_kill_switch(rotated, stored) == rotated + assert restore_kill_switch(None, stored) is None + + marker_only: Final = _config(auth={"type": "bearer", "token": REDACTED_BY_LITELM_STRING}) + assert restore_kill_switch(marker_only, None) == _config(auth={"type": "bearer", "token": ""}) + + +def test_restore_does_not_borrow_a_secret_from_a_different_auth_type() -> None: + incoming: Final = _config(auth={"type": "bearer", "token": REDACTED_BY_LITELM_STRING}) + stored: Final = _config(auth={"type": "api_key", "api_key": "k-456"}) + + assert restore_kill_switch(incoming, stored) == _config(auth={"type": "bearer", "token": ""}) + + +def test_redact_passes_through_configs_without_auth() -> None: + assert redact_kill_switch(None) is None + plain: Final = _config(headers={"X-Env": "prod"}) + assert redact_kill_switch(plain) is plain + + +@pytest.mark.asyncio +async def test_fire_sends_exactly_the_built_request_and_reports_the_2xx_reply_without_the_query() -> None: + client: Final = _RecordingClient(httpx.Response(202, text="stopping")) + config: Final = _config( + method="DELETE", + query_params={"force": "1", "token": "qs-secret"}, + headers={"X-Env": "prod"}, + body={"agent": "billing-bot"}, + auth={"type": "bearer", "token": "tok-123"}, + ) + + result: Final = await fire_kill_switch(agent_id="agent-1", config=config, http_client=client, timeout=3.5) + + assert client.sent == [ + _SentRequest( + method="DELETE", + url="https://ops.example.com/agents/kill?force=1&token=qs-secret", + headers={"X-Env": "prod", "Authorization": "Bearer tok-123"}, + json={"agent": "billing-bot"}, + timeout=3.5, + ) + ] + assert client.follow_redirects == [False], "a redirecting webhook must not be followed to another host" + assert result.succeeded is True + assert result.model_dump() == { + "agent_id": "agent-1", + "url": "https://ops.example.com/agents/kill", + "method": "DELETE", + "status_code": 202, + "response_body": "stopping", + "error": None, + } + + +@pytest.mark.asyncio +async def test_fire_reports_a_non_2xx_reply_as_failure_with_the_body() -> None: + client: Final = _RecordingClient(httpx.Response(503, text="x" * 5000)) + + result: Final = await fire_kill_switch(agent_id="agent-1", config=_config(), http_client=client) + + assert result.succeeded is False + assert result.status_code == 503 + assert result.response_body == "x" * 2000 + assert result.error is None + + +@pytest.mark.asyncio +async def test_fire_stops_reading_the_body_at_the_cap_instead_of_buffering_the_whole_reply() -> None: + stream: Final = _CountingStream(b"y" * 500, chunks=100) + client: Final = _RecordingClient(httpx.Response(200, stream=stream)) + + result: Final = await fire_kill_switch(agent_id="agent-1", config=_config(), http_client=client) + + assert result.response_body == "y" * 2000 + assert stream.pulled == 4, f"read {stream.pulled} of 100 chunks for a 2000 char cap" + + +@pytest.mark.asyncio +async def test_fire_reports_a_transport_error_by_type_without_raising_or_echoing_the_url() -> None: + client: Final = _RecordingClient(httpx.ConnectError("boom https://ops.example.com/agents/kill?token=qs-secret")) + + result: Final = await fire_kill_switch( + agent_id="agent-1", config=_config(query_params={"token": "qs-secret"}), http_client=client + ) + + assert result.succeeded is False + assert (result.status_code, result.response_body) == (None, None) + assert result.error == "ConnectError" + assert "qs-secret" not in result.model_dump_json() diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index 7bb79a115dd..f76a02e8361 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -3776,6 +3776,7 @@ AGENT_MANAGEMENT_ROUTES = [ "/v1/agents/abc-123", "/v1/agents/make_public", "/v1/agents/abc-123/make_public", + "/v1/agents/abc-123/kill_switch", ] AGENT_INFERENCE_ROUTES = [ diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentFormKit.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentFormKit.tsx index 8e100d0c3ed..3d863036234 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentFormKit.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentFormKit.tsx @@ -27,6 +27,7 @@ import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/component import { Input } from "@/components/ui/input"; import { Tooltip, TooltipContent, TooltipTrigger } from "@/components/ui/tooltip"; import { Field, FieldDescription, FieldError, FieldGroup, FieldLabel } from "@/components/ui/field"; +import type { KeyValueFormValue, KillSwitchConfig, KillSwitchFormValue } from "./kill_switch_config"; export interface AgentSkillFormValue { id?: string; @@ -54,6 +55,8 @@ export type AgentFormFieldValue = | string[] | AgentSkillFormValue[] | StaticHeaderFormValue[] + | KeyValueFormValue[] + | KillSwitchFormValue | McpServerSelection | Record | null @@ -82,6 +85,7 @@ export interface AgentFormValues { output_cost_per_token?: string | number; static_headers?: StaticHeaderFormValue[]; extra_headers?: string[]; + kill_switch?: KillSwitchFormValue; tpm_limit?: number | null; rpm_limit?: number | null; session_tpm_limit?: number | null; @@ -123,6 +127,7 @@ export interface AgentRequestPayload { litellm_params?: Record; object_permission?: Record; access_group_ids?: string[]; + kill_switch?: KillSwitchConfig | null; } interface AgentFormFieldProps { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentKillSwitchDangerZone.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentKillSwitchDangerZone.test.tsx new file mode 100644 index 00000000000..aa4b83478b6 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentKillSwitchDangerZone.test.tsx @@ -0,0 +1,124 @@ +import React from "react"; +import { fireEvent, render, screen, waitFor, within } from "@testing-library/react"; +import { describe, it, expect, vi, beforeEach } from "vitest"; +import AgentKillSwitchDangerZone from "./AgentKillSwitchDangerZone"; +import * as networking from "@/components/networking"; +import { toast } from "@/lib/toast"; + +vi.mock("@/components/networking", () => ({ + triggerAgentKillSwitchCall: vi.fn(), +})); + +vi.mock("@/lib/toast", () => ({ + toast: { success: vi.fn(), error: vi.fn() }, +})); + +const killSwitch = { url: "https://ops.example.com/kill", method: "DELETE" as const }; + +const renderZone = (props: Partial> = {}) => + render( + , + ); + +const openDialog = () => { + fireEvent.click(screen.getByRole("button", { name: "Fire Kill Switch" })); + return screen.getByRole("dialog"); +}; + +const dialogFireButton = () => within(screen.getByRole("dialog")).getByRole("button", { name: "Fire Kill Switch" }); + +describe("AgentKillSwitchDangerZone", () => { + beforeEach(() => { + vi.mocked(networking.triggerAgentKillSwitchCall).mockReset(); + vi.mocked(toast.success).mockReset(); + vi.mocked(toast.error).mockReset(); + }); + + it("renders nothing for non-admins", () => { + const { container } = renderZone({ isAdmin: false }); + + expect(container).toBeEmptyDOMElement(); + }); + + it("shows the webhook target and an outage warning inside a Danger Zone region", () => { + renderZone(); + + const region = screen.getByRole("region", { name: "Danger Zone" }); + expect(region).toHaveTextContent("DELETE https://ops.example.com/kill"); + expect(region).toHaveTextContent("can cause an outage"); + expect(screen.getByRole("button", { name: "Fire Kill Switch" })).toBeEnabled(); + }); + + it("shows an unconfigured notice without a fire button when no kill switch is set", () => { + renderZone({ killSwitch: null }); + + expect(screen.getByRole("region", { name: "Danger Zone" })).toHaveTextContent("Not configured"); + expect(screen.queryByRole("button", { name: "Fire Kill Switch" })).not.toBeInTheDocument(); + }); + + it("keeps the confirm button disabled until the exact agent name is typed", () => { + renderZone(); + openDialog(); + + expect(dialogFireButton()).toBeDisabled(); + + fireEvent.change(screen.getByLabelText("Confirm agent name"), { target: { value: "support-agen" } }); + expect(dialogFireButton()).toBeDisabled(); + + fireEvent.change(screen.getByLabelText("Confirm agent name"), { target: { value: "support-agent" } }); + expect(dialogFireButton()).toBeEnabled(); + expect(networking.triggerAgentKillSwitchCall).not.toHaveBeenCalled(); + }); + + it("fires the webhook after typed confirmation, closes the dialog and shows the sanitized result", async () => { + const firedResult = { + agent_id: "agent-1", + url: killSwitch.url, + method: "DELETE" as const, + status_code: 202, + response_body: '{"stopped": true}', + }; + vi.mocked(networking.triggerAgentKillSwitchCall).mockResolvedValue(firedResult); + renderZone(); + openDialog(); + + fireEvent.change(screen.getByLabelText("Confirm agent name"), { target: { value: "support-agent" } }); + fireEvent.click(dialogFireButton()); + + expect(await screen.findByRole("status")).toHaveTextContent('Last result: HTTP 202 {"stopped": true}'); + await waitFor(() => expect(screen.queryByRole("dialog")).not.toBeInTheDocument()); + expect(networking.triggerAgentKillSwitchCall).toHaveBeenCalledWith("sk-test", "agent-1"); + expect(toast.success).toHaveBeenCalledWith("Kill switch fired (HTTP 202)"); + }); + + it("does not call the webhook when the dialog is cancelled", () => { + renderZone(); + openDialog(); + + fireEvent.change(screen.getByLabelText("Confirm agent name"), { target: { value: "support-agent" } }); + fireEvent.click(screen.getByRole("button", { name: "Cancel" })); + + expect(networking.triggerAgentKillSwitchCall).not.toHaveBeenCalled(); + expect(screen.queryByRole("status")).not.toBeInTheDocument(); + }); + + it("surfaces a failed webhook as an error toast and keeps the dialog open", async () => { + vi.mocked(networking.triggerAgentKillSwitchCall).mockRejectedValue(new Error("Kill switch webhook returned 500")); + renderZone(); + openDialog(); + + fireEvent.change(screen.getByLabelText("Confirm agent name"), { target: { value: "support-agent" } }); + fireEvent.click(dialogFireButton()); + + await waitFor(() => expect(toast.error).toHaveBeenCalledWith("Kill switch webhook returned 500")); + expect(screen.getByRole("dialog")).toBeInTheDocument(); + expect(screen.queryByRole("status")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentKillSwitchDangerZone.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentKillSwitchDangerZone.tsx new file mode 100644 index 00000000000..8d68a8f921b --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentKillSwitchDangerZone.tsx @@ -0,0 +1,147 @@ +import { CircleAlert } from "lucide-react"; +import React, { useState } from "react"; +import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert"; +import { Button } from "@/components/ui/button"; +import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog"; +import { InputGroup, InputGroupAddon, InputGroupInput } from "@/components/ui/input-group"; +import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; +import { toast } from "@/lib/toast"; +import { AgentKillSwitchResult, triggerAgentKillSwitchCall } from "@/components/networking"; +import { KillSwitchConfig } from "./kill_switch_config"; + +interface AgentKillSwitchDangerZoneProps { + agentId: string; + agentName: string; + killSwitch: KillSwitchConfig | null | undefined; + accessToken: string | null; + isAdmin: boolean; +} + +const AgentKillSwitchDangerZone: React.FC = ({ + agentId, + agentName, + killSwitch, + accessToken, + isAdmin, +}) => { + const [isConfirmOpen, setIsConfirmOpen] = useState(false); + const [confirmationInput, setConfirmationInput] = useState(""); + const [isFiring, setIsFiring] = useState(false); + const [lastResult, setLastResult] = useState(null); + + if (!isAdmin) return null; + + const openConfirm = () => { + setConfirmationInput(""); + setIsConfirmOpen(true); + }; + + const fire = async () => { + if (!accessToken) return; + setIsFiring(true); + setLastResult(null); + try { + const result = await triggerAgentKillSwitchCall(accessToken, agentId); + setLastResult(result); + setIsConfirmOpen(false); + toast.success(`Kill switch fired (HTTP ${result.status_code})`); + } catch (error) { + toast.error(error instanceof Error ? error.message : "Failed to fire kill switch"); + } finally { + setIsFiring(false); + } + }; + + return ( +
+

+ Danger Zone +

+
+
+
+

Kill switch

+ {killSwitch ? ( + <> +

+ Calls the configured webhook to stop this agent's upstream runtime. This can cause an outage for + everyone using the agent and cannot be undone from LiteLLM +

+

+ {killSwitch.method ?? "POST"} {killSwitch.url} +

+ + ) : ( +

+ Not configured. Add a kill switch webhook under Settings to enable this action +

+ )} +
+ {killSwitch && ( + + )} +
+ {lastResult && ( +

+ Last result: HTTP {lastResult.status_code} + {lastResult.response_body ? ` ${lastResult.response_body}` : ""} +

+ )} +
+ + !open && !isFiring && setIsConfirmOpen(false)}> + + + Fire kill switch for {agentName}? + +
+ + + This can cause an outage + + LiteLLM will call {killSwitch?.method ?? "POST"} {killSwitch?.url} immediately. Whatever that webhook + does to the agent is outside LiteLLM's control and cannot be reverted here + + +
+

+ Type {agentName} to confirm: +

+ + + + + setConfirmationInput(e.target.value)} + placeholder={agentName} + autoFocus + /> + +
+
+ + + + +
+
+
+ ); +}; + +export default AgentKillSwitchDangerZone; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/KillSwitchFormFields.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/KillSwitchFormFields.tsx new file mode 100644 index 00000000000..578d80da6ed --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/KillSwitchFormFields.tsx @@ -0,0 +1,223 @@ +import React from "react"; +import { useFieldArray, useFormContext, useWatch } from "react-hook-form"; +import { Plus, Trash2 } from "lucide-react"; +import { Button } from "@/components/ui/button"; +import { Input } from "@/components/ui/input"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { Textarea } from "@/components/ui/textarea"; +import { Field, FieldTitle } from "@/components/ui/field"; +import { PasswordInput } from "@/components/shared/PasswordInput"; +import { AgentFormField, AgentFormValues, labelWithHint } from "./AgentFormKit"; +import { KILL_SWITCH_AUTH_TYPES, KILL_SWITCH_METHODS, validateKillSwitchBody } from "./kill_switch_config"; + +const KeyValueFieldArray = ({ + name, + addLabel, + keyPlaceholder, + valuePlaceholder, +}: { + name: "kill_switch.headers" | "kill_switch.query_params"; + addLabel: string; + keyPlaceholder: string; + valuePlaceholder: string; +}) => { + const { control } = useFormContext(); + const { fields, append, remove } = useFieldArray({ control, name }); + + return ( +
+ {fields.map((item, index) => ( +
+ + {({ value, onChange, ref, ...control }) => ( + + )} + + + {({ value, onChange, ref, ...control }) => ( + + )} + + +
+ ))} + +
+ ); +}; + +const TextField = ({ + name, + label, + placeholder, + required, + secret, +}: { + name: `kill_switch.${string}`; + label: React.ReactNode; + placeholder?: string; + required?: string; + secret?: boolean; +}) => ( + + {({ value, onChange, ref, ...control }) => + secret ? ( + + ) : ( + + ) + } + +); + +const KillSwitchAuthFields = () => { + const { control } = useFormContext(); + const authType = useWatch({ control, name: "kill_switch.auth_type" }); + + switch (authType) { + case "bearer": + return ; + case "api_key": + return ( + <> + + + + ); + case "basic": + return ( + <> + + + + ); + default: + return null; + } +}; + +const KillSwitchFormFields = () => ( + <> + + + + {({ value, onChange, ref: _ref, ...control }) => ( + + )} + + + + Headers + + + + + Query Parameters + + + + validateKillSwitchBody(typeof value === "string" ? value : "") }} + > + {({ value, onChange, ref, ...control }) => ( +