diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 85e96e22cb4..6a7006ea8d6 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -16,11 +16,7 @@ from pydantic import ( from typing_extensions import Required, TypedDict from litellm.types.integrations.slack_alerting import AlertType -from litellm.types.llms.openai import ( - AllMessageValues, - ChatCompletionRequest, - OpenAIFileObject, -) +from litellm.types.llms.openai import AllMessageValues, OpenAIFileObject from litellm.types.mcp import ( MCPAuthType, MCPSpecVersion, @@ -576,13 +572,47 @@ class LiteLLMPromptInjectionParams(LiteLLMPydanticObjectBase): ######### Request Class Definition ###### -class ProxyChatCompletionRequest(ChatCompletionRequest): +class ProxyChatCompletionRequest(LiteLLMPydanticObjectBase): + """ + Pydantic model for chat completion requests that includes both OpenAI standard fields + and LiteLLM-specific parameters. This replaces the previous TypedDict version. + """ + # Required fields (from ChatCompletionRequest) + model: str + messages: List[AllMessageValues] + + # Standard OpenAI completion parameters (all optional) + frequency_penalty: Optional[float] = None + logit_bias: Optional[Dict[str, float]] = None + logprobs: Optional[bool] = None + top_logprobs: Optional[int] = None + max_tokens: Optional[int] = None + n: Optional[int] = None + presence_penalty: Optional[float] = None + response_format: Optional[Dict[str, Any]] = None + seed: Optional[int] = None + service_tier: Optional[str] = None + stop: Optional[Union[str, List[str]]] = None + stream_options: Optional[Dict[str, Any]] = None + temperature: Optional[float] = None + top_p: Optional[float] = None + tools: Optional[List[Dict[str, Any]]] = None + tool_choice: Optional[Union[str, Dict[str, Any]]] = None + parallel_tool_calls: Optional[bool] = None + function_call: Optional[Union[str, Dict[str, Any]]] = None + functions: Optional[List[Dict[str, Any]]] = None + user: Optional[str] = None + stream: Optional[bool] = None + + # LiteLLM-specific metadata param (from original ChatCompletionRequest) + metadata: Optional[Dict[str, Any]] = None + # Optional LiteLLM params - guardrails: Optional[List[str]] - caching: Optional[bool] - num_retries: Optional[int] - context_window_fallback_dict: Optional[Dict[str, str]] - fallbacks: Optional[List[str]] + guardrails: Optional[List[str]] = None + caching: Optional[bool] = None + num_retries: Optional[int] = None + context_window_fallback_dict: Optional[Dict[str, str]] = None + fallbacks: Optional[List[str]] = None class ModelInfoDelete(LiteLLMPydanticObjectBase): diff --git a/litellm/proxy/common_utils/custom_openapi_spec.py b/litellm/proxy/common_utils/custom_openapi_spec.py index f2960bac297..5bd8534a9a6 100644 --- a/litellm/proxy/common_utils/custom_openapi_spec.py +++ b/litellm/proxy/common_utils/custom_openapi_spec.py @@ -72,7 +72,8 @@ class CustomOpenAPISpec: @staticmethod def add_request_body_to_paths(openapi_schema: Dict[str, Any], paths: List[str], schema_ref: str) -> None: """ - Add request body schema reference to specified paths. + Add request body with expanded form fields for better Swagger UI display. + This keeps the request body but expands it to show individual fields in the UI. Args: openapi_schema: The OpenAPI schema dict to modify @@ -81,16 +82,99 @@ class CustomOpenAPISpec: """ for path in paths: if path in openapi_schema.get("paths", {}) and "post" in openapi_schema["paths"][path]: + # Get the actual schema to extract ALL field definitions + schema_name = schema_ref.split("/")[-1] # Extract "ProxyChatCompletionRequest" from the ref + actual_schema = openapi_schema.get("components", {}).get("schemas", {}).get(schema_name, {}) + schema_properties = actual_schema.get("properties", {}) + required_fields = actual_schema.get("required", []) + + # Create an expanded inline schema instead of just a $ref + # This makes Swagger UI show all individual fields in the request body editor + expanded_schema = { + "type": "object", + "required": required_fields, + "properties": {} + } + + # Add all properties with their full definitions + for field_name, field_def in schema_properties.items(): + expanded_field = CustomOpenAPISpec._expand_field_definition(field_def) + + # Add a simple example for the messages field + if field_name == "messages": + expanded_field["example"] = [ + {"role": "user", "content": "Hello, how are you?"} + ] + + expanded_schema["properties"][field_name] = expanded_field + + # Include $defs from the original schema to support complex types like AllMessageValues + # This ensures that message types and other complex union types work properly + if "$defs" in actual_schema: + expanded_schema["$defs"] = actual_schema["$defs"] + + # Set the request body with the expanded schema openapi_schema["paths"][path]["post"]["requestBody"] = { "required": True, "content": { "application/json": { - "schema": { - "$ref": schema_ref - } + "schema": expanded_schema } } } + + # Keep any existing parameters (like path parameters) but remove conflicting query params + if "parameters" in openapi_schema["paths"][path]["post"]: + existing_params = openapi_schema["paths"][path]["post"]["parameters"] + # Only keep path parameters, remove query params that conflict with request body + filtered_params = [ + param for param in existing_params + if param.get("in") == "path" + ] + openapi_schema["paths"][path]["post"]["parameters"] = filtered_params + + @staticmethod + def _extract_field_schema(field_def: Dict[str, Any]) -> Dict[str, Any]: + """ + Extract a simple schema from a Pydantic field definition for parameter display. + + Args: + field_def: Pydantic field definition + + Returns: + Simplified schema for OpenAPI parameter + """ + # Handle simple types + if "type" in field_def: + return {"type": field_def["type"]} + + # Handle anyOf (Optional fields in Pydantic v2) + if "anyOf" in field_def: + any_of = field_def["anyOf"] + # Find the non-null type + for option in any_of: + if option.get("type") != "null": + return option + # Fallback to string if all else fails + return {"type": "string"} + + # Default fallback + return {"type": "string"} + + @staticmethod + def _expand_field_definition(field_def: Dict[str, Any]) -> Dict[str, Any]: + """ + Expand a Pydantic field definition for inline use in OpenAPI schema. + This creates a full field definition that Swagger UI can render as individual form fields. + + Args: + field_def: Pydantic field definition + + Returns: + Expanded field definition for OpenAPI schema + """ + # Return the field definition as-is since Pydantic already provides proper schemas + return field_def.copy() @staticmethod def add_request_schema( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index ef9ef3cb332..f29ab7c5882 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -731,6 +731,11 @@ def get_openapi_schema(): } } + # Add LLM API request schema bodies for documentation + from litellm.proxy.common_utils.custom_openapi_spec import CustomOpenAPISpec + + openapi_schema = CustomOpenAPISpec.add_llm_api_request_schema_body(openapi_schema) + app.openapi_schema = openapi_schema return app.openapi_schema @@ -759,6 +764,9 @@ def custom_openapi(): if os.getenv("DOCS_FILTERED", "False") == "True" and premium_user: app.openapi = custom_openapi # type: ignore +else: + # For regular users, use get_openapi_schema to include LLM API schemas + app.openapi = get_openapi_schema # type: ignore class UserAPIKeyCacheTTLEnum(enum.Enum): diff --git a/tests/test_litellm/proxy/test_swagger_chat_completions.py b/tests/test_litellm/proxy/test_swagger_chat_completions.py new file mode 100644 index 00000000000..b973eab6213 --- /dev/null +++ b/tests/test_litellm/proxy/test_swagger_chat_completions.py @@ -0,0 +1,310 @@ +""" +Unit test to validate that /chat/completions has the expected schema in Swagger after add_llm_api_request_schema_body runs. + +This test ensures that the ProxyChatCompletionRequest Pydantic model is properly added to the OpenAPI schema +for the /chat/completions endpoint, showing all expected fields in the Swagger documentation. +""" + +from unittest.mock import Mock, patch + +import pytest +from fastapi.testclient import TestClient + +from litellm.proxy.common_utils.custom_openapi_spec import CustomOpenAPISpec +from litellm.proxy.proxy_server import app + + +class TestSwaggerChatCompletions: + """Test suite for validating /chat/completions schema in Swagger documentation.""" + + @pytest.fixture + def client(self): + """FastAPI test client for the proxy server.""" + return TestClient(app) + + def test_openapi_schema_includes_chat_completions_request_body(self, client): + """ + Test that the OpenAPI schema includes ProxyChatCompletionRequest schema + for /chat/completions endpoints after add_llm_api_request_schema_body runs. + """ + # Clear any cached schema to ensure we get the latest version + from litellm.proxy.proxy_server import app + app.openapi_schema = None + + # Get the OpenAPI schema from the running app + response = client.get("/openapi.json") + assert response.status_code == 200 + + openapi_schema = response.json() + + # Verify the schema has the expected structure + assert "openapi" in openapi_schema + assert "paths" in openapi_schema + assert "components" in openapi_schema + assert "schemas" in openapi_schema["components"] + + # Check that ProxyChatCompletionRequest schema is in components + assert "ProxyChatCompletionRequest" in openapi_schema["components"]["schemas"] + + # Get the ProxyChatCompletionRequest schema + chat_completion_schema = openapi_schema["components"]["schemas"]["ProxyChatCompletionRequest"] + + # Verify it has the expected properties structure + assert "properties" in chat_completion_schema + properties = chat_completion_schema["properties"] + + # Check for core OpenAI chat completion fields + expected_core_fields = [ + "model", + "messages", + "temperature", + "top_p", + "max_tokens", + "stream", + "stop", + "presence_penalty", + "frequency_penalty", + "logit_bias", + "user", + "response_format", + "seed", + "tools", + "tool_choice", + "logprobs", + "top_logprobs" + ] + + for field in expected_core_fields: + assert field in properties, f"Expected field '{field}' not found in ProxyChatCompletionRequest schema" + + # Check for LiteLLM-specific fields added by ProxyChatCompletionRequest + expected_litellm_fields = [ + "guardrails", + "caching", + "num_retries", + "context_window_fallback_dict", + "fallbacks" + ] + + for field in expected_litellm_fields: + assert field in properties, f"Expected LiteLLM field '{field}' not found in ProxyChatCompletionRequest schema" + + # Verify model and messages are required fields + if "required" in chat_completion_schema: + required_fields = chat_completion_schema["required"] + assert "model" in required_fields, "Field 'model' should be required" + assert "messages" in required_fields, "Field 'messages' should be required" + + def test_chat_completions_endpoints_have_expanded_request_body(self, client): + """ + Test that /chat/completions endpoint has an expanded request body schema + with all individual fields visible (not just a $ref). + """ + # Clear any cached schema to ensure we get the latest version + from litellm.proxy.proxy_server import app + app.openapi_schema = None + + # Get the OpenAPI schema + response = client.get("/openapi.json") + assert response.status_code == 200 + + openapi_schema = response.json() + paths = openapi_schema["paths"] + + # Check main chat completion path + path_to_check = "/chat/completions" + assert path_to_check in paths, f"Path {path_to_check} not found in OpenAPI schema" + assert "post" in paths[path_to_check], f"POST method not found for path {path_to_check}" + + post_spec = paths[path_to_check]["post"] + + # Should have request body with expanded schema (not just $ref) + assert "requestBody" in post_spec, f"Path {path_to_check} should have requestBody" + request_body = post_spec["requestBody"] + + # Check request body structure + assert "content" in request_body + assert "application/json" in request_body["content"] + json_content = request_body["content"]["application/json"] + assert "schema" in json_content + + schema_def = json_content["schema"] + + # Should be an expanded object schema, not a $ref + assert schema_def.get("type") == "object", "Schema should be an expanded object type" + assert "properties" in schema_def, "Schema should have expanded properties" + assert "$ref" not in schema_def, "Schema should not be a reference (should be expanded inline)" + + # Should have all Pydantic fields as individual properties + properties = schema_def["properties"] + assert len(properties) >= 25, f"Expected at least 25 properties, got {len(properties)}" + + # Should have core OpenAI fields + core_fields = ["model", "messages", "temperature", "max_tokens", "stream"] + for field in core_fields: + assert field in properties, f"Core field '{field}' should be in expanded properties" + + # Should have LiteLLM-specific fields + litellm_fields = ["guardrails", "caching", "fallbacks", "num_retries"] + for field in litellm_fields: + assert field in properties, f"LiteLLM field '{field}' should be in expanded properties" + + # Check required fields + required_fields = schema_def.get("required", []) + assert "model" in required_fields, "Model should be marked as required" + assert "messages" in required_fields, "Messages should be marked as required" + + # Should have minimal parameters (only path parameters) + parameters = post_spec.get("parameters", []) + # All parameters should be path parameters, no query parameters + for param in parameters: + assert param.get("in") == "path", f"Only path parameters expected, found {param.get('in')} parameter: {param.get('name')}" + + @patch('litellm.proxy.common_utils.custom_openapi_spec.CustomOpenAPISpec.add_chat_completion_request_schema') + def test_add_llm_api_request_schema_body_calls_chat_completion_method(self, mock_add_chat): + """ + Test that add_llm_api_request_schema_body calls add_chat_completion_request_schema. + """ + # Create a mock schema + mock_schema = { + "openapi": "3.0.0", + "info": {"title": "Test API", "version": "1.0.0"}, + "paths": {} + } + + # Configure the mock to return the schema + mock_add_chat.return_value = mock_schema + + # Call the main method + result = CustomOpenAPISpec.add_llm_api_request_schema_body(mock_schema) + + # Verify the chat completion method was called + mock_add_chat.assert_called_once_with(mock_schema) + assert result == mock_schema + + def test_custom_openapi_spec_chat_completion_paths_constant(self): + """ + Test that the CHAT_COMPLETION_PATHS constant includes all expected endpoints. + """ + expected_paths = [ + "/v1/chat/completions", + "/chat/completions", + "/engines/{model}/chat/completions", + "/openai/deployments/{model}/chat/completions" + ] + + assert hasattr(CustomOpenAPISpec, 'CHAT_COMPLETION_PATHS') + actual_paths = CustomOpenAPISpec.CHAT_COMPLETION_PATHS + + for expected_path in expected_paths: + assert expected_path in actual_paths, f"Expected path '{expected_path}' not found in CHAT_COMPLETION_PATHS" + + def test_proxy_chat_completion_request_pydantic_model_works(self): + """ + Test that ProxyChatCompletionRequest properly generates schemas + and includes the expected LiteLLM-specific fields. + """ + from litellm.proxy._types import ProxyChatCompletionRequest + + # Check that we can get the schema + try: + # Try Pydantic v2 method first + schema = ProxyChatCompletionRequest.model_json_schema() + except AttributeError: + try: + # Fallback to Pydantic v1 method + schema = ProxyChatCompletionRequest.schema() + except AttributeError: + pytest.fail("Could not get schema from ProxyChatCompletionRequest using either Pydantic v1 or v2 methods") + + # Verify schema has properties + assert "properties" in schema + properties = schema["properties"] + + # Check for core required fields + assert "model" in properties, "Field 'model' should be in schema" + assert "messages" in properties, "Field 'messages' should be in schema" + + # Check for LiteLLM-specific fields + litellm_fields = ["guardrails", "caching", "num_retries", "context_window_fallback_dict", "fallbacks"] + for field in litellm_fields: + assert field in properties, f"LiteLLM field '{field}' should be in ProxyChatCompletionRequest schema" + + def test_messages_field_has_example(self, client): + """ + Test that the messages field in the expanded request body includes a helpful example. + """ + # Clear any cached schema to ensure we get the latest version + from litellm.proxy.proxy_server import app + app.openapi_schema = None + + # Get the OpenAPI schema + response = client.get("/openapi.json") + assert response.status_code == 200 + + openapi_schema = response.json() + + # Navigate to the chat completions request body schema + chat_completions_post = openapi_schema["paths"]["/chat/completions"]["post"] + request_body = chat_completions_post["requestBody"] + schema_def = request_body["content"]["application/json"]["schema"] + + # Check that messages field has an example + messages_field = schema_def["properties"]["messages"] + assert "example" in messages_field, "Messages field should have an example" + + # Verify the example structure + example = messages_field["example"] + assert isinstance(example, list), "Messages example should be a list" + assert len(example) >= 1, "Messages example should have at least 1 message" + + # Check that example messages have proper structure + for message in example: + assert "role" in message, "Each example message should have a role" + assert "content" in message, "Each example message should have content" + assert message["role"] in ["user", "assistant", "system"], f"Invalid role: {message['role']}" + assert isinstance(message["content"], str), "Message content should be a string" + + def test_request_body_accepts_actual_chat_request(self, client): + """ + Test that the expanded request body schema accepts a real chat completion request. + This ensures our schema modifications don't break actual API functionality. + """ + # Test data that should be valid according to our expanded schema + test_request = { + "model": "gpt-4o", + "messages": [ + {"role": "user", "content": "Hello, how are you?"}, + {"role": "assistant", "content": "I'm doing well, thank you!"} + ], + "temperature": 0.7, + "max_tokens": 100, + "guardrails": ["no-harmful-content"], + "caching": True + } + + # This should validate against our schema without errors + # Note: We're not actually calling the endpoint (which would require API keys) + # but testing that the request structure is accepted by the schema + + # Get the OpenAPI schema to verify our test data matches + response = client.get("/openapi.json") + assert response.status_code == 200 + + openapi_schema = response.json() + chat_completions_post = openapi_schema["paths"]["/chat/completions"]["post"] + + # Should have expanded request body (not just $ref) + assert "requestBody" in chat_completions_post + request_body = chat_completions_post["requestBody"] + schema_def = request_body["content"]["application/json"]["schema"] + + # Verify our test request has fields that exist in the schema + properties = schema_def["properties"] + for field_name in test_request.keys(): + assert field_name in properties, f"Field '{field_name}' should be in expanded schema properties" + + # Verify required fields are present in test request + required_fields = schema_def.get("required", []) + for required_field in required_fields: + assert required_field in test_request, f"Required field '{required_field}' should be in test request" \ No newline at end of file