mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
[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:
parent
1270df08a4
commit
60306d34a0
4 changed files with 447 additions and 15 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
310
tests/test_litellm/proxy/test_swagger_chat_completions.py
Normal file
310
tests/test_litellm/proxy/test_swagger_chat_completions.py
Normal 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"
|
||||
Loading…
Add table
Reference in a new issue