From 46db2f7b06da0c14401cf8971fab9e34f95588d2 Mon Sep 17 00:00:00 2001 From: Alex Mazzeo Date: Mon, 27 Jul 2026 13:42:11 -0700 Subject: [PATCH] Unify payload visitation into a single implementation by combining explicit roots and discovered system nexus roots --- scripts/gen_payload_visitor.py | 14 +-- temporalio/bridge/_visitor.py | 22 ++++ temporalio/nexus/system/__init__.py | 2 +- temporalio/nexus/system/_payload_visitor.py | 131 -------------------- 4 files changed, 24 insertions(+), 145 deletions(-) delete mode 100644 temporalio/nexus/system/_payload_visitor.py diff --git a/scripts/gen_payload_visitor.py b/scripts/gen_payload_visitor.py index c2bc15837..fa28cb455 100644 --- a/scripts/gen_payload_visitor.py +++ b/scripts/gen_payload_visitor.py @@ -425,28 +425,18 @@ def walk(self, desc: Descriptor) -> bool: def write_bridge_visitors() -> None: out_path = base_dir / "temporalio" / "bridge" / "_visitor.py" - # Build root descriptors: WorkflowActivation, WorkflowActivationCompletion, - # and all messages from selected API modules roots: list[Descriptor] = [ WorkflowActivation.DESCRIPTOR, WorkflowActivationCompletion.DESCRIPTOR, - ] + ] + discover_system_nexus_roots() code = VisitorGenerator().generate(roots) out_path.write_text(code) -def write_system_nexus_payload_visitors() -> None: - out_path = base_dir / "temporalio" / "nexus" / "system" / "_payload_visitor.py" - code = VisitorGenerator().generate(discover_system_nexus_roots()) - out_path.write_text(code) - - if __name__ == "__main__": print("Generating temporalio/bridge/_visitor.py...", file=sys.stderr) write_bridge_visitors() - print("Generating temporalio/nexus/system/_payload_visitor.py...", file=sys.stderr) - write_system_nexus_payload_visitors() subprocess.run( [ "uv", @@ -457,7 +447,6 @@ def write_system_nexus_payload_visitors() -> None: "I", "--fix", "temporalio/bridge/_visitor.py", - "temporalio/nexus/system/_payload_visitor.py", ], cwd=base_dir, check=True, @@ -469,7 +458,6 @@ def write_system_nexus_payload_visitors() -> None: "ruff", "format", "temporalio/bridge/_visitor.py", - "temporalio/nexus/system/_payload_visitor.py", ], cwd=base_dir, check=True, diff --git a/temporalio/bridge/_visitor.py b/temporalio/bridge/_visitor.py index a9956e12b..974b4f77d 100644 --- a/temporalio/bridge/_visitor.py +++ b/temporalio/bridge/_visitor.py @@ -549,3 +549,25 @@ async def _visit_coresdk_workflow_completion_WorkflowActivationCompletion( await self._visit_coresdk_workflow_completion_Success(fs, o.successful) elif o.HasField("failed"): await self._visit_coresdk_workflow_completion_Failure(fs, o.failed) + + async def _visit_temporal_api_common_v1_Header(self, fs: VisitorFunctions, o: Any): + for v in o.fields.values(): + await self._visit_temporal_api_common_v1_Payload(fs, v) + + async def _visit_temporal_api_workflowservice_v1_SignalWithStartWorkflowExecutionRequest( + self, fs: VisitorFunctions, o: Any + ): + if o.HasField("input"): + await self._visit_temporal_api_common_v1_Payloads(fs, o.input) + if o.HasField("signal_input"): + await self._visit_temporal_api_common_v1_Payloads(fs, o.signal_input) + if o.HasField("memo"): + await self._visit_temporal_api_common_v1_Memo(fs, o.memo) + if o.HasField("search_attributes"): + await self._visit_temporal_api_common_v1_SearchAttributes( + fs, o.search_attributes + ) + if o.HasField("header"): + await self._visit_temporal_api_common_v1_Header(fs, o.header) + if o.HasField("user_metadata"): + await self._visit_temporal_api_sdk_v1_UserMetadata(fs, o.user_metadata) diff --git a/temporalio/nexus/system/__init__.py b/temporalio/nexus/system/__init__.py index 14a43cb72..7c83229c1 100644 --- a/temporalio/nexus/system/__init__.py +++ b/temporalio/nexus/system/__init__.py @@ -106,7 +106,7 @@ async def _maybe_visit_payload( # pyright: ignore[reportUnusedFunction] payload_converter = _SystemNexusOuterPayloadConverter() value = payload_converter.from_payload(payload) - from ._payload_visitor import PayloadVisitor + from temporalio.bridge._visitor import PayloadVisitor await PayloadVisitor(skip_search_attributes=skip_search_attributes).visit( visitor_functions, value diff --git a/temporalio/nexus/system/_payload_visitor.py b/temporalio/nexus/system/_payload_visitor.py deleted file mode 100644 index 4f194168f..000000000 --- a/temporalio/nexus/system/_payload_visitor.py +++ /dev/null @@ -1,131 +0,0 @@ -from __future__ import annotations - -# This file is generated by gen_payload_visitor.py. Changes should be made there. -from typing import Any - -import temporalio.nexus.system -from temporalio.api.common.v1.message_pb2 import Payload -from temporalio.bridge._visitor_functions import ( - BoundedVisitorFunctions, - PayloadSequence, - VisitorFunctions, -) - - -class PayloadVisitor: - """A visitor for payloads. - Applies a function to every payload in a tree of messages. - """ - - def __init__( - self, - *, - skip_search_attributes: bool = False, - skip_headers: bool = False, - concurrency_limit: int = 1, - ): - """Creates a new payload visitor. - - Args: - skip_search_attributes: If True, search attributes are not visited. - skip_headers: If True, headers are not visited. - concurrency_limit: Maximum number of payload visits that may run - concurrently during a single call to visit(). Defaults to 1 - (sequential). - """ - if concurrency_limit < 1: - raise ValueError("concurrency_limit must be positive") - self.skip_search_attributes = skip_search_attributes - self.skip_headers = skip_headers - self._concurrency_limit = concurrency_limit - - async def visit(self, fs: VisitorFunctions, root: Any) -> None: - """Visits the given root message with the given function.""" - method_name = "_visit_" + root.DESCRIPTOR.full_name.replace(".", "_") - method = getattr(self, method_name, None) - if method is None: - raise ValueError(f"Unknown root message type: {root.DESCRIPTOR.full_name}") - if self._concurrency_limit == 1: - await method(fs, root) - return - - bounded = BoundedVisitorFunctions(fs, self._concurrency_limit) - try: - await method(bounded, root) - finally: - await bounded.drain() - - async def _visit_nexus_operation_input_payload( - self, - fs: VisitorFunctions, - endpoint: str, - payload: Payload, - ) -> None: - new_payload = await temporalio.nexus.system._maybe_visit_payload( - endpoint, - payload, - fs, - self.skip_search_attributes, - ) - if new_payload is None: - await self._visit_temporal_api_common_v1_Payload(fs, payload) - return - - if new_payload is not payload: - payload.CopyFrom(new_payload) - await fs.visit_system_nexus_envelope(payload) - - async def _visit_temporal_api_common_v1_Payload( - self, fs: VisitorFunctions, o: Payload - ): - await fs.visit_payload(o) - - async def _visit_temporal_api_common_v1_Payloads( - self, fs: VisitorFunctions, o: Any - ): - await fs.visit_payloads(o.payloads) - - async def _visit_payload_container(self, fs: VisitorFunctions, o: PayloadSequence): - await fs.visit_payloads(o) - - async def _visit_temporal_api_common_v1_Memo(self, fs: VisitorFunctions, o: Any): - for v in o.fields.values(): - await self._visit_temporal_api_common_v1_Payload(fs, v) - - async def _visit_temporal_api_common_v1_SearchAttributes( - self, fs: VisitorFunctions, o: Any - ): - if self.skip_search_attributes: - return - for v in o.indexed_fields.values(): - await self._visit_temporal_api_common_v1_Payload(fs, v) - - async def _visit_temporal_api_common_v1_Header(self, fs: VisitorFunctions, o: Any): - for v in o.fields.values(): - await self._visit_temporal_api_common_v1_Payload(fs, v) - - async def _visit_temporal_api_sdk_v1_UserMetadata( - self, fs: VisitorFunctions, o: Any - ): - if o.HasField("summary"): - await self._visit_temporal_api_common_v1_Payload(fs, o.summary) - if o.HasField("details"): - await self._visit_temporal_api_common_v1_Payload(fs, o.details) - - async def _visit_temporal_api_workflowservice_v1_SignalWithStartWorkflowExecutionRequest( - self, fs: VisitorFunctions, o: Any - ): - if o.HasField("input"): - await self._visit_temporal_api_common_v1_Payloads(fs, o.input) - if o.HasField("signal_input"): - await self._visit_temporal_api_common_v1_Payloads(fs, o.signal_input) - if o.HasField("memo"): - await self._visit_temporal_api_common_v1_Memo(fs, o.memo) - if o.HasField("search_attributes"): - await self._visit_temporal_api_common_v1_SearchAttributes( - fs, o.search_attributes - ) - if o.HasField("header"): - await self._visit_temporal_api_common_v1_Header(fs, o.header) - if o.HasField("user_metadata"): - await self._visit_temporal_api_sdk_v1_UserMetadata(fs, o.user_metadata)