diff --git a/temporalio/worker/_nexus.py b/temporalio/worker/_nexus.py index 9f21811a1..90ba40382 100644 --- a/temporalio/worker/_nexus.py +++ b/temporalio/worker/_nexus.py @@ -31,6 +31,7 @@ import temporalio.common import temporalio.converter import temporalio.nexus +import temporalio.nexus.system from temporalio.bridge._visitor import PayloadVisitor from temporalio.bridge._visitor_functions import PayloadSequence, VisitorFunctions from temporalio.bridge.worker import PollShutdownError @@ -428,7 +429,7 @@ async def _start_operation( _worker_shutdown_event=self._worker_shutdown_event, ).set() input = LazyValue( - serializer=_DummyPayloadSerializer( + serializer=_NexusPayloadSerializer( data_converter=self._data_converter, payload=start_request.payload, ), @@ -527,7 +528,7 @@ def _is_payload_validation_error(err: BaseException) -> TypeGuard[ApplicationErr @dataclass -class _DummyPayloadSerializer: +class _NexusPayloadSerializer: data_converter: temporalio.converter.DataConverter payload: temporalio.api.common.v1.Payload @@ -577,7 +578,13 @@ async def deserialize( ) from err try: - [input] = dc.payload_converter.from_payloads( + payload_converter = dc.payload_converter + if temporalio.nexus.system._is_system_payload(payload): + payload_converter = temporalio.nexus.system._get_payload_converter( + dc.payload_converter, + dc.failure_converter, + ) + [input] = payload_converter.from_payloads( [payload], type_hints=[as_type] if as_type else None, ) diff --git a/tests/nexus/test_temporal_system_nexus.py b/tests/nexus/test_temporal_system_nexus.py index c75fd64ff..1fa413452 100644 --- a/tests/nexus/test_temporal_system_nexus.py +++ b/tests/nexus/test_temporal_system_nexus.py @@ -6,6 +6,7 @@ from datetime import timedelta from typing import Any, cast +import nexusrpc import pytest from google.protobuf.descriptor import FieldDescriptor from google.protobuf.message import Message @@ -41,6 +42,7 @@ WorkflowInterceptorClassInput, WorkflowOutboundInterceptor, ) +from temporalio.worker._nexus import _NexusPayloadSerializer from temporalio.worker._workflow_instance import UnsandboxedWorkflowRunner from tests.test_extstore import InMemoryTestDriver @@ -166,6 +168,7 @@ def test_signal_with_start_serialization_context() -> None: class RejectOuterSystemNexusCodec(PayloadCodec): def __init__(self) -> None: self.encode_count = 0 + self.decode_count = 0 async def encode( self, payloads: Sequence[temporalio.api.common.v1.Payload] @@ -202,6 +205,7 @@ async def decode( raise RuntimeError( "outer system nexus envelope should not be codec decoded" ) + self.decode_count += 1 decoded.append(payload) return decoded @@ -333,6 +337,88 @@ def _new_unmarked_system_nexus_request_payload() -> temporalio.api.common.v1.Pay return payload +async def test_nexus_payload_serializer_decodes_system_input() -> None: + """A marked system request is decoded into its generated Nexus model.""" + data_converter = temporalio.converter.default() + request = workflow_service_models.SignalWithStartWorkflowRequest( + workflow="test-workflow", + args=["workflow-input"], + id="target-workflow-id", + task_queue="target-task-queue", + signal="test-signal", + namespace="target-namespace", + ) + payload = nexus_system._get_payload_converter( + data_converter.payload_converter, + data_converter.failure_converter, + ).to_payload(request) + assert payload is not None + assert payload.metadata[SYSTEM_NEXUS_PAYLOAD_METADATA_KEY] == b"true" + assert payload.metadata["encoding"] == b"binary/protobuf" + + decoded = await _NexusPayloadSerializer( + data_converter=data_converter, + payload=payload, + ).deserialize( + nexusrpc.Content(headers={}, data=b""), + as_type=workflow_service_models.SignalWithStartWorkflowRequest, + ) + + assert decoded == request + + +async def test_nexus_payload_serializer_codec_skips_outer_envelope() -> None: + """The codec decodes nested user payloads but not the outer system envelope.""" + codec = RejectOuterSystemNexusCodec() + data_converter = dataclasses.replace( + temporalio.converter.default(), + payload_codec=codec, + ) + request = workflow_service_models.SignalWithStartWorkflowRequest( + workflow="test-workflow", + args=["workflow-input"], + id="target-workflow-id", + task_queue="target-task-queue", + signal="test-signal", + namespace="target-namespace", + ) + payload = nexus_system._get_payload_converter( + data_converter.payload_converter, + data_converter.failure_converter, + ).to_payload(request) + assert payload is not None + + decoded = await _NexusPayloadSerializer( + data_converter=data_converter, + payload=payload, + ).deserialize( + nexusrpc.Content(headers={}, data=b""), + as_type=workflow_service_models.SignalWithStartWorkflowRequest, + ) + + assert codec.decode_count == 1 + assert codec.encode_count == 0 + assert decoded == request + + +async def test_nexus_payload_serializer_uses_user_converter() -> None: + """An unmarked Nexus input continues to use the configured user converter.""" + data_converter = temporalio.converter.default() + payload = data_converter.payload_converter.to_payload("ordinary-input") + assert payload is not None + assert SYSTEM_NEXUS_PAYLOAD_METADATA_KEY not in payload.metadata + + decoded = await _NexusPayloadSerializer( + data_converter=data_converter, + payload=payload, + ).deserialize( + nexusrpc.Content(headers={}, data=b""), + as_type=str, + ) + + assert decoded == "ordinary-input" + + async def test_schedule_marked_system_nexus_payload_ignores_endpoint() -> None: completion = _new_schedule_nexus_completion( "not-the-system-endpoint", diff --git a/tests/nexus/test_workflow_caller_errors.py b/tests/nexus/test_workflow_caller_errors.py index 5e883275c..78d5f9e67 100644 --- a/tests/nexus/test_workflow_caller_errors.py +++ b/tests/nexus/test_workflow_caller_errors.py @@ -44,7 +44,7 @@ from temporalio.service import RPCError, RPCStatusCode from temporalio.testing import WorkflowEnvironment from temporalio.worker import Worker -from temporalio.worker._nexus import _DummyPayloadSerializer +from temporalio.worker._nexus import _NexusPayloadSerializer from tests.helpers import LogCapturer, assert_eq_eventually from tests.helpers.nexus import make_nexus_endpoint_name @@ -884,7 +884,7 @@ def from_payloads( async def _deserialize_input(data_converter: DataConverter) -> Any: [payload] = DataConverter.default.payload_converter.to_payloads(["input"]) - serializer = _DummyPayloadSerializer(data_converter=data_converter, payload=payload) + serializer = _NexusPayloadSerializer(data_converter=data_converter, payload=payload) return await serializer.deserialize(nexusrpc.Content(headers={}, data=b""))