From 455566419ff62d0b9ff437d604baec6eb55dbbd8 Mon Sep 17 00:00:00 2001 From: Gerrit Date: Mon, 27 Jul 2026 22:13:47 +0200 Subject: [PATCH] Add test interceptor to Python client. --- generate/python_client.tpl | 11 ++- python/metalstack/client/__init__.py | 1 + python/metalstack/client/client.py | 97 ++++++++++---------- python/metalstack/client/test_interceptor.py | 57 ++++++++++++ 4 files changed, 116 insertions(+), 50 deletions(-) create mode 100644 python/metalstack/client/test_interceptor.py diff --git a/generate/python_client.tpl b/generate/python_client.tpl index 4e2d39ed..4384c0a9 100644 --- a/generate/python_client.tpl +++ b/generate/python_client.tpl @@ -8,9 +8,11 @@ import metalstack.{{ $name | trimSuffix "v2" }}.v2.{{ $svc.FileName | trimSuffix {{ end }} {{ end }} + class Client: - def __init__(self, baseurl: str, timeout: float = 10): + def __init__(self, baseurl: str, timeout: float = 10, interceptors: list = []): self._baseurl = baseurl + self._interceptors = list(interceptors) transport = pyqwest.SyncHTTPTransport( http_version=pyqwest.HTTPVersion.HTTP2, @@ -21,17 +23,18 @@ class Client: {{ range $name, $api := . }} def {{ $name | lower }}(self): - return self._{{ $name | title }}(baseurl=self._baseurl, client=self._client) + return self._{{ $name | title }}(baseurl=self._baseurl, client=self._client, interceptors=self._interceptors) {{ end }} {{ range $name, $api := . }} class _{{ $name | title }}: - def __init__(self, baseurl: str, client: pyqwest.SyncClient = None): + def __init__(self, baseurl: str, client: pyqwest.SyncClient = None, interceptors: list = []): self._baseurl = baseurl self._client = client + self._interceptors = list(interceptors) {{ range $svc := $api.Services }} def {{ $svc.FileName | trimSuffix ".proto" | lower }}(self): - return {{ $name | trimSuffix "v2" }}_{{ $svc.FileName | trimSuffix ".proto" | lower }}_connect.{{ $svc.Name }}ClientSync(address=self._baseurl, http_client=self._client) + return {{ $name | trimSuffix "v2" }}_{{ $svc.FileName | trimSuffix ".proto" | lower }}_connect.{{ $svc.Name }}ClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) {{ end }} {{ end }} diff --git a/python/metalstack/client/__init__.py b/python/metalstack/client/__init__.py index e69de29b..8b137891 100644 --- a/python/metalstack/client/__init__.py +++ b/python/metalstack/client/__init__.py @@ -0,0 +1 @@ + diff --git a/python/metalstack/client/client.py b/python/metalstack/client/client.py index 9f15fc31..02719011 100755 --- a/python/metalstack/client/client.py +++ b/python/metalstack/client/client.py @@ -46,9 +46,11 @@ + class Client: - def __init__(self, baseurl: str, timeout: float = 10): + def __init__(self, baseurl: str, timeout: float = 10, interceptors: list = []): self._baseurl = baseurl + self._interceptors = list(interceptors) transport = pyqwest.SyncHTTPTransport( http_version=pyqwest.HTTPVersion.HTTP2, @@ -59,151 +61,154 @@ def __init__(self, baseurl: str, timeout: float = 10): def adminv2(self): - return self._Adminv2(baseurl=self._baseurl, client=self._client) + return self._Adminv2(baseurl=self._baseurl, client=self._client, interceptors=self._interceptors) def apiv2(self): - return self._Apiv2(baseurl=self._baseurl, client=self._client) + return self._Apiv2(baseurl=self._baseurl, client=self._client, interceptors=self._interceptors) def infrav2(self): - return self._Infrav2(baseurl=self._baseurl, client=self._client) + return self._Infrav2(baseurl=self._baseurl, client=self._client, interceptors=self._interceptors) class _Adminv2: - def __init__(self, baseurl: str, client: pyqwest.SyncClient = None): + def __init__(self, baseurl: str, client: pyqwest.SyncClient = None, interceptors: list = []): self._baseurl = baseurl self._client = client + self._interceptors = list(interceptors) def audit(self): - return admin_audit_connect.AuditServiceClientSync(address=self._baseurl, http_client=self._client) + return admin_audit_connect.AuditServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def component(self): - return admin_component_connect.ComponentServiceClientSync(address=self._baseurl, http_client=self._client) + return admin_component_connect.ComponentServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def filesystem(self): - return admin_filesystem_connect.FilesystemServiceClientSync(address=self._baseurl, http_client=self._client) + return admin_filesystem_connect.FilesystemServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def image(self): - return admin_image_connect.ImageServiceClientSync(address=self._baseurl, http_client=self._client) + return admin_image_connect.ImageServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def ip(self): - return admin_ip_connect.IPServiceClientSync(address=self._baseurl, http_client=self._client) + return admin_ip_connect.IPServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def machine(self): - return admin_machine_connect.MachineServiceClientSync(address=self._baseurl, http_client=self._client) + return admin_machine_connect.MachineServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def network(self): - return admin_network_connect.NetworkServiceClientSync(address=self._baseurl, http_client=self._client) + return admin_network_connect.NetworkServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def partition(self): - return admin_partition_connect.PartitionServiceClientSync(address=self._baseurl, http_client=self._client) + return admin_partition_connect.PartitionServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def project(self): - return admin_project_connect.ProjectServiceClientSync(address=self._baseurl, http_client=self._client) + return admin_project_connect.ProjectServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def size(self): - return admin_size_connect.SizeServiceClientSync(address=self._baseurl, http_client=self._client) + return admin_size_connect.SizeServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def size_imageconstraint(self): - return admin_size_imageconstraint_connect.SizeImageConstraintServiceClientSync(address=self._baseurl, http_client=self._client) + return admin_size_imageconstraint_connect.SizeImageConstraintServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def size_reservation(self): - return admin_size_reservation_connect.SizeReservationServiceClientSync(address=self._baseurl, http_client=self._client) + return admin_size_reservation_connect.SizeReservationServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def switch(self): - return admin_switch_connect.SwitchServiceClientSync(address=self._baseurl, http_client=self._client) + return admin_switch_connect.SwitchServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def task(self): - return admin_task_connect.TaskServiceClientSync(address=self._baseurl, http_client=self._client) + return admin_task_connect.TaskServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def tenant(self): - return admin_tenant_connect.TenantServiceClientSync(address=self._baseurl, http_client=self._client) + return admin_tenant_connect.TenantServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def token(self): - return admin_token_connect.TokenServiceClientSync(address=self._baseurl, http_client=self._client) + return admin_token_connect.TokenServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def vpn(self): - return admin_vpn_connect.VPNServiceClientSync(address=self._baseurl, http_client=self._client) + return admin_vpn_connect.VPNServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) class _Apiv2: - def __init__(self, baseurl: str, client: pyqwest.SyncClient = None): + def __init__(self, baseurl: str, client: pyqwest.SyncClient = None, interceptors: list = []): self._baseurl = baseurl self._client = client + self._interceptors = list(interceptors) def audit(self): - return api_audit_connect.AuditServiceClientSync(address=self._baseurl, http_client=self._client) + return api_audit_connect.AuditServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def filesystem(self): - return api_filesystem_connect.FilesystemServiceClientSync(address=self._baseurl, http_client=self._client) + return api_filesystem_connect.FilesystemServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def health(self): - return api_health_connect.HealthServiceClientSync(address=self._baseurl, http_client=self._client) + return api_health_connect.HealthServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def image(self): - return api_image_connect.ImageServiceClientSync(address=self._baseurl, http_client=self._client) + return api_image_connect.ImageServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def ip(self): - return api_ip_connect.IPServiceClientSync(address=self._baseurl, http_client=self._client) + return api_ip_connect.IPServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def machine(self): - return api_machine_connect.MachineServiceClientSync(address=self._baseurl, http_client=self._client) + return api_machine_connect.MachineServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def method(self): - return api_method_connect.MethodServiceClientSync(address=self._baseurl, http_client=self._client) + return api_method_connect.MethodServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def network(self): - return api_network_connect.NetworkServiceClientSync(address=self._baseurl, http_client=self._client) + return api_network_connect.NetworkServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def partition(self): - return api_partition_connect.PartitionServiceClientSync(address=self._baseurl, http_client=self._client) + return api_partition_connect.PartitionServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def project(self): - return api_project_connect.ProjectServiceClientSync(address=self._baseurl, http_client=self._client) + return api_project_connect.ProjectServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def size(self): - return api_size_connect.SizeServiceClientSync(address=self._baseurl, http_client=self._client) + return api_size_connect.SizeServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def size_imageconstraint(self): - return api_size_imageconstraint_connect.SizeImageConstraintServiceClientSync(address=self._baseurl, http_client=self._client) + return api_size_imageconstraint_connect.SizeImageConstraintServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def size_reservation(self): - return api_size_reservation_connect.SizeReservationServiceClientSync(address=self._baseurl, http_client=self._client) + return api_size_reservation_connect.SizeReservationServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def tenant(self): - return api_tenant_connect.TenantServiceClientSync(address=self._baseurl, http_client=self._client) + return api_tenant_connect.TenantServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def token(self): - return api_token_connect.TokenServiceClientSync(address=self._baseurl, http_client=self._client) + return api_token_connect.TokenServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def user(self): - return api_user_connect.UserServiceClientSync(address=self._baseurl, http_client=self._client) + return api_user_connect.UserServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def version(self): - return api_version_connect.VersionServiceClientSync(address=self._baseurl, http_client=self._client) + return api_version_connect.VersionServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) class _Infrav2: - def __init__(self, baseurl: str, client: pyqwest.SyncClient = None): + def __init__(self, baseurl: str, client: pyqwest.SyncClient = None, interceptors: list = []): self._baseurl = baseurl self._client = client + self._interceptors = list(interceptors) def bmc(self): - return infra_bmc_connect.BMCServiceClientSync(address=self._baseurl, http_client=self._client) + return infra_bmc_connect.BMCServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def boot(self): - return infra_boot_connect.BootServiceClientSync(address=self._baseurl, http_client=self._client) + return infra_boot_connect.BootServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def component(self): - return infra_component_connect.ComponentServiceClientSync(address=self._baseurl, http_client=self._client) + return infra_component_connect.ComponentServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def event(self): - return infra_event_connect.EventServiceClientSync(address=self._baseurl, http_client=self._client) + return infra_event_connect.EventServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) def switch(self): - return infra_switch_connect.SwitchServiceClientSync(address=self._baseurl, http_client=self._client) + return infra_switch_connect.SwitchServiceClientSync(address=self._baseurl, http_client=self._client, interceptors=self._interceptors) diff --git a/python/metalstack/client/test_interceptor.py b/python/metalstack/client/test_interceptor.py new file mode 100644 index 00000000..45fca252 --- /dev/null +++ b/python/metalstack/client/test_interceptor.py @@ -0,0 +1,57 @@ +from __future__ import annotations + +from dataclasses import dataclass + +from connectrpc.interceptor import UnaryInterceptorSync +from google.protobuf.json_format import MessageToDict +from google.protobuf.message import Message + + +def _messages_equal(a: Message, b: Message) -> bool: + return MessageToDict(a, preserving_proto_field_name=True) == MessageToDict( + b, preserving_proto_field_name=True + ) + + +@dataclass +class RpcCall: + request: Message | None = None + response: Message | None = None + error: Exception | None = None + + +class TestClientInterceptor(UnaryInterceptorSync): + def __init__(self, calls: list[RpcCall] | None = None): + self._calls: list[RpcCall] = list(calls) if calls else [] + self._call_count: int = 0 + + def intercept_unary_sync(self, call_next, request, ctx): + if self._call_count >= len(self._calls): + raise AssertionError( + f"unexpected RPC call #{self._call_count}: {request.DESCRIPTOR.full_name}" + ) + + expected = self._calls[self._call_count] + self._call_count += 1 + + if expected.request is not None and not _messages_equal(expected.request, request): + import json + + raise AssertionError( + f"request mismatch for {request.DESCRIPTOR.full_name}\n" + f" expected: {json.dumps(MessageToDict(expected.request, preserving_proto_field_name=True), indent=2)}\n" + f" got: {json.dumps(MessageToDict(request, preserving_proto_field_name=True), indent=2)}" + ) + + if expected.error is not None: + raise expected.error + + return expected.response + + def assert_all_calls_used(self): + remaining = len(self._calls) - self._call_count + if remaining != 0: + raise AssertionError( + f"{remaining} expected RPC call(s) were not made " + f"(made {self._call_count} of {len(self._calls)})" + )