mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
test(interactions): follow operation schemas and declared path parameters
This commit is contained in:
parent
e5098f6080
commit
82ef0b9135
1 changed files with 137 additions and 39 deletions
|
|
@ -9,7 +9,8 @@ Run with: pytest tests/unit/interactions/test_openapi_compliance.py -v
|
|||
|
||||
import json
|
||||
import os
|
||||
from typing import Any, Dict
|
||||
import re
|
||||
from typing import Any, Dict, Final
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
|
|
@ -44,6 +45,56 @@ def _declared_type_value(variant_schema: Dict[str, Any]) -> Any:
|
|||
return type_property.get("const") or (enum_values[0] if len(enum_values) == 1 else None)
|
||||
|
||||
|
||||
def _resolve_local_ref(spec_dict: dict[str, Any], schema: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Resolve component references used by operations, schemas, and parameters."""
|
||||
if "$ref" not in schema:
|
||||
return schema
|
||||
reference: Final = schema["$ref"]
|
||||
assert reference.startswith("#/components/"), f"Expected a local component reference: {reference}"
|
||||
category, name = reference.removeprefix("#/components/").split("/")
|
||||
return spec_dict["components"][category][name.replace("~1", "/").replace("~0", "~")]
|
||||
|
||||
|
||||
def _interaction_operation(
|
||||
spec_dict: dict[str, Any], method: str, *, individual: bool = False
|
||||
) -> tuple[str, dict[str, Any]]:
|
||||
"""Match collection or item routes exactly, independent of placeholder names."""
|
||||
pattern: Final = r"(?:/[^/]+)*/interactions" + (r"/(\{[^/{}]+\})" if individual else "")
|
||||
matches: Final = tuple(
|
||||
(path, path_item, match)
|
||||
for path, path_item in spec_dict["paths"].items()
|
||||
if (match := re.fullmatch(pattern, path)) and method in path_item
|
||||
)
|
||||
assert len(matches) == 1, f"Expected one {method.upper()} interactions endpoint, got {matches}"
|
||||
path, path_item, match = matches[0]
|
||||
operation: Final = path_item[method]
|
||||
if individual:
|
||||
parameter_name: Final = match.group(1)[1:-1]
|
||||
parameters: Final = {
|
||||
(parameter["name"], parameter["in"]): parameter
|
||||
for raw_parameter in (*path_item.get("parameters", ()), *operation.get("parameters", ()))
|
||||
for parameter in (_resolve_local_ref(spec_dict, raw_parameter),)
|
||||
}
|
||||
parameter: Final = parameters.get((parameter_name, "path"))
|
||||
assert parameter is not None, f"{path} must declare its interaction ID path parameter"
|
||||
assert parameter.get("required") is True, f"{path} must require its interaction ID"
|
||||
parameter_schema: Final = _resolve_local_ref(spec_dict, parameter["schema"])
|
||||
assert parameter_schema.get("type") == "string", f"{path} must accept a string interaction ID"
|
||||
return path, operation
|
||||
|
||||
|
||||
def _model_request_schema(spec_dict: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Find the model variant of the JSON body declared by the create operation."""
|
||||
_, operation = _interaction_operation(spec_dict, "post")
|
||||
request_body: Final = _resolve_local_ref(spec_dict, operation["requestBody"])
|
||||
assert request_body.get("required") is True, "Creating an interaction must require a request body"
|
||||
schema: Final = _resolve_local_ref(spec_dict, request_body["content"]["application/json"]["schema"])
|
||||
variants: Final = tuple(_resolve_local_ref(spec_dict, variant) for variant in schema.get("oneOf", (schema,)))
|
||||
model_variants: Final = tuple(variant for variant in variants if "model" in variant.get("properties", {}))
|
||||
assert len(model_variants) == 1, f"Expected one model request variant, got {model_variants}"
|
||||
return model_variants[0]
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def spec_dict() -> Dict[str, Any]:
|
||||
"""Load raw spec dict for manual validation."""
|
||||
|
|
@ -60,12 +111,15 @@ class TestRequestCompliance:
|
|||
"""Tests that our request bodies match the OpenAPI spec."""
|
||||
|
||||
def test_create_model_interaction_request_schema(self, spec_dict):
|
||||
"""Verify CreateModelInteractionParams schema fields."""
|
||||
schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"]
|
||||
"""Verify the model request schema declared by POST /interactions."""
|
||||
schema = _model_request_schema(spec_dict)
|
||||
|
||||
# Required fields per spec
|
||||
assert "model" in schema["required"]
|
||||
assert "input" in schema["required"]
|
||||
for field in ("model", "input"):
|
||||
assert field in schema["properties"]
|
||||
assert schema["properties"][field].get("readOnly") is not True
|
||||
assert _resolve_local_ref(spec_dict, schema["properties"][field]).get("readOnly") is not True
|
||||
|
||||
# Check our supported optional fields exist in spec
|
||||
our_optional_fields = [
|
||||
|
|
@ -88,13 +142,8 @@ class TestRequestCompliance:
|
|||
|
||||
def test_input_types_match_spec(self, spec_dict):
|
||||
"""Verify input field supports string, Content, Content[], Turn[]."""
|
||||
schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"]
|
||||
input_schema = schema["properties"]["input"]
|
||||
|
||||
# The input property may be inline oneOf or a $ref to InteractionsInput
|
||||
if "$ref" in input_schema:
|
||||
ref_name = input_schema["$ref"].split("/")[-1]
|
||||
input_schema = spec_dict["components"]["schemas"][ref_name]
|
||||
schema = _model_request_schema(spec_dict)
|
||||
input_schema = _resolve_local_ref(spec_dict, schema["properties"]["input"])
|
||||
|
||||
# Should be oneOf with multiple types
|
||||
assert "oneOf" in input_schema
|
||||
|
|
@ -295,45 +344,94 @@ class TestEndpointCompliance:
|
|||
|
||||
def test_create_endpoint_exists(self, spec_dict):
|
||||
"""Verify POST /interactions endpoint exists."""
|
||||
paths = spec_dict["paths"]
|
||||
|
||||
# Find the create interactions endpoint
|
||||
create_path = None
|
||||
for path, methods in paths.items():
|
||||
if "interactions" in path and "post" in methods:
|
||||
create_path = path
|
||||
break
|
||||
|
||||
assert create_path is not None, "POST /interactions endpoint not found"
|
||||
create_path, _ = _interaction_operation(spec_dict, "post")
|
||||
print(f"✓ Create endpoint: POST {create_path}")
|
||||
|
||||
def test_get_endpoint_exists(self, spec_dict):
|
||||
"""Verify GET /interactions/{id} endpoint exists."""
|
||||
paths = spec_dict["paths"]
|
||||
|
||||
get_path = None
|
||||
for path, methods in paths.items():
|
||||
if "{id}" in path and "interactions" in path and "get" in methods:
|
||||
get_path = path
|
||||
break
|
||||
|
||||
assert get_path is not None, "GET /interactions/{id} endpoint not found"
|
||||
get_path, _ = _interaction_operation(spec_dict, "get", individual=True)
|
||||
print(f"✓ Get endpoint: GET {get_path}")
|
||||
|
||||
def test_delete_endpoint_exists(self, spec_dict):
|
||||
"""Verify DELETE /interactions/{id} endpoint exists."""
|
||||
paths = spec_dict["paths"]
|
||||
|
||||
delete_path = None
|
||||
for path, methods in paths.items():
|
||||
if "{id}" in path and "interactions" in path and "delete" in methods:
|
||||
delete_path = path
|
||||
break
|
||||
|
||||
assert delete_path is not None, "DELETE /interactions/{id} endpoint not found"
|
||||
delete_path, _ = _interaction_operation(spec_dict, "delete", individual=True)
|
||||
print(f"✓ Delete endpoint: DELETE {delete_path}")
|
||||
|
||||
|
||||
class TestOperationResolution:
|
||||
"""Keep structural resolution strict without depending on generated names."""
|
||||
|
||||
@pytest.mark.parametrize("as_union", [False, True])
|
||||
def test_model_schema_comes_from_create_operation(self, as_union):
|
||||
model_schema: Final = {"properties": {"model": {"type": "string"}}, "required": ["model"]}
|
||||
reference: Final = {"$ref": "#/components/schemas/RenamedModelRequest"}
|
||||
body_schema: Final = (
|
||||
{"oneOf": [{"properties": {"agent": {"type": "string"}}}, reference]} if as_union else reference
|
||||
)
|
||||
spec: Final = {
|
||||
"paths": {
|
||||
"/{version}/interactions": {
|
||||
"post": {
|
||||
"requestBody": {"required": True, "content": {"application/json": {"schema": body_schema}}}
|
||||
}
|
||||
}
|
||||
},
|
||||
"components": {
|
||||
"schemas": {"RenamedModelRequest": model_schema, "CreateModelInteractionParams": {"properties": {}}}
|
||||
},
|
||||
}
|
||||
assert _model_request_schema(spec) is model_schema
|
||||
|
||||
@pytest.mark.parametrize("method,shared", [("get", False), ("delete", True)])
|
||||
def test_item_route_accepts_a_renamed_declared_identifier(self, method, shared):
|
||||
parameter: Final = {"name": "renamedId", "in": "path", "required": True, "schema": {"type": "string"}}
|
||||
parameters: Final = [{"$ref": "#/components/parameters/Identifier"}]
|
||||
operation: Final = {"parameters": [] if shared else parameters}
|
||||
path: Final = "/{version}/interactions/{renamedId}"
|
||||
spec: Final = {
|
||||
"paths": {path: {"parameters": parameters if shared else [], method: operation}},
|
||||
"components": {"parameters": {"Identifier": parameter}},
|
||||
}
|
||||
assert _interaction_operation(spec, method, individual=True) == (path, operation)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"path,parameter,error",
|
||||
[
|
||||
(
|
||||
"/interactions/{id}/cancel",
|
||||
{"required": True, "type": "string"},
|
||||
"Expected one GET interactions endpoint",
|
||||
),
|
||||
(
|
||||
"/other_interactions/{id}",
|
||||
{"required": True, "type": "string"},
|
||||
"Expected one GET interactions endpoint",
|
||||
),
|
||||
("/interactions/{id}", {"required": False, "type": "string"}, "must require its interaction ID"),
|
||||
("/interactions/{id}", {"required": True, "type": "integer"}, "must accept a string interaction ID"),
|
||||
],
|
||||
)
|
||||
def test_item_route_rejects_incompatible_contracts(self, path, parameter, error):
|
||||
spec: Final = {
|
||||
"paths": {
|
||||
path: {
|
||||
"get": {
|
||||
"parameters": [
|
||||
{
|
||||
"name": "id",
|
||||
"in": "path",
|
||||
"required": parameter["required"],
|
||||
"schema": {"type": parameter["type"]},
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
with pytest.raises(AssertionError, match=error):
|
||||
_interaction_operation(spec, "get", individual=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Quick manual test
|
||||
import httpx
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue