From 4cf4eaa960980487a912695f09009753246c3376 Mon Sep 17 00:00:00 2001 From: Tim Conley Date: Tue, 21 Jul 2026 12:59:46 -0700 Subject: [PATCH] Mark system Nexus envelope payloads --- CHANGELOG.md | 3 ++ scripts/gen_payload_visitor.py | 32 +++++++------- temporalio/bridge/_visitor.py | 24 +++++----- temporalio/nexus/system/__init__.py | 31 +++++++++++-- temporalio/nexus/system/_payload_visitor.py | 22 +++++---- tests/nexus/test_temporal_system_nexus.py | 25 ++++++++--- tests/worker/test_visitor.py | 49 +++++++++++++++++++++ 7 files changed, 136 insertions(+), 50 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 31036bfed..7a2862b17 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -45,6 +45,9 @@ to include examples, links to docs, or any other relevant information. ### Fixed +- Marked system Nexus envelope payloads so nested payloads can be detected and + visited after the envelope is already stored as a payload. + ### Security ## [1.30.0] - 2026-07-01 diff --git a/scripts/gen_payload_visitor.py b/scripts/gen_payload_visitor.py index efe9c0df2..4729af5ee 100644 --- a/scripts/gen_payload_visitor.py +++ b/scripts/gen_payload_visitor.py @@ -188,22 +188,9 @@ async def visit( 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) + await self._visit_temporal_api_common_v1_Payload(fs, payload) """ @@ -219,7 +206,18 @@ def __init__(self): self.methods: list[str] = [ """\ async def _visit_temporal_api_common_v1_Payload(self, fs: VisitorFunctions, o: Payload): - await fs.visit_payload(o) + new_payload = await temporalio.nexus.system.maybe_visit_payload( + o, + fs, + self.skip_search_attributes, + ) + if new_payload is None: + await fs.visit_payload(o) + return + + if new_payload is not o: + o.CopyFrom(new_payload) + await fs.visit_system_nexus_envelope(o) """, """\ async def _visit_temporal_api_common_v1_Payloads(self, fs: VisitorFunctions, o: Any): @@ -403,11 +401,11 @@ def walk(self, desc: Descriptor) -> bool: ) ) elif item[0] == "system_nexus": - _, field_name, endpoint_expr, payload_expr = item + _, field_name, _endpoint_expr, payload_expr = item lines.append( f' if o.HasField("{field_name}"):\n' " await self._visit_nexus_operation_input_payload(\n" - f" fs, {endpoint_expr}, {payload_expr}\n" + f" fs, {payload_expr}\n" " )" ) else: # oneof_group diff --git a/temporalio/bridge/_visitor.py b/temporalio/bridge/_visitor.py index 4e258b9a1..fce017a63 100644 --- a/temporalio/bridge/_visitor.py +++ b/temporalio/bridge/_visitor.py @@ -58,27 +58,25 @@ async def visit(self, fs: VisitorFunctions, root: Any) -> None: async def _visit_nexus_operation_input_payload( self, fs: VisitorFunctions, - endpoint: str, payload: Payload, ) -> None: + await self._visit_temporal_api_common_v1_Payload(fs, payload) + + async def _visit_temporal_api_common_v1_Payload( + self, fs: VisitorFunctions, o: Payload + ): new_payload = await temporalio.nexus.system.maybe_visit_payload( - endpoint, - payload, + o, fs, self.skip_search_attributes, ) if new_payload is None: - await self._visit_temporal_api_common_v1_Payload(fs, payload) + await fs.visit_payload(o) 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) + if new_payload is not o: + o.CopyFrom(new_payload) + await fs.visit_system_nexus_envelope(o) async def _visit_temporal_api_common_v1_Payloads( self, fs: VisitorFunctions, o: Any @@ -474,7 +472,7 @@ async def _visit_coresdk_workflow_commands_ScheduleNexusOperation( self, fs: VisitorFunctions, o: Any ): if o.HasField("input"): - await self._visit_nexus_operation_input_payload(fs, o.endpoint, o.input) + await self._visit_nexus_operation_input_payload(fs, o.input) async def _visit_coresdk_workflow_commands_WorkflowCommand( self, fs: VisitorFunctions, o: Any diff --git a/temporalio/nexus/system/__init__.py b/temporalio/nexus/system/__init__.py index 21c5a1408..98691f942 100644 --- a/temporalio/nexus/system/__init__.py +++ b/temporalio/nexus/system/__init__.py @@ -2,12 +2,18 @@ from __future__ import annotations +from collections.abc import Sequence +from typing import Any + import temporalio.api.common.v1 +import temporalio.common import temporalio.converter from temporalio.bridge._visitor_functions import VisitorFunctions from temporalio.converter import BinaryProtoPayloadConverter, CompositePayloadConverter TEMPORAL_SYSTEM_ENDPOINT = "__temporal_system" +_SYSTEM_PAYLOAD_METADATA_KEY = "__temporal_system_payload" +_SYSTEM_PAYLOAD_METADATA_VALUE = b"true" class SystemNexusPayloadConverter(CompositePayloadConverter): @@ -17,20 +23,39 @@ def __init__(self) -> None: """Create a payload converter for system Nexus outer envelopes.""" super().__init__(BinaryProtoPayloadConverter()) + def to_payloads( + self, values: Sequence[Any] + ) -> list[temporalio.api.common.v1.Payload]: + """See base class.""" + payloads = super().to_payloads(values) + for value, payload in zip(values, payloads): + if isinstance(value, temporalio.common.RawValue): + continue + payload.metadata[_SYSTEM_PAYLOAD_METADATA_KEY] = ( + _SYSTEM_PAYLOAD_METADATA_VALUE + ) + return payloads + def is_system_endpoint(endpoint: str) -> bool: """Return whether a Nexus endpoint is the Temporal system endpoint.""" return endpoint == TEMPORAL_SYSTEM_ENDPOINT +def _is_system_payload(payload: temporalio.api.common.v1.Payload) -> bool: + return ( + payload.metadata.get(_SYSTEM_PAYLOAD_METADATA_KEY) + == _SYSTEM_PAYLOAD_METADATA_VALUE + ) + + async def maybe_visit_payload( - endpoint: str, payload: temporalio.api.common.v1.Payload, visitor_functions: VisitorFunctions, skip_search_attributes: bool, ) -> temporalio.api.common.v1.Payload | None: - """Visit nested payloads if the payload is for the Temporal system endpoint.""" - if not is_system_endpoint(endpoint): + """Visit nested payloads if the payload is a Temporal system Nexus envelope.""" + if not _is_system_payload(payload): return None payload_converter = get_payload_converter() diff --git a/temporalio/nexus/system/_payload_visitor.py b/temporalio/nexus/system/_payload_visitor.py index 5b4178ff1..bd4f86e71 100644 --- a/temporalio/nexus/system/_payload_visitor.py +++ b/temporalio/nexus/system/_payload_visitor.py @@ -58,27 +58,25 @@ async def visit(self, fs: VisitorFunctions, root: Any) -> None: async def _visit_nexus_operation_input_payload( self, fs: VisitorFunctions, - endpoint: str, payload: Payload, ) -> None: + await self._visit_temporal_api_common_v1_Payload(fs, payload) + + async def _visit_temporal_api_common_v1_Payload( + self, fs: VisitorFunctions, o: Payload + ): new_payload = await temporalio.nexus.system.maybe_visit_payload( - endpoint, - payload, + o, fs, self.skip_search_attributes, ) if new_payload is None: - await self._visit_temporal_api_common_v1_Payload(fs, payload) + await fs.visit_payload(o) 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) + if new_payload is not o: + o.CopyFrom(new_payload) + await fs.visit_system_nexus_envelope(o) async def _visit_temporal_api_common_v1_Payloads( self, fs: VisitorFunctions, o: Any diff --git a/tests/nexus/test_temporal_system_nexus.py b/tests/nexus/test_temporal_system_nexus.py index b689ee8d9..7253da5fd 100644 --- a/tests/nexus/test_temporal_system_nexus.py +++ b/tests/nexus/test_temporal_system_nexus.py @@ -34,6 +34,7 @@ from tests.test_extstore import InMemoryTestDriver interceptor_traces: list[tuple[str, object]] = [] +SYSTEM_NEXUS_PAYLOAD_METADATA_KEY = "__temporal_system_payload" @workflow.defn @@ -191,9 +192,15 @@ def _new_system_nexus_request_payload() -> temporalio.api.common.v1.Payload: return payload -async def test_schedule_system_nexus_endpoint_ignores_operation_registry() -> None: +def _new_unmarked_system_nexus_request_payload() -> temporalio.api.common.v1.Payload: + payload = _new_system_nexus_request_payload() + del payload.metadata[SYSTEM_NEXUS_PAYLOAD_METADATA_KEY] + return payload + + +async def test_schedule_marked_system_nexus_payload_ignores_endpoint() -> None: completion = _new_schedule_nexus_completion( - nexus_system.TEMPORAL_SYSTEM_ENDPOINT, + "not-the-system-endpoint", _new_system_nexus_request_payload(), ) visitor = _MarkingPayloadVisitor() @@ -211,10 +218,12 @@ async def test_schedule_system_nexus_endpoint_ignores_operation_registry() -> No assert visitor.system_envelope_count == 1 -async def test_schedule_non_system_nexus_visits_input_as_regular_payload() -> None: +async def test_schedule_unmarked_system_nexus_payload_visits_input_as_regular_payload() -> ( + None +): completion = _new_schedule_nexus_completion( - "not-the-system-endpoint", - _new_system_nexus_request_payload(), + nexus_system.TEMPORAL_SYSTEM_ENDPOINT, + _new_unmarked_system_nexus_request_payload(), ) visitor = _MarkingPayloadVisitor() @@ -222,6 +231,11 @@ async def test_schedule_non_system_nexus_visits_input_as_regular_payload() -> No schedule = completion.successful.commands[0].schedule_nexus_operation assert schedule.input.metadata["visited"] == b"true" + decoded = nexus_system.get_payload_converter().from_payload(schedule.input) + assert isinstance( + decoded, workflowservice_pb2.SignalWithStartWorkflowExecutionRequest + ) + assert "visited" not in decoded.input.payloads[0].metadata assert visitor.visited_payload_count == 1 assert visitor.system_envelope_count == 0 @@ -345,6 +359,7 @@ def test_system_nexus_proto_roundtrip(message_type: type[Message]) -> None: assert payload is not None assert payload.metadata["encoding"] == b"binary/protobuf" assert payload.metadata["messageType"] == message_type.DESCRIPTOR.full_name.encode() + assert payload.metadata[SYSTEM_NEXUS_PAYLOAD_METADATA_KEY] == b"true" roundtripped = payload_converter.from_payload(payload, message_type) assert isinstance(roundtripped, message_type) assert roundtripped == proto_value diff --git a/tests/worker/test_visitor.py b/tests/worker/test_visitor.py index bd4004625..09cf80f49 100644 --- a/tests/worker/test_visitor.py +++ b/tests/worker/test_visitor.py @@ -212,6 +212,55 @@ async def test_visit_payloads_on_other_commands(): assert ur.completed.metadata["visited"] +async def test_system_nexus_envelope_is_detected_in_generic_payload_field(): + class SystemNexusVisitor(Visitor): + def __init__(self) -> None: + self.visited_payload_count = 0 + self.system_envelope_count = 0 + + async def visit_payload(self, payload: Payload) -> None: + self.visited_payload_count += 1 + await super().visit_payload(payload) + + async def visit_payloads(self, payloads: MutableSequence[Payload]) -> None: + self.visited_payload_count += len(payloads) + await super().visit_payloads(payloads) + + async def visit_system_nexus_envelope(self, payload: Payload) -> None: + _ = payload + self.system_envelope_count += 1 + + system_request = workflowservice_pb2.SignalWithStartWorkflowExecutionRequest( + input=Payloads(payloads=[Payload(data=b"workflow-input")]), + ) + system_payload = nexus_system.get_payload_converter().to_payload(system_request) + assert system_payload is not None + comp = WorkflowActivationCompletion( + run_id="3", + successful=Success( + commands=[ + WorkflowCommand( + update_response=UpdateResponse(completed=system_payload), + ) + ] + ), + ) + visitor = SystemNexusVisitor() + + await PayloadVisitor().visit(visitor, comp) + + completed = comp.successful.commands[0].update_response.completed + assert completed.metadata["__temporal_system_payload"] == b"true" + assert "visited" not in completed.metadata + decoded = nexus_system.get_payload_converter().from_payload(completed) + assert isinstance( + decoded, workflowservice_pb2.SignalWithStartWorkflowExecutionRequest + ) + assert decoded.input.payloads[0].metadata["visited"] == b"True" + assert visitor.visited_payload_count == 1 + assert visitor.system_envelope_count == 1 + + async def test_concurrent_throughput(): """Demonstrate that concurrent visitation is faster than serialized for I/O-bound codecs.""" N_CMDS = 10