[Bug Fix] Allow using Swagger for /chat/completions (#13469)

* fix get_openapi_schema

* fixes for ProxyChatCompletionRequest

* TestSwaggerChatCompletions

* fix working request body

* fix - add "messages"

* fix messages

* TestSwaggerChatCompletions

* test_messages_field_has_example

* ruff check fix
This commit is contained in:
Ishaan Jaff 2025-08-09 15:35:45 -07:00 • committed by GitHub
parent 1270df08a4
commit 60306d34a0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 447 additions and 15 deletions

View file

@ -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):

View file

@ -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(

View file

@ -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):

View file

@ -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"