From d9609c834b0afe48033443646dbb85683d8bc6db Mon Sep 17 00:00:00 2001 From: reachsridhard Date: Sun, 26 Jul 2026 07:08:16 +0530 Subject: [PATCH 01/15] feat: replace check_expected_keys with typed pydantic models (#1034) --- pyproject.toml | 2 +- tavern/_core/pydantic_models.py | 163 +++++++++++++ tavern/_plugins/grpc/client.py | 18 +- tavern/_plugins/grpc/request.py | 7 +- tavern/_plugins/grpc/response.py | 4 +- tavern/_plugins/mqtt/client.py | 53 ++--- tavern/_plugins/mqtt/request.py | 7 +- tavern/_plugins/rest/request.py | 24 +- tests/unit/test_pydantic_models.py | 353 +++++++++++++++++++++++++++++ 9 files changed, 550 insertions(+), 81 deletions(-) create mode 100644 tavern/_core/pydantic_models.py create mode 100644 tests/unit/test_pydantic_models.py diff --git a/pyproject.toml b/pyproject.toml index 38fec1ba3..72cfb1d3a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -27,6 +27,7 @@ dependencies = [ "PyYAML>=6.0.1,<7", "jmespath>=1,<2", "jsonschema>=4,<5", + "pydantic>=2,<3", "pyjwt>=2.5.0,<3", "pykwalify>=1.8.0,<2", "pytest>=8,<10", @@ -108,7 +109,6 @@ dev = [ "tox-uv>=1.28.0", "pytest-asyncio>=1.3.0", "hypothesis>=6,<7", - "pydantic", "flask-httpauth>=4.8.1,<6", ] diff --git a/tavern/_core/pydantic_models.py b/tavern/_core/pydantic_models.py new file mode 100644 index 000000000..79eb0c3b5 --- /dev/null +++ b/tavern/_core/pydantic_models.py @@ -0,0 +1,163 @@ +"""Pydantic models for validating request/client specs. + +Replaces the older ``check_expected_keys`` pattern with pydantic models +that use ``extra="forbid"`` to reject unexpected keys, providing the same +validation with better error messages and type safety. +""" + +from collections.abc import Mapping +from typing import Optional, Union + +from pydantic import BaseModel, ConfigDict, Field, ValidationError + +from tavern._core import exceptions + +# Type alias for JSON-compatible values (any valid JSON type) +JSONType = Union[dict, list, str, int, float, bool, None] + + +class _BaseKeyValidator(BaseModel): + """Base model that forbids extra keys and raises UnexpectedKeysError on validation failure.""" + + model_config = ConfigDict( + extra="forbid", arbitrary_types_allowed=True, populate_by_name=True + ) + + @classmethod + def validate_keys(cls, data: Mapping) -> dict: + """Validate that ``data`` contains only expected keys and types. + + Args: + data: Dictionary to validate against this model's fields. + + Returns: + The validated data as a dict. + + Raises: + exceptions.UnexpectedKeysError: If unexpected keys are present or + a value has an invalid type. + """ + try: + return cls(**dict(data)).model_dump(exclude_unset=True, by_alias=True) + except ValidationError as e: + # Extract unexpected field names from the error + unexpected = set() + for err in e.errors(): + if err["type"] == "extra_forbidden": + unexpected.add(err["loc"][-1]) + if unexpected: + msg = f"Unexpected keys {unexpected}" + else: + msg = str(e) + raise exceptions.UnexpectedKeysError(msg) from e + + +# --- REST request spec --- +class RestRequestSpec(_BaseKeyValidator): + method: Optional[str] = None + url: Optional[str] = None + headers: Optional[dict] = None + data: Optional[Union[dict, list, str, bytes]] = None + params: Optional[dict] = None + auth: Optional[Union[list, str]] = None + json_body: Optional[JSONType] = Field(default=None, alias="json") + verify: Optional[Union[bool, str]] = None + files: Optional[Union[dict, list]] = None + file_body: Optional[str] = None + stream: Optional[bool] = None + timeout: Optional[Union[float, list]] = None + cookies: Optional[dict] = None + cert: Optional[Union[str, list]] = None + follow_redirects: Optional[bool] = None + + +# --- MQTT request spec --- +class MQTTRequestSpec(_BaseKeyValidator): + topic: Optional[str] = None + payload: Optional[Union[str, bytes, int, float]] = None + json_body: Optional[JSONType] = Field(default=None, alias="json") + qos: Optional[int] = None + retain: Optional[bool] = None + + +# --- MQTT client config blocks --- +class MQTTClientArgs(_BaseKeyValidator): + client_id: Optional[str] = None + clean_session: Optional[bool] = None + transport: Optional[str] = None + + +class MQTTConnectArgs(_BaseKeyValidator): + host: Optional[str] = None + port: Optional[int] = None + keepalive: Optional[int] = None + timeout: Optional[Union[int, float]] = None + + +class MQTTAuthArgs(_BaseKeyValidator): + username: Optional[str] = None + password: Optional[str] = None + + +class MQTTTLSArgs(_BaseKeyValidator): + enable: Optional[bool] = None + ca_certs: Optional[str] = None + cert_reqs: Optional[str] = None + certfile: Optional[str] = None + keyfile: Optional[str] = None + tls_version: Optional[str] = None + ciphers: Optional[str] = None + + +class MQTTSSLContextArgs(_BaseKeyValidator): + ca_certs: Optional[str] = None + certfile: Optional[str] = None + keyfile: Optional[str] = None + password: Optional[str] = None + tls_version: Optional[str] = None + ciphers: Optional[str] = None + alpn_protocols: Optional[list[str]] = None + + +class MQTTClientTopLevel(_BaseKeyValidator): + client: Optional[dict] = None + connect: Optional[dict] = None + tls: Optional[dict] = None + auth: Optional[dict] = None + ssl_context: Optional[dict] = None + + +# --- gRPC request spec --- +class GRPCRequestSpec(_BaseKeyValidator): + host: Optional[str] = None + service: Optional[str] = None + body: Optional[Union[dict, str]] = None + + +# --- gRPC response spec --- +class GRPCResponseSpec(_BaseKeyValidator): + body: Optional[dict] = None + status: Optional[Union[str, int, list[str], list[int]]] = None + details: Optional[str] = None + save: Optional[dict] = None + + +# --- gRPC client config blocks --- +class GRPCConnectArgs(_BaseKeyValidator): + host: Optional[str] = None + port: Optional[int] = None + options: Optional[dict] = None + timeout: Optional[int] = None + secure: Optional[bool] = None + + +class GRPCProtoArgs(_BaseKeyValidator): + source: Optional[str] = None + module: Optional[str] = None + + +class GRPCClientTopLevel(_BaseKeyValidator): + connect: Optional[dict] = None + proto: Optional[dict] = None + metadata: Optional[dict] = None + attempt_reflection: Optional[bool] = None diff --git a/tavern/_plugins/grpc/client.py b/tavern/_plugins/grpc/client.py index 5720a1ccb..b65a2d5af 100644 --- a/tavern/_plugins/grpc/client.py +++ b/tavern/_plugins/grpc/client.py @@ -19,7 +19,11 @@ from grpc_status import rpc_status from tavern._core import exceptions -from tavern._core.dict_util import check_expected_keys +from tavern._core.pydantic_models import ( + GRPCClientTopLevel, + GRPCConnectArgs, + GRPCProtoArgs, +) from tavern._plugins.grpc.protos import _generate_proto_import, _import_grpc_module logger: logging.Logger = logging.getLogger(__name__) @@ -41,23 +45,17 @@ class _ChannelVals: class GRPCClient: def __init__(self, **kwargs) -> None: logger.debug("Initialising GRPC client with %s", kwargs) - expected_blocks = { - "connect": {"host", "port", "options", "timeout", "secure"}, - "proto": {"source", "module"}, - "metadata": {}, - "attempt_reflection": {}, - } # check main block first - check_expected_keys(expected_blocks.keys(), kwargs) + GRPCClientTopLevel.validate_keys(kwargs) _connect_args = kwargs.pop("connect", {}) - check_expected_keys(expected_blocks["connect"], _connect_args) + GRPCConnectArgs.validate_keys(_connect_args) metadata = kwargs.pop("metadata", {}) self._metadata = list(metadata.items()) _proto_args = kwargs.pop("proto", {}) - check_expected_keys(expected_blocks["proto"], _proto_args) + GRPCProtoArgs.validate_keys(_proto_args) self._attempt_reflection = bool(kwargs.pop("attempt_reflection", False)) diff --git a/tavern/_plugins/grpc/request.py b/tavern/_plugins/grpc/request.py index 9fdfa9f11..bd84dc36a 100644 --- a/tavern/_plugins/grpc/request.py +++ b/tavern/_plugins/grpc/request.py @@ -7,7 +7,8 @@ from box import Box from tavern._core import exceptions -from tavern._core.dict_util import check_expected_keys, format_keys +from tavern._core.dict_util import format_keys +from tavern._core.pydantic_models import GRPCRequestSpec from tavern._core.pytest.config import TestConfig from tavern._plugins.grpc.client import GRPCClient from tavern.request import BaseRequest @@ -48,9 +49,7 @@ class GRPCRequest(BaseRequest): def __init__( self, client: GRPCClient, request_spec: dict, test_block_config: TestConfig ) -> None: - expected = {"host", "service", "body"} - - check_expected_keys(expected, request_spec) + GRPCRequestSpec.validate_keys(request_spec) grpc_args = get_grpc_args(request_spec, test_block_config) diff --git a/tavern/_plugins/grpc/response.py b/tavern/_plugins/grpc/response.py index 6df716875..c6176ad81 100644 --- a/tavern/_plugins/grpc/response.py +++ b/tavern/_plugins/grpc/response.py @@ -7,8 +7,8 @@ from google.protobuf import json_format from tavern._core import exceptions -from tavern._core.dict_util import check_expected_keys from tavern._core.exceptions import TestFailError +from tavern._core.pydantic_models import GRPCResponseSpec from tavern._core.pytest.config import TestConfig from tavern._core.schema.extensions import to_grpc_status from tavern._plugins.grpc.client import GRPCClient @@ -50,7 +50,7 @@ def __init__( expected: _GRPCExpected | Mapping, test_block_config: TestConfig, ) -> None: - check_expected_keys({"body", "status", "details", "save"}, expected) + GRPCResponseSpec.validate_keys(expected) super().__init__( name, expected, diff --git a/tavern/_plugins/mqtt/client.py b/tavern/_plugins/mqtt/client.py index fd9176702..044ff5c85 100644 --- a/tavern/_plugins/mqtt/client.py +++ b/tavern/_plugins/mqtt/client.py @@ -12,7 +12,14 @@ from paho.mqtt.client import MQTTMessageInfo from tavern._core import exceptions -from tavern._core.dict_util import check_expected_keys +from tavern._core.pydantic_models import ( + MQTTAuthArgs, + MQTTClientArgs, + MQTTClientTopLevel, + MQTTConnectArgs, + MQTTSSLContextArgs, + MQTTTLSArgs, +) # MQTT error values _err_vals = { @@ -121,38 +128,6 @@ def _check_and_update_common_tls_args( class MQTTClient: def __init__(self, **kwargs) -> None: - expected_blocks = { - "client": { - "client_id", - "clean_session", - # Can't really use this easily... - # "userdata", - # Force mqttv311 - fix if this becomes an issue - # "protocol", - "transport", - }, - "connect": {"host", "port", "keepalive", "timeout"}, - "tls": { - "enable", - "ca_certs", - "cert_reqs", - "certfile", - "keyfile", - "tls_version", - "ciphers", - }, - "auth": {"username", "password"}, - "ssl_context": { - "ca_certs", - "certfile", - "keyfile", - "password", - "tls_version", - "ciphers", - "alpn_protocols", - }, - } - sanitised_kwargs = copy.deepcopy(kwargs) if auth := kwargs.get("auth"): if "password" in auth: @@ -161,17 +136,17 @@ def __init__(self, **kwargs) -> None: logger.debug("Initialising MQTT client with %s", sanitised_kwargs) # check main block first - check_expected_keys(expected_blocks.keys(), kwargs) + MQTTClientTopLevel.validate_keys(kwargs) # then check constructor/connect/tls_set args self._client_args = kwargs.pop("client", {}) - check_expected_keys(expected_blocks["client"], self._client_args) + MQTTClientArgs.validate_keys(self._client_args) self._connect_args = kwargs.pop("connect", {}) - check_expected_keys(expected_blocks["connect"], self._connect_args) + MQTTConnectArgs.validate_keys(self._connect_args) self._auth_args = kwargs.pop("auth", {}) - check_expected_keys(expected_blocks["auth"], self._auth_args) + MQTTAuthArgs.validate_keys(self._auth_args) if "host" not in self._connect_args: msg = "Need 'host' in 'connect' block for mqtt" @@ -189,12 +164,12 @@ def __init__(self, **kwargs) -> None: ) raise exceptions.MQTTTLSError(msg) - check_expected_keys(expected_blocks["tls"], file_tls_args) + MQTTTLSArgs.validate_keys(file_tls_args) self._tls_args = _handle_tls_args(file_tls_args) logger.debug("TLS is %s", "enabled" if self._tls_args else "disabled") # If there is any SSL kwarg, enable tls through the SSL context - check_expected_keys(expected_blocks["ssl_context"], file_ssl_context_args) + MQTTSSLContextArgs.validate_keys(file_ssl_context_args) self._ssl_context_args = _handle_ssl_context_args(file_ssl_context_args) logger.debug("Paho client args: %s", self._client_args) diff --git a/tavern/_plugins/mqtt/request.py b/tavern/_plugins/mqtt/request.py index 57630be50..1b57a9c16 100644 --- a/tavern/_plugins/mqtt/request.py +++ b/tavern/_plugins/mqtt/request.py @@ -5,8 +5,9 @@ from box.box import Box from tavern._core import exceptions -from tavern._core.dict_util import check_expected_keys, format_keys +from tavern._core.dict_util import format_keys from tavern._core.extfunctions import update_from_ext +from tavern._core.pydantic_models import MQTTRequestSpec from tavern._core.pytest.config import TestConfig from tavern._core.report import attach_yaml from tavern._plugins.mqtt.client import MQTTClient @@ -42,9 +43,7 @@ class MQTTRequest(BaseRequest): def __init__( self, client: MQTTClient, rspec: dict, test_block_config: TestConfig ) -> None: - expected = {"topic", "payload", "json", "qos", "retain"} - - check_expected_keys(expected, rspec) + MQTTRequestSpec.validate_keys(rspec) publish_args = get_publish_args(rspec, test_block_config) diff --git a/tavern/_plugins/rest/request.py b/tavern/_plugins/rest/request.py index 65607b83b..0b0eada70 100644 --- a/tavern/_plugins/rest/request.py +++ b/tavern/_plugins/rest/request.py @@ -12,7 +12,7 @@ from box.box import Box from tavern._core import exceptions -from tavern._core.dict_util import check_expected_keys, deep_dict_merge, format_keys +from tavern._core.dict_util import deep_dict_merge, format_keys from tavern._core.extfunctions import update_from_ext from tavern._core.files import ( _find_file_in_include_path, @@ -21,6 +21,7 @@ guess_filespec, ) from tavern._core.general import valid_http_methods +from tavern._core.pydantic_models import RestRequestSpec from tavern._core.pytest.config import TestConfig from tavern._core.report import attach_yaml from tavern.request import BaseRequest @@ -423,26 +424,7 @@ def __init__( if rspec.pop("clear_session_cookies", False): session.cookies.clear_session_cookies() - expected = { - "method", - "url", - "headers", - "data", - "params", - "auth", - "json", - "verify", - "files", - "file_body", - "stream", - "timeout", - "cookies", - "cert", - # "hooks", - "follow_redirects", - } - - check_expected_keys(expected, rspec) + RestRequestSpec.validate_keys(rspec) request_args = get_request_args(rspec, test_block_config) update_from_ext( diff --git a/tests/unit/test_pydantic_models.py b/tests/unit/test_pydantic_models.py new file mode 100644 index 000000000..9b5f40237 --- /dev/null +++ b/tests/unit/test_pydantic_models.py @@ -0,0 +1,353 @@ +"""Tests for pydantic-based key validation models.""" + +import pytest + +from tavern._core import exceptions +from tavern._core.pydantic_models import ( + GRPCClientTopLevel, + GRPCConnectArgs, + GRPCProtoArgs, + GRPCRequestSpec, + GRPCResponseSpec, + MQTTAuthArgs, + MQTTClientArgs, + MQTTClientTopLevel, + MQTTConnectArgs, + MQTTRequestSpec, + MQTTSSLContextArgs, + MQTTTLSArgs, + RestRequestSpec, +) + + +class TestRestRequestSpec: + def test_valid_keys(self): + data = {"method": "GET", "url": "http://example.com", "json": {"a": 1}} + result = RestRequestSpec.validate_keys(data) + assert "method" in result + assert "url" in result + assert "json" in result + + def test_unexpected_key(self): + data = {"method": "GET", "url": "http://example.com", "bad_key": "value"} + with pytest.raises(exceptions.UnexpectedKeysError): + RestRequestSpec.validate_keys(data) + + def test_empty_dict(self): + result = RestRequestSpec.validate_keys({}) + assert result == {} + + +class TestMQTTRequestSpec: + def test_valid_keys(self): + data = {"topic": "test/topic", "payload": "hello", "qos": 1} + result = MQTTRequestSpec.validate_keys(data) + assert "topic" in result + + def test_unexpected_key(self): + data = {"topic": "test/topic", "bad_key": "value"} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTRequestSpec.validate_keys(data) + + +class TestMQTTClientSpecs: + def test_top_level_valid(self): + data = {"client": {}, "connect": {"host": "localhost"}, "auth": {}} + result = MQTTClientTopLevel.validate_keys(data) + assert "client" in result + + def test_top_level_unexpected(self): + data = {"client": {}, "bad_block": {}} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTClientTopLevel.validate_keys(data) + + def test_connect_args_valid(self): + data = {"host": "localhost", "port": 1883, "keepalive": 60} + result = MQTTConnectArgs.validate_keys(data) + assert "host" in result + + def test_connect_args_unexpected(self): + data = {"host": "localhost", "bad_key": "value"} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTConnectArgs.validate_keys(data) + + def test_client_args_valid(self): + data = {"client_id": "test_id", "transport": "tcp"} + result = MQTTClientArgs.validate_keys(data) + assert "client_id" in result + + def test_client_args_unexpected(self): + data = {"client_id": "test_id", "bad_key": "value"} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTClientArgs.validate_keys(data) + + def test_auth_args_valid(self): + data = {"username": "user", "password": "pass"} + result = MQTTAuthArgs.validate_keys(data) + assert "username" in result + + def test_auth_args_unexpected(self): + data = {"username": "user", "bad_key": "value"} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTAuthArgs.validate_keys(data) + + def test_tls_args_valid(self): + data = {"enable": True, "ca_certs": "/path/to/ca"} + result = MQTTTLSArgs.validate_keys(data) + assert "enable" in result + + def test_tls_args_unexpected(self): + data = {"enable": True, "bad_key": "value"} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTTLSArgs.validate_keys(data) + + def test_ssl_context_args_valid(self): + data = {"ca_certs": "/path/to/ca", "alpn_protocols": ["h2"]} + result = MQTTSSLContextArgs.validate_keys(data) + assert "ca_certs" in result + + def test_ssl_context_args_unexpected(self): + data = {"ca_certs": "/path/to/ca", "bad_key": "value"} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTSSLContextArgs.validate_keys(data) + + +class TestGRPCSpecs: + def test_request_spec_valid(self): + data = {"host": "localhost:50051", "service": "MyService/Method", "body": {}} + result = GRPCRequestSpec.validate_keys(data) + assert "host" in result + + def test_request_spec_unexpected(self): + data = {"host": "localhost:50051", "bad_key": "value"} + with pytest.raises(exceptions.UnexpectedKeysError): + GRPCRequestSpec.validate_keys(data) + + def test_response_spec_valid(self): + data = {"body": {}, "status": 0, "details": "ok", "save": {}} + result = GRPCResponseSpec.validate_keys(data) + assert "body" in result + + def test_response_spec_unexpected(self): + data = {"body": {}, "bad_key": "value"} + with pytest.raises(exceptions.UnexpectedKeysError): + GRPCResponseSpec.validate_keys(data) + + def test_client_top_level_valid(self): + data = {"connect": {"host": "localhost"}, "proto": {"source": "test.proto"}} + result = GRPCClientTopLevel.validate_keys(data) + assert "connect" in result + + def test_client_top_level_unexpected(self): + data = {"connect": {}, "bad_block": {}} + with pytest.raises(exceptions.UnexpectedKeysError): + GRPCClientTopLevel.validate_keys(data) + + def test_connect_args_valid(self): + data = {"host": "localhost", "port": 50051, "timeout": 5, "secure": False} + result = GRPCConnectArgs.validate_keys(data) + assert "host" in result + + def test_connect_args_unexpected(self): + data = {"host": "localhost", "bad_key": "value"} + with pytest.raises(exceptions.UnexpectedKeysError): + GRPCConnectArgs.validate_keys(data) + + def test_proto_args_valid(self): + data = {"source": "test.proto"} + result = GRPCProtoArgs.validate_keys(data) + assert "source" in result + + def test_proto_args_unexpected(self): + data = {"source": "test.proto", "bad_key": "value"} + with pytest.raises(exceptions.UnexpectedKeysError): + GRPCProtoArgs.validate_keys(data) + + +class TestTypeValidation: + """Tests verifying that pydantic models enforce type checking, not just key validation.""" + + def test_rest_method_must_be_string(self): + data = {"method": 123} + with pytest.raises(exceptions.UnexpectedKeysError): + RestRequestSpec.validate_keys(data) + + def test_rest_stream_must_be_bool(self): + data = {"url": "http://example.com", "stream": "not_a_bool"} + with pytest.raises(exceptions.UnexpectedKeysError): + RestRequestSpec.validate_keys(data) + + def test_rest_verify_can_be_bool(self): + data = {"url": "http://example.com", "verify": False} + result = RestRequestSpec.validate_keys(data) + assert result["verify"] is False + + def test_rest_verify_can_be_string(self): + data = {"url": "http://example.com", "verify": "/path/to/ca.pem"} + result = RestRequestSpec.validate_keys(data) + assert result["verify"] == "/path/to/ca.pem" + + def test_rest_json_can_be_dict(self): + data = {"url": "http://example.com", "json": {"key": "value"}} + result = RestRequestSpec.validate_keys(data) + assert result["json"] == {"key": "value"} + + def test_rest_json_can_be_list(self): + data = {"url": "http://example.com", "json": [1, 2, 3]} + result = RestRequestSpec.validate_keys(data) + assert result["json"] == [1, 2, 3] + + def test_rest_json_can_be_string(self): + data = {"url": "http://example.com", "json": "hello"} + result = RestRequestSpec.validate_keys(data) + assert result["json"] == "hello" + + def test_rest_json_can_be_int(self): + data = {"url": "http://example.com", "json": 42} + result = RestRequestSpec.validate_keys(data) + assert result["json"] == 42 + + def test_rest_json_can_be_bool(self): + data = {"url": "http://example.com", "json": True} + result = RestRequestSpec.validate_keys(data) + assert result["json"] is True + + def test_rest_headers_must_be_dict(self): + data = {"url": "http://example.com", "headers": "not_a_dict"} + with pytest.raises(exceptions.UnexpectedKeysError): + RestRequestSpec.validate_keys(data) + + def test_rest_timeout_can_be_float(self): + data = {"url": "http://example.com", "timeout": 30.0} + result = RestRequestSpec.validate_keys(data) + assert result["timeout"] == 30.0 + + def test_rest_timeout_can_be_list(self): + data = {"url": "http://example.com", "timeout": [5.0, 30.0]} + result = RestRequestSpec.validate_keys(data) + assert result["timeout"] == [5.0, 30.0] + + def test_mqtt_topic_must_be_string(self): + data = {"topic": 123} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTRequestSpec.validate_keys(data) + + def test_mqtt_qos_must_be_int(self): + data = {"topic": "test/topic", "qos": "not_an_int"} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTRequestSpec.validate_keys(data) + + def test_mqtt_qos_accepts_int(self): + data = {"topic": "test/topic", "qos": 1} + result = MQTTRequestSpec.validate_keys(data) + assert result["qos"] == 1 + + def test_mqtt_retain_must_be_bool(self): + data = {"topic": "test/topic", "retain": "not_a_bool"} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTRequestSpec.validate_keys(data) + + def test_mqtt_retain_accepts_bool(self): + data = {"topic": "test/topic", "retain": True} + result = MQTTRequestSpec.validate_keys(data) + assert result["retain"] is True + + def test_mqtt_connect_host_must_be_string(self): + data = {"host": 123} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTConnectArgs.validate_keys(data) + + def test_mqtt_connect_port_must_be_int(self): + data = {"host": "localhost", "port": "not_an_int"} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTConnectArgs.validate_keys(data) + + def test_mqtt_connect_port_accepts_int(self): + data = {"host": "localhost", "port": 1883} + result = MQTTConnectArgs.validate_keys(data) + assert result["port"] == 1883 + + def test_mqtt_tls_enable_must_be_bool(self): + data = {"enable": "not_a_bool"} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTTLSArgs.validate_keys(data) + + def test_mqtt_tls_enable_accepts_bool(self): + data = {"enable": True} + result = MQTTTLSArgs.validate_keys(data) + assert result["enable"] is True + + def test_mqtt_ssl_alpn_protocols_must_be_list(self): + data = {"alpn_protocols": "h2"} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTSSLContextArgs.validate_keys(data) + + def test_mqtt_ssl_alpn_protocols_accepts_list(self): + data = {"alpn_protocols": ["h2", "http/1.1"]} + result = MQTTSSLContextArgs.validate_keys(data) + assert result["alpn_protocols"] == ["h2", "http/1.1"] + + def test_mqtt_client_top_level_blocks_must_be_dict(self): + data = {"client": "not_a_dict"} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTClientTopLevel.validate_keys(data) + + def test_grpc_request_host_must_be_string(self): + data = {"host": 123} + with pytest.raises(exceptions.UnexpectedKeysError): + GRPCRequestSpec.validate_keys(data) + + def test_grpc_request_body_can_be_dict(self): + data = {"host": "localhost:50051", "service": "MyService/Method", "body": {"key": "value"}} + result = GRPCRequestSpec.validate_keys(data) + assert result["body"] == {"key": "value"} + + def test_grpc_request_body_can_be_string(self): + data = {"host": "localhost:50051", "service": "MyService/Method", "body": "raw string"} + result = GRPCRequestSpec.validate_keys(data) + assert result["body"] == "raw string" + + def test_grpc_response_status_can_be_int(self): + data = {"status": 0} + result = GRPCResponseSpec.validate_keys(data) + assert result["status"] == 0 + + def test_grpc_response_status_can_be_string(self): + data = {"status": "OK"} + result = GRPCResponseSpec.validate_keys(data) + assert result["status"] == "OK" + + def test_grpc_response_status_can_be_list_of_strings(self): + data = {"status": ["OK", "CANCELLED"]} + result = GRPCResponseSpec.validate_keys(data) + assert result["status"] == ["OK", "CANCELLED"] + + def test_grpc_response_status_can_be_list_of_ints(self): + data = {"status": [0, 1]} + result = GRPCResponseSpec.validate_keys(data) + assert result["status"] == [0, 1] + + def test_grpc_connect_port_must_be_int(self): + data = {"host": "localhost", "port": "not_an_int"} + with pytest.raises(exceptions.UnexpectedKeysError): + GRPCConnectArgs.validate_keys(data) + + def test_grpc_connect_secure_must_be_bool(self): + data = {"host": "localhost", "secure": "not_a_bool"} + with pytest.raises(exceptions.UnexpectedKeysError): + GRPCConnectArgs.validate_keys(data) + + def test_grpc_connect_secure_accepts_bool(self): + data = {"host": "localhost", "secure": True} + result = GRPCConnectArgs.validate_keys(data) + assert result["secure"] is True + + def test_grpc_client_attempt_reflection_must_be_bool(self): + data = {"attempt_reflection": "not_a_bool"} + with pytest.raises(exceptions.UnexpectedKeysError): + GRPCClientTopLevel.validate_keys(data) + + def test_grpc_client_attempt_reflection_accepts_bool(self): + data = {"attempt_reflection": True} + result = GRPCClientTopLevel.validate_keys(data) + assert result["attempt_reflection"] is True From 8da3e878aaa0d1477e3352e17c4be42e50f81a61 Mon Sep 17 00:00:00 2001 From: reachsridhard Date: Sun, 26 Jul 2026 13:02:42 +0530 Subject: [PATCH 02/15] fix: ruff format and uv-lock pre-commit fixes --- tests/unit/test_pydantic_models.py | 12 ++++++++++-- uv.lock | 4 ++-- 2 files changed, 12 insertions(+), 4 deletions(-) diff --git a/tests/unit/test_pydantic_models.py b/tests/unit/test_pydantic_models.py index 9b5f40237..bf9085aa0 100644 --- a/tests/unit/test_pydantic_models.py +++ b/tests/unit/test_pydantic_models.py @@ -298,12 +298,20 @@ def test_grpc_request_host_must_be_string(self): GRPCRequestSpec.validate_keys(data) def test_grpc_request_body_can_be_dict(self): - data = {"host": "localhost:50051", "service": "MyService/Method", "body": {"key": "value"}} + data = { + "host": "localhost:50051", + "service": "MyService/Method", + "body": {"key": "value"}, + } result = GRPCRequestSpec.validate_keys(data) assert result["body"] == {"key": "value"} def test_grpc_request_body_can_be_string(self): - data = {"host": "localhost:50051", "service": "MyService/Method", "body": "raw string"} + data = { + "host": "localhost:50051", + "service": "MyService/Method", + "body": "raw string", + } result = GRPCRequestSpec.validate_keys(data) assert result["body"] == "raw string" diff --git a/uv.lock b/uv.lock index 4388721c4..d232b4103 100644 --- a/uv.lock +++ b/uv.lock @@ -2667,6 +2667,7 @@ source = { editable = "." } dependencies = [ { name = "jmespath" }, { name = "jsonschema" }, + { name = "pydantic" }, { name = "pyjwt" }, { name = "pykwalify" }, { name = "pytest" }, @@ -2713,7 +2714,6 @@ dev = [ { name = "pre-commit" }, { name = "protobuf-protoc-bin" }, { name = "py" }, - { name = "pydantic" }, { name = "pytest-asyncio" }, { name = "pytest-cov" }, { name = "pytest-xdist" }, @@ -2745,6 +2745,7 @@ requires-dist = [ { name = "paho-mqtt", marker = "extra == 'mqtt'", specifier = ">=1.3.1,<=1.6.1" }, { name = "proto-plus", marker = "extra == 'grpc'" }, { name = "protobuf", marker = "extra == 'grpc'", specifier = ">=5,<6" }, + { name = "pydantic", specifier = ">=2,<3" }, { name = "pyjwt", specifier = ">=2.5.0,<3" }, { name = "pykwalify", specifier = ">=1.8.0,<2" }, { name = "pytest", specifier = ">=8,<10" }, @@ -2775,7 +2776,6 @@ dev = [ { name = "pre-commit" }, { name = "protobuf-protoc-bin", specifier = "==29.5" }, { name = "py" }, - { name = "pydantic" }, { name = "pytest-asyncio", specifier = ">=1.3.0" }, { name = "pytest-cov" }, { name = "pytest-xdist" }, From 6be5f9b582d1fce2f9544694394dadc08f9a9079 Mon Sep 17 00:00:00 2001 From: reachsridhard Date: Sun, 26 Jul 2026 13:23:50 +0530 Subject: [PATCH 03/15] fix: handle pre-resolution YAML tags and $ext in pydantic types TypeConvertToken objects (from !force_format_include, !int, etc.) and dict values (from $ext function calls) exist at validation time before resolution. Add these to Union types so pydantic accepts them while still rejecting clearly wrong types (e.g. int for headers, list for method). --- tavern/_core/pydantic_models.py | 137 +++++++++++++++++--------------- 1 file changed, 71 insertions(+), 66 deletions(-) diff --git a/tavern/_core/pydantic_models.py b/tavern/_core/pydantic_models.py index 79eb0c3b5..001ae6c0b 100644 --- a/tavern/_core/pydantic_models.py +++ b/tavern/_core/pydantic_models.py @@ -11,9 +11,12 @@ from pydantic import BaseModel, ConfigDict, Field, ValidationError from tavern._core import exceptions +from tavern._core.loader import TypeConvertToken -# Type alias for JSON-compatible values (any valid JSON type) -JSONType = Union[dict, list, str, int, float, bool, None] +# Type alias for JSON-compatible values (any valid JSON type), plus +# TypeConvertToken for pre-resolution YAML tags like !force_format_include +# and dict for pre-resolution $ext function calls. +JSONType = Union[dict, list, str, int, float, bool, None, TypeConvertToken] class _BaseKeyValidator(BaseModel): @@ -54,110 +57,112 @@ def validate_keys(cls, data: Mapping) -> dict: # --- REST request spec --- class RestRequestSpec(_BaseKeyValidator): - method: Optional[str] = None - url: Optional[str] = None - headers: Optional[dict] = None - data: Optional[Union[dict, list, str, bytes]] = None - params: Optional[dict] = None - auth: Optional[Union[list, str]] = None + method: Optional[Union[str, dict, TypeConvertToken]] = None + url: Optional[Union[str, dict, TypeConvertToken]] = None + headers: Optional[Union[dict, TypeConvertToken]] = None + data: Optional[Union[dict, list, str, bytes, int, float, TypeConvertToken]] = None + params: Optional[Union[dict, TypeConvertToken]] = None + auth: Optional[Union[list, str, dict, TypeConvertToken]] = None json_body: Optional[JSONType] = Field(default=None, alias="json") - verify: Optional[Union[bool, str]] = None - files: Optional[Union[dict, list]] = None - file_body: Optional[str] = None - stream: Optional[bool] = None - timeout: Optional[Union[float, list]] = None - cookies: Optional[dict] = None - cert: Optional[Union[str, list]] = None - follow_redirects: Optional[bool] = None + verify: Optional[Union[bool, str, dict, TypeConvertToken]] = None + files: Optional[Union[dict, list, TypeConvertToken]] = None + file_body: Optional[Union[str, dict, TypeConvertToken]] = None + stream: Optional[Union[bool, dict, TypeConvertToken]] = None + timeout: Optional[Union[float, int, list, str, dict, TypeConvertToken]] = None + cookies: Optional[Union[dict, TypeConvertToken]] = None + cert: Optional[Union[str, list, int, dict, TypeConvertToken]] = None + follow_redirects: Optional[Union[bool, dict, TypeConvertToken]] = None # --- MQTT request spec --- class MQTTRequestSpec(_BaseKeyValidator): - topic: Optional[str] = None - payload: Optional[Union[str, bytes, int, float]] = None + topic: Optional[Union[str, dict, TypeConvertToken]] = None + payload: Optional[Union[str, bytes, int, float, dict, TypeConvertToken]] = None json_body: Optional[JSONType] = Field(default=None, alias="json") - qos: Optional[int] = None - retain: Optional[bool] = None + qos: Optional[Union[int, dict, TypeConvertToken]] = None + retain: Optional[Union[bool, dict, TypeConvertToken]] = None # --- MQTT client config blocks --- class MQTTClientArgs(_BaseKeyValidator): - client_id: Optional[str] = None - clean_session: Optional[bool] = None - transport: Optional[str] = None + client_id: Optional[Union[str, dict, TypeConvertToken]] = None + clean_session: Optional[Union[bool, dict, TypeConvertToken]] = None + transport: Optional[Union[str, dict, TypeConvertToken]] = None class MQTTConnectArgs(_BaseKeyValidator): - host: Optional[str] = None - port: Optional[int] = None - keepalive: Optional[int] = None - timeout: Optional[Union[int, float]] = None + host: Optional[Union[str, dict, TypeConvertToken]] = None + port: Optional[Union[int, dict, TypeConvertToken]] = None + keepalive: Optional[Union[int, dict, TypeConvertToken]] = None + timeout: Optional[Union[int, float, dict, TypeConvertToken]] = None class MQTTAuthArgs(_BaseKeyValidator): - username: Optional[str] = None - password: Optional[str] = None + username: Optional[Union[str, dict, TypeConvertToken]] = None + password: Optional[Union[str, dict, TypeConvertToken]] = None class MQTTTLSArgs(_BaseKeyValidator): - enable: Optional[bool] = None - ca_certs: Optional[str] = None - cert_reqs: Optional[str] = None - certfile: Optional[str] = None - keyfile: Optional[str] = None - tls_version: Optional[str] = None - ciphers: Optional[str] = None + enable: Optional[Union[bool, dict, TypeConvertToken]] = None + ca_certs: Optional[Union[str, dict, TypeConvertToken]] = None + cert_reqs: Optional[Union[str, dict, TypeConvertToken]] = None + certfile: Optional[Union[str, dict, TypeConvertToken]] = None + keyfile: Optional[Union[str, dict, TypeConvertToken]] = None + tls_version: Optional[Union[str, dict, TypeConvertToken]] = None + ciphers: Optional[Union[str, dict, TypeConvertToken]] = None class MQTTSSLContextArgs(_BaseKeyValidator): - ca_certs: Optional[str] = None - certfile: Optional[str] = None - keyfile: Optional[str] = None - password: Optional[str] = None - tls_version: Optional[str] = None - ciphers: Optional[str] = None - alpn_protocols: Optional[list[str]] = None + ca_certs: Optional[Union[str, dict, TypeConvertToken]] = None + certfile: Optional[Union[str, dict, TypeConvertToken]] = None + keyfile: Optional[Union[str, dict, TypeConvertToken]] = None + password: Optional[Union[str, dict, TypeConvertToken]] = None + tls_version: Optional[Union[str, dict, TypeConvertToken]] = None + ciphers: Optional[Union[str, dict, TypeConvertToken]] = None + alpn_protocols: Optional[Union[list[str], dict, TypeConvertToken]] = None class MQTTClientTopLevel(_BaseKeyValidator): - client: Optional[dict] = None - connect: Optional[dict] = None - tls: Optional[dict] = None - auth: Optional[dict] = None - ssl_context: Optional[dict] = None + client: Optional[Union[dict, TypeConvertToken]] = None + connect: Optional[Union[dict, TypeConvertToken]] = None + tls: Optional[Union[dict, TypeConvertToken]] = None + auth: Optional[Union[dict, TypeConvertToken]] = None + ssl_context: Optional[Union[dict, TypeConvertToken]] = None # --- gRPC request spec --- class GRPCRequestSpec(_BaseKeyValidator): - host: Optional[str] = None - service: Optional[str] = None - body: Optional[Union[dict, str]] = None + host: Optional[Union[str, dict, TypeConvertToken]] = None + service: Optional[Union[str, dict, TypeConvertToken]] = None + body: Optional[Union[dict, str, TypeConvertToken]] = None # --- gRPC response spec --- class GRPCResponseSpec(_BaseKeyValidator): - body: Optional[dict] = None - status: Optional[Union[str, int, list[str], list[int]]] = None - details: Optional[str] = None - save: Optional[dict] = None + body: Optional[Union[dict, TypeConvertToken]] = None + status: Optional[ + Union[str, int, list[str], list[int], dict, TypeConvertToken] + ] = None + details: Optional[Union[str, dict, TypeConvertToken]] = None + save: Optional[Union[dict, TypeConvertToken]] = None # --- gRPC client config blocks --- class GRPCConnectArgs(_BaseKeyValidator): - host: Optional[str] = None - port: Optional[int] = None - options: Optional[dict] = None - timeout: Optional[int] = None - secure: Optional[bool] = None + host: Optional[Union[str, dict, TypeConvertToken]] = None + port: Optional[Union[int, dict, TypeConvertToken]] = None + options: Optional[Union[dict, TypeConvertToken]] = None + timeout: Optional[Union[int, dict, TypeConvertToken]] = None + secure: Optional[Union[bool, dict, TypeConvertToken]] = None class GRPCProtoArgs(_BaseKeyValidator): - source: Optional[str] = None - module: Optional[str] = None + source: Optional[Union[str, dict, TypeConvertToken]] = None + module: Optional[Union[str, dict, TypeConvertToken]] = None class GRPCClientTopLevel(_BaseKeyValidator): - connect: Optional[dict] = None - proto: Optional[dict] = None - metadata: Optional[dict] = None - attempt_reflection: Optional[bool] = None + connect: Optional[Union[dict, TypeConvertToken]] = None + proto: Optional[Union[dict, TypeConvertToken]] = None + metadata: Optional[Union[dict, TypeConvertToken]] = None + attempt_reflection: Optional[Union[bool, dict, TypeConvertToken]] = None From 92aaa009a5419c260109a34c45d48687d26d7921 Mon Sep 17 00:00:00 2001 From: reachsridhard Date: Sun, 26 Jul 2026 13:25:31 +0530 Subject: [PATCH 04/15] fix: ruff format status field wrapping --- tavern/_core/pydantic_models.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/tavern/_core/pydantic_models.py b/tavern/_core/pydantic_models.py index 001ae6c0b..2ee58068f 100644 --- a/tavern/_core/pydantic_models.py +++ b/tavern/_core/pydantic_models.py @@ -140,9 +140,9 @@ class GRPCRequestSpec(_BaseKeyValidator): # --- gRPC response spec --- class GRPCResponseSpec(_BaseKeyValidator): body: Optional[Union[dict, TypeConvertToken]] = None - status: Optional[ - Union[str, int, list[str], list[int], dict, TypeConvertToken] - ] = None + status: Optional[Union[str, int, list[str], list[int], dict, TypeConvertToken]] = ( + None + ) details: Optional[Union[str, dict, TypeConvertToken]] = None save: Optional[Union[dict, TypeConvertToken]] = None From d265c5cce2bf0bed88a18f16fe1b0c2aefc4b186 Mon Sep 17 00:00:00 2001 From: reachsridhard Date: Sun, 26 Jul 2026 13:31:28 +0530 Subject: [PATCH 05/15] fix: allow list type for cookies field Integration tests use cookies as a list (e.g. cookie name lists, cookie override dicts in list form, empty list to send no cookies). --- tavern/_core/pydantic_models.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tavern/_core/pydantic_models.py b/tavern/_core/pydantic_models.py index 2ee58068f..1dc8be14c 100644 --- a/tavern/_core/pydantic_models.py +++ b/tavern/_core/pydantic_models.py @@ -69,7 +69,7 @@ class RestRequestSpec(_BaseKeyValidator): file_body: Optional[Union[str, dict, TypeConvertToken]] = None stream: Optional[Union[bool, dict, TypeConvertToken]] = None timeout: Optional[Union[float, int, list, str, dict, TypeConvertToken]] = None - cookies: Optional[Union[dict, TypeConvertToken]] = None + cookies: Optional[Union[dict, list, TypeConvertToken]] = None cert: Optional[Union[str, list, int, dict, TypeConvertToken]] = None follow_redirects: Optional[Union[bool, dict, TypeConvertToken]] = None From 0e63af2abb487a5f6b84edd3f4ca63a341f015ca Mon Sep 17 00:00:00 2001 From: reachsridhard Date: Tue, 25 Aug 2026 01:35:39 +0530 Subject: [PATCH 06/15] refactor: remove dict from boolean field types per maintainer feedback Boolean fields (enable, clean_session, stream, follow_redirects, retain, secure, attempt_reflection) should never receive dicts. Keep TypeConvertToken for !bool YAML tag support. --- tavern/_core/pydantic_models.py | 14 +++++++------- tests/unit/test_pydantic_models.py | 5 +++++ 2 files changed, 12 insertions(+), 7 deletions(-) diff --git a/tavern/_core/pydantic_models.py b/tavern/_core/pydantic_models.py index 1dc8be14c..0c9a1ceda 100644 --- a/tavern/_core/pydantic_models.py +++ b/tavern/_core/pydantic_models.py @@ -67,11 +67,11 @@ class RestRequestSpec(_BaseKeyValidator): verify: Optional[Union[bool, str, dict, TypeConvertToken]] = None files: Optional[Union[dict, list, TypeConvertToken]] = None file_body: Optional[Union[str, dict, TypeConvertToken]] = None - stream: Optional[Union[bool, dict, TypeConvertToken]] = None + stream: Optional[Union[bool, TypeConvertToken]] = None timeout: Optional[Union[float, int, list, str, dict, TypeConvertToken]] = None cookies: Optional[Union[dict, list, TypeConvertToken]] = None cert: Optional[Union[str, list, int, dict, TypeConvertToken]] = None - follow_redirects: Optional[Union[bool, dict, TypeConvertToken]] = None + follow_redirects: Optional[Union[bool, TypeConvertToken]] = None # --- MQTT request spec --- @@ -80,13 +80,13 @@ class MQTTRequestSpec(_BaseKeyValidator): payload: Optional[Union[str, bytes, int, float, dict, TypeConvertToken]] = None json_body: Optional[JSONType] = Field(default=None, alias="json") qos: Optional[Union[int, dict, TypeConvertToken]] = None - retain: Optional[Union[bool, dict, TypeConvertToken]] = None + retain: Optional[Union[bool, TypeConvertToken]] = None # --- MQTT client config blocks --- class MQTTClientArgs(_BaseKeyValidator): client_id: Optional[Union[str, dict, TypeConvertToken]] = None - clean_session: Optional[Union[bool, dict, TypeConvertToken]] = None + clean_session: Optional[Union[bool, TypeConvertToken]] = None transport: Optional[Union[str, dict, TypeConvertToken]] = None @@ -103,7 +103,7 @@ class MQTTAuthArgs(_BaseKeyValidator): class MQTTTLSArgs(_BaseKeyValidator): - enable: Optional[Union[bool, dict, TypeConvertToken]] = None + enable: Optional[Union[bool, TypeConvertToken]] = None ca_certs: Optional[Union[str, dict, TypeConvertToken]] = None cert_reqs: Optional[Union[str, dict, TypeConvertToken]] = None certfile: Optional[Union[str, dict, TypeConvertToken]] = None @@ -153,7 +153,7 @@ class GRPCConnectArgs(_BaseKeyValidator): port: Optional[Union[int, dict, TypeConvertToken]] = None options: Optional[Union[dict, TypeConvertToken]] = None timeout: Optional[Union[int, dict, TypeConvertToken]] = None - secure: Optional[Union[bool, dict, TypeConvertToken]] = None + secure: Optional[Union[bool, TypeConvertToken]] = None class GRPCProtoArgs(_BaseKeyValidator): @@ -165,4 +165,4 @@ class GRPCClientTopLevel(_BaseKeyValidator): connect: Optional[Union[dict, TypeConvertToken]] = None proto: Optional[Union[dict, TypeConvertToken]] = None metadata: Optional[Union[dict, TypeConvertToken]] = None - attempt_reflection: Optional[Union[bool, dict, TypeConvertToken]] = None + attempt_reflection: Optional[Union[bool, TypeConvertToken]] = None diff --git a/tests/unit/test_pydantic_models.py b/tests/unit/test_pydantic_models.py index bf9085aa0..e9b81166a 100644 --- a/tests/unit/test_pydantic_models.py +++ b/tests/unit/test_pydantic_models.py @@ -277,6 +277,11 @@ def test_mqtt_tls_enable_accepts_bool(self): result = MQTTTLSArgs.validate_keys(data) assert result["enable"] is True + def test_mqtt_tls_enable_rejects_dict(self): + data = {"enable": {"$ext": {"function": "some_func"}}} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTTLSArgs.validate_keys(data) + def test_mqtt_ssl_alpn_protocols_must_be_list(self): data = {"alpn_protocols": "h2"} with pytest.raises(exceptions.UnexpectedKeysError): From b964f2bccb17e7f61f132b2084e29d0d5d69ace2 Mon Sep 17 00:00:00 2001 From: reachsridhard Date: Tue, 25 Aug 2026 01:55:05 +0530 Subject: [PATCH 07/15] Apply suggestion from @coderabbitai[bot] Co-authored-by: coderabbitai[bot] <136622811+coderabbitai[bot]@users.noreply.github.com> --- tavern/_core/pydantic_models.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tavern/_core/pydantic_models.py b/tavern/_core/pydantic_models.py index 0c9a1ceda..93d563f54 100644 --- a/tavern/_core/pydantic_models.py +++ b/tavern/_core/pydantic_models.py @@ -23,10 +23,10 @@ class _BaseKeyValidator(BaseModel): """Base model that forbids extra keys and raises UnexpectedKeysError on validation failure.""" model_config = ConfigDict( - extra="forbid", arbitrary_types_allowed=True, populate_by_name=True + extra="forbid", arbitrary_types_allowed=True, populate_by_name=True, hide_input_in_errors=True ) - @classmethod + `@classmethod` def validate_keys(cls, data: Mapping) -> dict: """Validate that ``data`` contains only expected keys and types. From 1263000e97093ffd76c438a656156d79ad3bc9d3 Mon Sep 17 00:00:00 2001 From: reachsridhard Date: Tue, 25 Aug 2026 01:56:12 +0530 Subject: [PATCH 08/15] Apply suggestion from @coderabbitai[bot] Co-authored-by: coderabbitai[bot] <136622811+coderabbitai[bot]@users.noreply.github.com> --- tests/unit/test_pydantic_models.py | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_pydantic_models.py b/tests/unit/test_pydantic_models.py index e9b81166a..fbfc9f3c3 100644 --- a/tests/unit/test_pydantic_models.py +++ b/tests/unit/test_pydantic_models.py @@ -302,6 +302,14 @@ def test_grpc_request_host_must_be_string(self): with pytest.raises(exceptions.UnexpectedKeysError): GRPCRequestSpec.validate_keys(data) + def test_grpc_request_body_can_be_dict(self): + data = { + "host": "localhost:50051", + "service": "MyService/Method", + "body": {"key": "value"}, + } + result = GRPCRequestSpec.validate_keys(data) + assert result["body"] == {"key": "value"} def test_grpc_request_body_can_be_dict(self): data = { "host": "localhost:50051", @@ -319,7 +327,6 @@ def test_grpc_request_body_can_be_string(self): } result = GRPCRequestSpec.validate_keys(data) assert result["body"] == "raw string" - def test_grpc_response_status_can_be_int(self): data = {"status": 0} result = GRPCResponseSpec.validate_keys(data) From 9a894dbde811044e055438c305a303bee2ca8723 Mon Sep 17 00:00:00 2001 From: reachsridhard Date: Fri, 4 Sep 2026 20:18:24 +0530 Subject: [PATCH 09/15] fix: stricter type checking on pydantic models --- tavern/_core/pydantic_models.py | 171 +++++++++++++++++------------ tavern/_plugins/mqtt/client.py | 3 +- tests/unit/test_pydantic_models.py | 110 ++++++++++++++++++- 3 files changed, 206 insertions(+), 78 deletions(-) diff --git a/tavern/_core/pydantic_models.py b/tavern/_core/pydantic_models.py index 93d563f54..39b1dab4b 100644 --- a/tavern/_core/pydantic_models.py +++ b/tavern/_core/pydantic_models.py @@ -11,22 +11,42 @@ from pydantic import BaseModel, ConfigDict, Field, ValidationError from tavern._core import exceptions -from tavern._core.loader import TypeConvertToken +from tavern._core.loader import ( + BoolToken, + FloatToken, + ForceIncludeToken, + IntToken, +) # Type alias for JSON-compatible values (any valid JSON type), plus -# TypeConvertToken for pre-resolution YAML tags like !force_format_include -# and dict for pre-resolution $ext function calls. -JSONType = Union[dict, list, str, int, float, bool, None, TypeConvertToken] +# conversion tokens for pre-resolution YAML tags like !force_format_include +# and !int/!float/!bool, and dict for pre-resolution $ext function calls. +JSONType = Union[ + dict, + list, + str, + int, + float, + bool, + None, + IntToken, + FloatToken, + BoolToken, + ForceIncludeToken, +] class _BaseKeyValidator(BaseModel): """Base model that forbids extra keys and raises UnexpectedKeysError on validation failure.""" model_config = ConfigDict( - extra="forbid", arbitrary_types_allowed=True, populate_by_name=True, hide_input_in_errors=True + extra="forbid", + arbitrary_types_allowed=True, + populate_by_name=True, + hide_input_in_errors=True, ) - `@classmethod` + @classmethod def validate_keys(cls, data: Mapping) -> dict: """Validate that ``data`` contains only expected keys and types. @@ -45,124 +65,129 @@ def validate_keys(cls, data: Mapping) -> dict: except ValidationError as e: # Extract unexpected field names from the error unexpected = set() + missing = set() for err in e.errors(): if err["type"] == "extra_forbidden": unexpected.add(err["loc"][-1]) + elif err["type"] == "missing": + missing.add(err["loc"][-1]) if unexpected: msg = f"Unexpected keys {unexpected}" + raise exceptions.UnexpectedKeysError(msg) from e + elif missing: + msg = f"Missing keys {missing}" + raise exceptions.MissingKeysError(msg) from e else: msg = str(e) - raise exceptions.UnexpectedKeysError(msg) from e + raise exceptions.UnexpectedKeysError(msg) from e # --- REST request spec --- class RestRequestSpec(_BaseKeyValidator): - method: Optional[Union[str, dict, TypeConvertToken]] = None - url: Optional[Union[str, dict, TypeConvertToken]] = None - headers: Optional[Union[dict, TypeConvertToken]] = None - data: Optional[Union[dict, list, str, bytes, int, float, TypeConvertToken]] = None - params: Optional[Union[dict, TypeConvertToken]] = None - auth: Optional[Union[list, str, dict, TypeConvertToken]] = None + method: Optional[str] = None + url: Union[str, dict] # required; dict for pre-resolution $ext + headers: Optional[dict] = None # dict for pre-resolution $ext + data: Optional[Union[dict, list, str, bytes, int, float]] = None + params: Optional[dict] = None # dict for pre-resolution $ext + auth: Optional[Union[list, str, dict]] = None # dict for pre-resolution $ext json_body: Optional[JSONType] = Field(default=None, alias="json") - verify: Optional[Union[bool, str, dict, TypeConvertToken]] = None - files: Optional[Union[dict, list, TypeConvertToken]] = None - file_body: Optional[Union[str, dict, TypeConvertToken]] = None - stream: Optional[Union[bool, TypeConvertToken]] = None - timeout: Optional[Union[float, int, list, str, dict, TypeConvertToken]] = None - cookies: Optional[Union[dict, list, TypeConvertToken]] = None - cert: Optional[Union[str, list, int, dict, TypeConvertToken]] = None - follow_redirects: Optional[Union[bool, TypeConvertToken]] = None + verify: Optional[Union[bool, str]] = None + files: Optional[Union[dict, list]] = None # dict for pre-resolution $ext + file_body: Optional[str] = None + stream: Optional[bool] = None + timeout: Optional[Union[float, int, list, str, dict]] = None # dict for $ext + cookies: Optional[Union[dict, list]] = None + cert: Optional[Union[str, list, int, dict]] = None # dict for $ext + follow_redirects: Optional[bool] = None # --- MQTT request spec --- class MQTTRequestSpec(_BaseKeyValidator): - topic: Optional[Union[str, dict, TypeConvertToken]] = None - payload: Optional[Union[str, bytes, int, float, dict, TypeConvertToken]] = None + topic: Optional[str] = None + payload: Optional[Union[str, bytes, int, float]] = None json_body: Optional[JSONType] = Field(default=None, alias="json") - qos: Optional[Union[int, dict, TypeConvertToken]] = None - retain: Optional[Union[bool, TypeConvertToken]] = None + qos: Optional[Union[int, IntToken]] = None + retain: Optional[Union[bool, BoolToken]] = None # --- MQTT client config blocks --- class MQTTClientArgs(_BaseKeyValidator): - client_id: Optional[Union[str, dict, TypeConvertToken]] = None - clean_session: Optional[Union[bool, TypeConvertToken]] = None - transport: Optional[Union[str, dict, TypeConvertToken]] = None + client_id: Optional[str] = None + clean_session: Optional[Union[bool, BoolToken]] = None + transport: Optional[str] = None class MQTTConnectArgs(_BaseKeyValidator): - host: Optional[Union[str, dict, TypeConvertToken]] = None - port: Optional[Union[int, dict, TypeConvertToken]] = None - keepalive: Optional[Union[int, dict, TypeConvertToken]] = None - timeout: Optional[Union[int, float, dict, TypeConvertToken]] = None + host: str # required + port: Optional[Union[int, IntToken]] = None + keepalive: Optional[Union[int, IntToken]] = None + timeout: Optional[Union[int, float, IntToken, FloatToken]] = None class MQTTAuthArgs(_BaseKeyValidator): - username: Optional[Union[str, dict, TypeConvertToken]] = None - password: Optional[Union[str, dict, TypeConvertToken]] = None + username: str # required + password: Optional[str] = None class MQTTTLSArgs(_BaseKeyValidator): - enable: Optional[Union[bool, TypeConvertToken]] = None - ca_certs: Optional[Union[str, dict, TypeConvertToken]] = None - cert_reqs: Optional[Union[str, dict, TypeConvertToken]] = None - certfile: Optional[Union[str, dict, TypeConvertToken]] = None - keyfile: Optional[Union[str, dict, TypeConvertToken]] = None - tls_version: Optional[Union[str, dict, TypeConvertToken]] = None - ciphers: Optional[Union[str, dict, TypeConvertToken]] = None + enable: Optional[Union[bool, BoolToken]] = None + ca_certs: Optional[str] = None + cert_reqs: Optional[str] = None + certfile: Optional[str] = None + keyfile: Optional[str] = None + tls_version: Optional[str] = None + ciphers: Optional[str] = None class MQTTSSLContextArgs(_BaseKeyValidator): - ca_certs: Optional[Union[str, dict, TypeConvertToken]] = None - certfile: Optional[Union[str, dict, TypeConvertToken]] = None - keyfile: Optional[Union[str, dict, TypeConvertToken]] = None - password: Optional[Union[str, dict, TypeConvertToken]] = None - tls_version: Optional[Union[str, dict, TypeConvertToken]] = None - ciphers: Optional[Union[str, dict, TypeConvertToken]] = None - alpn_protocols: Optional[Union[list[str], dict, TypeConvertToken]] = None + ca_certs: Optional[str] = None + certfile: Optional[str] = None + keyfile: Optional[str] = None + password: Optional[str] = None + tls_version: Optional[str] = None + ciphers: Optional[str] = None + alpn_protocols: Optional[list[str]] = None class MQTTClientTopLevel(_BaseKeyValidator): - client: Optional[Union[dict, TypeConvertToken]] = None - connect: Optional[Union[dict, TypeConvertToken]] = None - tls: Optional[Union[dict, TypeConvertToken]] = None - auth: Optional[Union[dict, TypeConvertToken]] = None - ssl_context: Optional[Union[dict, TypeConvertToken]] = None + client: Optional[dict] = None + connect: Optional[dict] = None + tls: Optional[dict] = None + auth: Optional[dict] = None + ssl_context: Optional[dict] = None # --- gRPC request spec --- class GRPCRequestSpec(_BaseKeyValidator): - host: Optional[Union[str, dict, TypeConvertToken]] = None - service: Optional[Union[str, dict, TypeConvertToken]] = None - body: Optional[Union[dict, str, TypeConvertToken]] = None + host: Optional[str] = None + service: str # required + body: Optional[Union[dict, str]] = None # --- gRPC response spec --- class GRPCResponseSpec(_BaseKeyValidator): - body: Optional[Union[dict, TypeConvertToken]] = None - status: Optional[Union[str, int, list[str], list[int], dict, TypeConvertToken]] = ( - None - ) - details: Optional[Union[str, dict, TypeConvertToken]] = None - save: Optional[Union[dict, TypeConvertToken]] = None + body: Optional[dict] = None + status: Optional[Union[str, int, list[str], list[int]]] = None + details: Optional[Union[str, dict]] = None + save: Optional[dict] = None # --- gRPC client config blocks --- class GRPCConnectArgs(_BaseKeyValidator): - host: Optional[Union[str, dict, TypeConvertToken]] = None - port: Optional[Union[int, dict, TypeConvertToken]] = None - options: Optional[Union[dict, TypeConvertToken]] = None - timeout: Optional[Union[int, dict, TypeConvertToken]] = None - secure: Optional[Union[bool, TypeConvertToken]] = None + host: Optional[str] = None + port: Optional[Union[int, IntToken]] = None + options: Optional[dict] = None + timeout: Optional[Union[int, IntToken]] = None + secure: Optional[Union[bool, BoolToken]] = None class GRPCProtoArgs(_BaseKeyValidator): - source: Optional[Union[str, dict, TypeConvertToken]] = None - module: Optional[Union[str, dict, TypeConvertToken]] = None + source: Optional[str] = None + module: Optional[str] = None class GRPCClientTopLevel(_BaseKeyValidator): - connect: Optional[Union[dict, TypeConvertToken]] = None - proto: Optional[Union[dict, TypeConvertToken]] = None - metadata: Optional[Union[dict, TypeConvertToken]] = None - attempt_reflection: Optional[Union[bool, TypeConvertToken]] = None + connect: Optional[dict] = None + proto: Optional[dict] = None + metadata: Optional[dict] = None + attempt_reflection: Optional[Union[bool, BoolToken]] = None diff --git a/tavern/_plugins/mqtt/client.py b/tavern/_plugins/mqtt/client.py index 044ff5c85..f009aeb96 100644 --- a/tavern/_plugins/mqtt/client.py +++ b/tavern/_plugins/mqtt/client.py @@ -146,7 +146,8 @@ def __init__(self, **kwargs) -> None: MQTTConnectArgs.validate_keys(self._connect_args) self._auth_args = kwargs.pop("auth", {}) - MQTTAuthArgs.validate_keys(self._auth_args) + if self._auth_args: + MQTTAuthArgs.validate_keys(self._auth_args) if "host" not in self._connect_args: msg = "Need 'host' in 'connect' block for mqtt" diff --git a/tests/unit/test_pydantic_models.py b/tests/unit/test_pydantic_models.py index fbfc9f3c3..c8c8b1609 100644 --- a/tests/unit/test_pydantic_models.py +++ b/tests/unit/test_pydantic_models.py @@ -34,8 +34,9 @@ def test_unexpected_key(self): RestRequestSpec.validate_keys(data) def test_empty_dict(self): - result = RestRequestSpec.validate_keys({}) - assert result == {} + """url is required, so an empty dict should fail""" + with pytest.raises(exceptions.MissingKeysError): + RestRequestSpec.validate_keys({}) class TestMQTTRequestSpec: @@ -168,7 +169,7 @@ class TestTypeValidation: """Tests verifying that pydantic models enforce type checking, not just key validation.""" def test_rest_method_must_be_string(self): - data = {"method": 123} + data = {"url": "http://example.com", "method": 123} with pytest.raises(exceptions.UnexpectedKeysError): RestRequestSpec.validate_keys(data) @@ -298,7 +299,7 @@ def test_mqtt_client_top_level_blocks_must_be_dict(self): MQTTClientTopLevel.validate_keys(data) def test_grpc_request_host_must_be_string(self): - data = {"host": 123} + data = {"service": "MyService/Method", "host": 123} with pytest.raises(exceptions.UnexpectedKeysError): GRPCRequestSpec.validate_keys(data) @@ -371,3 +372,104 @@ def test_grpc_client_attempt_reflection_accepts_bool(self): data = {"attempt_reflection": True} result = GRPCClientTopLevel.validate_keys(data) assert result["attempt_reflection"] is True + + +class TestRequiredFields: + """Tests verifying that required fields are enforced.""" + + def test_rest_url_is_required(self): + data = {"method": "GET"} + with pytest.raises(exceptions.MissingKeysError): + RestRequestSpec.validate_keys(data) + + def test_mqtt_connect_host_is_required(self): + data = {"port": 1883} + with pytest.raises(exceptions.MissingKeysError): + MQTTConnectArgs.validate_keys(data) + + def test_mqtt_auth_username_is_required(self): + data = {"password": "pass"} + with pytest.raises(exceptions.MissingKeysError): + MQTTAuthArgs.validate_keys(data) + + def test_grpc_request_service_is_required(self): + data = {"host": "localhost:50051"} + with pytest.raises(exceptions.MissingKeysError): + GRPCRequestSpec.validate_keys(data) + + +class TestSpecificTokenTypes: + """Tests verifying that conversion tokens are type-specific.""" + + def test_mqtt_qos_accepts_int_token(self): + from tavern._core.loader import IntToken + + data = {"topic": "test/topic", "qos": IntToken("{qos:d}")} + result = MQTTRequestSpec.validate_keys(data) + assert "qos" in result + + def test_mqtt_qos_rejects_bool_token(self): + from tavern._core.loader import BoolToken + + data = {"topic": "test/topic", "qos": BoolToken("{qos:d}")} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTRequestSpec.validate_keys(data) + + def test_mqtt_retain_accepts_bool_token(self): + from tavern._core.loader import BoolToken + + data = {"topic": "test/topic", "retain": BoolToken("{retain:d}")} + result = MQTTRequestSpec.validate_keys(data) + assert "retain" in result + + def test_mqtt_retain_rejects_int_token(self): + from tavern._core.loader import IntToken + + data = {"topic": "test/topic", "retain": IntToken("{retain:d}")} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTRequestSpec.validate_keys(data) + + def test_mqtt_connect_port_accepts_int_token(self): + from tavern._core.loader import IntToken + + data = {"host": "localhost", "port": IntToken("{port:d}")} + result = MQTTConnectArgs.validate_keys(data) + assert "port" in result + + def test_grpc_connect_port_rejects_bool_token(self): + from tavern._core.loader import BoolToken + + data = {"host": "localhost", "port": BoolToken("{port:d}")} + with pytest.raises(exceptions.UnexpectedKeysError): + GRPCConnectArgs.validate_keys(data) + + def test_rest_method_rejects_dict(self): + data = { + "url": "http://example.com", + "method": {"$ext": {"function": "mod:func"}}, + } + with pytest.raises(exceptions.UnexpectedKeysError): + RestRequestSpec.validate_keys(data) + + def test_rest_verify_rejects_dict(self): + data = { + "url": "http://example.com", + "verify": {"$ext": {"function": "mod:func"}}, + } + with pytest.raises(exceptions.UnexpectedKeysError): + RestRequestSpec.validate_keys(data) + + def test_mqtt_topic_rejects_dict(self): + data = {"topic": {"$ext": {"function": "mod:func"}}} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTRequestSpec.validate_keys(data) + + def test_mqtt_payload_rejects_dict(self): + data = {"topic": "test/topic", "payload": {"key": "value"}} + with pytest.raises(exceptions.UnexpectedKeysError): + MQTTRequestSpec.validate_keys(data) + + def test_grpc_response_status_rejects_dict(self): + data = {"status": {"key": "value"}} + with pytest.raises(exceptions.UnexpectedKeysError): + GRPCResponseSpec.validate_keys(data) From 11fd9bf5ed0d436062bb0b66884e6bbab946e8da Mon Sep 17 00:00:00 2001 From: reachsridhard Date: Fri, 4 Sep 2026 20:20:53 +0530 Subject: [PATCH 10/15] ci: restart checks From 56cbd2d27197f10b6b2a2aba879cb75f9f8db378 Mon Sep 17 00:00:00 2001 From: reachsridhard Date: Fri, 4 Sep 2026 20:27:30 +0530 Subject: [PATCH 11/15] fix: remove duplicate test and fix formatting --- tests/unit/test_pydantic_models.py | 9 +-------- 1 file changed, 1 insertion(+), 8 deletions(-) diff --git a/tests/unit/test_pydantic_models.py b/tests/unit/test_pydantic_models.py index c8c8b1609..ec9c9639b 100644 --- a/tests/unit/test_pydantic_models.py +++ b/tests/unit/test_pydantic_models.py @@ -303,14 +303,6 @@ def test_grpc_request_host_must_be_string(self): with pytest.raises(exceptions.UnexpectedKeysError): GRPCRequestSpec.validate_keys(data) - def test_grpc_request_body_can_be_dict(self): - data = { - "host": "localhost:50051", - "service": "MyService/Method", - "body": {"key": "value"}, - } - result = GRPCRequestSpec.validate_keys(data) - assert result["body"] == {"key": "value"} def test_grpc_request_body_can_be_dict(self): data = { "host": "localhost:50051", @@ -328,6 +320,7 @@ def test_grpc_request_body_can_be_string(self): } result = GRPCRequestSpec.validate_keys(data) assert result["body"] == "raw string" + def test_grpc_response_status_can_be_int(self): data = {"status": 0} result = GRPCResponseSpec.validate_keys(data) From b2afe8b94e269627aadda86754311227284b735c Mon Sep 17 00:00:00 2001 From: reachsridhard Date: Fri, 4 Sep 2026 23:27:11 +0530 Subject: [PATCH 12/15] fix: add BoolToken to verify, stream, follow_redirects fields --- tavern/_core/pydantic_models.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/tavern/_core/pydantic_models.py b/tavern/_core/pydantic_models.py index 39b1dab4b..0a815472a 100644 --- a/tavern/_core/pydantic_models.py +++ b/tavern/_core/pydantic_models.py @@ -91,14 +91,14 @@ class RestRequestSpec(_BaseKeyValidator): params: Optional[dict] = None # dict for pre-resolution $ext auth: Optional[Union[list, str, dict]] = None # dict for pre-resolution $ext json_body: Optional[JSONType] = Field(default=None, alias="json") - verify: Optional[Union[bool, str]] = None + verify: Optional[Union[bool, BoolToken, str]] = None files: Optional[Union[dict, list]] = None # dict for pre-resolution $ext file_body: Optional[str] = None - stream: Optional[bool] = None + stream: Optional[Union[bool, BoolToken]] = None timeout: Optional[Union[float, int, list, str, dict]] = None # dict for $ext cookies: Optional[Union[dict, list]] = None cert: Optional[Union[str, list, int, dict]] = None # dict for $ext - follow_redirects: Optional[bool] = None + follow_redirects: Optional[Union[bool, BoolToken]] = None # --- MQTT request spec --- From d8f1fbe261b9d2e161a669cee4fd4a9a9799a7a9 Mon Sep 17 00:00:00 2001 From: reachsridhard Date: Fri, 11 Sep 2026 16:29:21 +0530 Subject: [PATCH 13/15] Update tavern/_core/pydantic_models.py Co-authored-by: michaelboulton --- tavern/_core/pydantic_models.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/tavern/_core/pydantic_models.py b/tavern/_core/pydantic_models.py index 0a815472a..e1a77a978 100644 --- a/tavern/_core/pydantic_models.py +++ b/tavern/_core/pydantic_models.py @@ -150,11 +150,11 @@ class MQTTSSLContextArgs(_BaseKeyValidator): class MQTTClientTopLevel(_BaseKeyValidator): - client: Optional[dict] = None - connect: Optional[dict] = None - tls: Optional[dict] = None - auth: Optional[dict] = None - ssl_context: Optional[dict] = None + client: Optional[MQTTClientArgs] = None + connect: Optional[MQTTConnectArgs] = None + tls: Optional[MQTTTLSArgs] = None + auth: Optional[MQTTAuthArgs] = None + ssl_context: Optional[MQTTAuthArgs] = None # --- gRPC request spec --- From 1005846e35e52d500986b188116beb41442c2827 Mon Sep 17 00:00:00 2001 From: reachsridhard Date: Fri, 11 Sep 2026 16:29:42 +0530 Subject: [PATCH 14/15] Update tavern/_core/pydantic_models.py Co-authored-by: michaelboulton --- tavern/_core/pydantic_models.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tavern/_core/pydantic_models.py b/tavern/_core/pydantic_models.py index e1a77a978..9d26c11cc 100644 --- a/tavern/_core/pydantic_models.py +++ b/tavern/_core/pydantic_models.py @@ -187,7 +187,7 @@ class GRPCProtoArgs(_BaseKeyValidator): class GRPCClientTopLevel(_BaseKeyValidator): - connect: Optional[dict] = None - proto: Optional[dict] = None + connect: Optional[GRPCConnectArgs] = None + proto: Optional[GRPCProtoArgs] = None metadata: Optional[dict] = None attempt_reflection: Optional[Union[bool, BoolToken]] = None From 43cdd833929717f0bfd060d9f017d20955c85b56 Mon Sep 17 00:00:00 2001 From: reachsridhard Date: Sat, 26 Sep 2026 12:57:44 +0530 Subject: [PATCH 15/15] fix: use MQTTSSLContextArgs for ssl_context and fix auth test --- tavern/_core/pydantic_models.py | 2 +- tests/unit/test_pydantic_models.py | 6 +++++- 2 files changed, 6 insertions(+), 2 deletions(-) diff --git a/tavern/_core/pydantic_models.py b/tavern/_core/pydantic_models.py index 9d26c11cc..e45a6d5d7 100644 --- a/tavern/_core/pydantic_models.py +++ b/tavern/_core/pydantic_models.py @@ -154,7 +154,7 @@ class MQTTClientTopLevel(_BaseKeyValidator): connect: Optional[MQTTConnectArgs] = None tls: Optional[MQTTTLSArgs] = None auth: Optional[MQTTAuthArgs] = None - ssl_context: Optional[MQTTAuthArgs] = None + ssl_context: Optional[MQTTSSLContextArgs] = None # --- gRPC request spec --- diff --git a/tests/unit/test_pydantic_models.py b/tests/unit/test_pydantic_models.py index ec9c9639b..7ae996746 100644 --- a/tests/unit/test_pydantic_models.py +++ b/tests/unit/test_pydantic_models.py @@ -53,7 +53,11 @@ def test_unexpected_key(self): class TestMQTTClientSpecs: def test_top_level_valid(self): - data = {"client": {}, "connect": {"host": "localhost"}, "auth": {}} + data = { + "client": {}, + "connect": {"host": "localhost"}, + "auth": {"username": "user"}, + } result = MQTTClientTopLevel.validate_keys(data) assert "client" in result