Skip to content
Open
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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",
]

Expand Down
163 changes: 163 additions & 0 deletions tavern/_core/pydantic_models.py
Original file line number Diff line number Diff line change
@@ -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
Comment thread
reachsridhard marked this conversation as resolved.


# --- 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
18 changes: 8 additions & 10 deletions tavern/_plugins/grpc/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)
Expand All @@ -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))

Expand Down
7 changes: 3 additions & 4 deletions tavern/_plugins/grpc/request.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)

Expand Down
4 changes: 2 additions & 2 deletions tavern/_plugins/grpc/response.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down
53 changes: 14 additions & 39 deletions tavern/_plugins/mqtt/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = {
Expand Down Expand Up @@ -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:
Expand All @@ -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"
Expand All @@ -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)
Expand Down
7 changes: 3 additions & 4 deletions tavern/_plugins/mqtt/request.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)

Expand Down
Loading
Loading