diff --git a/pyproject.toml b/pyproject.toml index 386e614d..8773581b 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 00000000..93d563f5 --- /dev/null +++ b/tavern/_core/pydantic_models.py @@ -0,0 +1,168 @@ +"""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 +from tavern._core.loader import TypeConvertToken + +# 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): + """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 + ) + + `@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[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, 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 + + +# --- MQTT request spec --- +class MQTTRequestSpec(_BaseKeyValidator): + 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[Union[int, 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, TypeConvertToken]] = None + transport: Optional[Union[str, dict, TypeConvertToken]] = 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 + + +class MQTTAuthArgs(_BaseKeyValidator): + username: Optional[Union[str, dict, TypeConvertToken]] = None + password: Optional[Union[str, dict, TypeConvertToken]] = 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 + + +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 + + +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 + + +# --- 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 + + +# --- 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 + + +# --- 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 + + +class GRPCProtoArgs(_BaseKeyValidator): + source: Optional[Union[str, dict, TypeConvertToken]] = None + module: Optional[Union[str, dict, TypeConvertToken]] = 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 diff --git a/tavern/_plugins/grpc/client.py b/tavern/_plugins/grpc/client.py index 5720a1cc..b65a2d5a 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 9fdfa9f1..bd84dc36 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 6df71687..c6176ad8 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 fd917670..044ff5c8 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 b4d1e999..e18d8be4 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 65607b83..0b0eada7 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 00000000..fbfc9f3c --- /dev/null +++ b/tests/unit/test_pydantic_models.py @@ -0,0 +1,373 @@ +"""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_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): + 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_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 diff --git a/uv.lock b/uv.lock index 4388721c..d232b410 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" },