diff --git a/tests/unit/interactions/test_openapi_compliance.py b/tests/unit/interactions/test_openapi_compliance.py index d3f1183cea6..8e460cc7c38 100644 --- a/tests/unit/interactions/test_openapi_compliance.py +++ b/tests/unit/interactions/test_openapi_compliance.py @@ -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