Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
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
13 changes: 10 additions & 3 deletions temporalio/worker/_nexus.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
),
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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,
)
Expand Down
86 changes: 86 additions & 0 deletions tests/nexus/test_temporal_system_nexus.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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",
Expand Down
4 changes: 2 additions & 2 deletions tests/nexus/test_workflow_caller_errors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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""))


Expand Down
Loading