diff --git a/tests/unit/interactions/test_openapi_compliance.py b/tests/unit/interactions/test_openapi_compliance.py index d3f1183cea6..433b94f5983 100644 --- a/tests/unit/interactions/test_openapi_compliance.py +++ b/tests/unit/interactions/test_openapi_compliance.py @@ -9,6 +9,7 @@ Run with: pytest tests/unit/interactions/test_openapi_compliance.py -v import json import os +import re from typing import Any, Dict from unittest.mock import MagicMock, patch @@ -44,6 +45,24 @@ 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 _schema_named(spec_dict: dict[str, Any], ref: str) -> dict[str, Any]: + return spec_dict["components"]["schemas"][ref.split("/")[-1]] + + +def _create_request_schema(spec_dict: dict[str, Any]) -> dict[str, Any]: + """The model-interaction body POSTed to /interactions, found by endpoint since Google renames + the schema and now wraps it in a oneOf with the agent variant.""" + for path, methods in spec_dict["paths"].items(): + if path.endswith("/interactions") and "post" in methods: + body = methods["post"]["requestBody"]["content"]["application/json"]["schema"] + candidates = body.get("oneOf") or [body] + for candidate in candidates: + schema = _schema_named(spec_dict, candidate["$ref"]) if "$ref" in candidate else candidate + if "model" in schema.get("properties", {}): + return schema + raise AssertionError("no model-interaction request body found on POST /interactions") + + @pytest.fixture(scope="module") def spec_dict() -> Dict[str, Any]: """Load raw spec dict for manual validation.""" @@ -56,16 +75,22 @@ def openapi_spec(spec_dict: Dict[str, Any]) -> OpenAPI: return OpenAPI.from_dict(spec_dict) +def _interaction_item_path(paths: dict[str, Any], method: str) -> Any: + """`/…/interactions/{param}` for a method; the parameter's name is cosmetic and has changed.""" + item = re.compile(r"/interactions/\{[^}/]+\}$") + return next((p for p, methods in paths.items() if item.search(p) and method in methods), None) + + 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-interaction request body fields.""" + schema = _create_request_schema(spec_dict) - # Required fields per spec + # `model` is required; `input` is optional now that a call can resume via previous_interaction_id assert "model" in schema["required"] - assert "input" in schema["required"] + assert "input" in schema["properties"] # Check our supported optional fields exist in spec our_optional_fields = [ @@ -88,8 +113,7 @@ 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"] + input_schema = _create_request_schema(spec_dict)["properties"]["input"] # The input property may be inline oneOf or a $ref to InteractionsInput if "$ref" in input_schema: @@ -311,11 +335,7 @@ class TestEndpointCompliance: """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 + get_path = _interaction_item_path(paths, "get") assert get_path is not None, "GET /interactions/{id} endpoint not found" print(f"✓ Get endpoint: GET {get_path}") @@ -324,11 +344,7 @@ class TestEndpointCompliance: """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 + delete_path = _interaction_item_path(paths, "delete") assert delete_path is not None, "DELETE /interactions/{id} endpoint not found" print(f"✓ Delete endpoint: DELETE {delete_path}")