Skip to content

Commit 6ecff04

Browse files
Support payload-free system Nexus roots (#1783)
* Support payload-free system Nexus roots * Fixing time-skipping test error * Update visit exception thrown * sdk core update * Revert "sdk core update" This reverts commit d2f3219.
1 parent 5c0f0d0 commit 6ecff04

6 files changed

Lines changed: 110 additions & 33 deletions

File tree

‎scripts/gen_payload_visitor.py‎

Lines changed: 12 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -124,12 +124,20 @@ def generate(self, roots: list[Descriptor]) -> str:
124124
125125
The generated code defines async visitor functions for each reachable
126126
protobuf message type starting from WorkflowActivation, including support
127-
for repeated fields and map entries, and a convenience entrypoint
128-
function `visit`.
127+
for repeated fields and map entries. Payload-free roots get no-op methods
128+
so the `visit` entrypoint recognizes them as supported.
129129
"""
130130

131-
for r in roots:
132-
self.walk(r)
131+
for root in roots:
132+
if not self.walk(root):
133+
self.methods.append(
134+
f"""\
135+
async def _visit_{name_for(root)}(
136+
self, fs: VisitorFunctions, o: Any
137+
) -> None:
138+
pass
139+
"""
140+
)
133141

134142
header = """
135143
from __future__ import annotations

‎temporalio/bridge/_visitor.py‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -645,3 +645,8 @@ async def _visit_temporal_api_workflowservice_v1_SignalWithStartWorkflowExecutio
645645
await self._visit_temporal_api_common_v1_Header(fs, o.header)
646646
if o.HasField("user_metadata"):
647647
await self._visit_temporal_api_sdk_v1_UserMetadata(fs, o.user_metadata)
648+
649+
async def _visit_temporal_api_workflowservice_v1_SignalWithStartWorkflowExecutionResponse(
650+
self, fs: VisitorFunctions, o: Any
651+
) -> None:
652+
pass

‎temporalio/nexus/system/__init__.py‎

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515
import temporalio.api.common.v1
1616
import temporalio.common
1717
import temporalio.converter
18+
import temporalio.exceptions
1819
from temporalio.bridge._visitor_functions import VisitorFunctions
1920
from temporalio.converter import BinaryProtoPayloadConverter, CompositePayloadConverter
2021
from temporalio.converter._payload_converter import (
@@ -154,7 +155,14 @@ async def maybe_visit_payload(
154155

155156
payload_visitor = PayloadVisitor(skip_search_attributes=skip_search_attributes)
156157
checkpoint = visitor_functions.checkpoint()
157-
await payload_visitor.visit(visitor_functions, value)
158+
try:
159+
await payload_visitor.visit(visitor_functions, value)
160+
except ValueError as err:
161+
if not str(err).startswith("Unknown root message type: "):
162+
raise
163+
raise temporalio.exceptions.ApplicationError(
164+
f"Unknown Temporal system payload: {value.DESCRIPTOR.full_name}"
165+
) from err
158166
if checkpoint is not None:
159167
await visitor_functions.drain_since(checkpoint)
160168
return payload_converter.to_payload(value)

‎tests/__init__.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1 +1 @@
1-
DEV_SERVER_DOWNLOAD_VERSION = "v1.7.4-standalone-nexus-operations"
1+
DEV_SERVER_DOWNLOAD_VERSION = "v1.8.3-server-1.32.0-162.0"

‎tests/nexus/test_temporal_operation.py‎

Lines changed: 13 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -44,9 +44,6 @@
4444
# See https://github.com/temporalio/sdk-python/issues/1704.
4545
pytestmark = pytest.mark.requires_local_server
4646

47-
# Query response links require a newer server than the shared test environment.
48-
_QUERY_LINK_DEV_SERVER_DOWNLOAD_VERSION = "v1.8.3-server-1.32.0-162.0"
49-
5047

5148
@dataclass
5249
class Input:
@@ -854,14 +851,7 @@ async def run(self, input: Input) -> bool:
854851
return await client.execute_operation(TestService.query_op, input.value)
855852

856853

857-
async def test_temporal_operation_query_workflow() -> None:
858-
async with await WorkflowEnvironment.start_local(
859-
dev_server_download_version=_QUERY_LINK_DEV_SERVER_DOWNLOAD_VERSION
860-
) as env:
861-
await _assert_temporal_operation_query_workflow(env.client, env)
862-
863-
864-
async def _assert_temporal_operation_query_workflow(
854+
async def test_temporal_operation_query_workflow(
865855
client: Client, env: WorkflowEnvironment
866856
) -> None:
867857
task_queue = str(uuid.uuid4())
@@ -900,15 +890,17 @@ async def _assert_temporal_operation_query_workflow(
900890
target_history = await target_handle.fetch_history()
901891
assert not any(event.links for event in target_history.events)
902892

903-
assert target_handle.result_run_id is not None
904-
assert Link(
905-
workflow=Link.Workflow(
906-
namespace=client.namespace,
907-
workflow_id=target_workflow_id,
908-
run_id=target_handle.result_run_id,
909-
reason="Query processed",
910-
)
911-
) in list(completed_event.links)
893+
# The Java time-skipping test server does not return Nexus operation links.
894+
if not env.supports_time_skipping:
895+
assert target_handle.result_run_id is not None
896+
assert Link(
897+
workflow=Link.Workflow(
898+
namespace=client.namespace,
899+
workflow_id=target_workflow_id,
900+
run_id=target_handle.result_run_id,
901+
reason="Query processed",
902+
)
903+
) in list(completed_event.links)
912904
finally:
913905
await target_handle.cancel()
914906

@@ -1314,14 +1306,8 @@ async def test_temporal_operation_start_activity_raises_error(
13141306
id=str(uuid.uuid4()),
13151307
)
13161308

1317-
operation_err = err.value.__cause__
1318-
assert isinstance(operation_err, temporalio.exceptions.ApplicationError)
1319-
assert operation_err.type == "OperationError"
1320-
assert "nexus operation completed unsuccessfully" in str(operation_err)
1321-
1322-
application_err = operation_err.__cause__
1309+
application_err = err.value.__cause__
13231310
assert isinstance(application_err, temporalio.exceptions.ApplicationError)
1324-
13251311
assert application_err.type == "test-activity-error-type"
13261312
assert "test-activity-error-message" in str(application_err)
13271313
assert application_err.__cause__ is None

‎tests/worker/test_visitor.py‎

Lines changed: 70 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@
99
import temporalio.api.workflowservice.v1.request_response_pb2 as workflowservice_pb2
1010
import temporalio.bridge.worker
1111
import temporalio.converter
12+
import temporalio.exceptions
1213
import temporalio.nexus.system as nexus_system
1314
from temporalio.api.common.v1.message_pb2 import (
1415
Payload,
@@ -266,6 +267,75 @@ async def visit_system_nexus_envelope(self, payload: Payload) -> None:
266267
assert visitor.system_envelope_count == 1
267268

268269

270+
async def test_system_nexus_envelope_without_payloads_is_visited():
271+
class SystemNexusVisitor(Visitor):
272+
def __init__(self) -> None:
273+
self.system_envelope_count = 0
274+
275+
async def visit_system_nexus_envelope(self, payload: Payload) -> None:
276+
_ = payload
277+
self.system_envelope_count += 1
278+
279+
response = workflowservice_pb2.SignalWithStartWorkflowExecutionResponse(
280+
run_id="test-run-id"
281+
)
282+
data_converter = temporalio.converter.default()
283+
payload_converter = nexus_system._get_payload_converter(
284+
data_converter.payload_converter,
285+
data_converter.failure_converter,
286+
)
287+
system_payload = payload_converter.to_payload(response)
288+
assert system_payload is not None
289+
completion = WorkflowActivationCompletion(
290+
run_id="3",
291+
successful=Success(
292+
commands=[
293+
WorkflowCommand(
294+
update_response=UpdateResponse(completed=system_payload),
295+
)
296+
]
297+
),
298+
)
299+
visitor = SystemNexusVisitor()
300+
301+
await PayloadVisitor().visit(visitor, completion)
302+
303+
completed = completion.successful.commands[0].update_response.completed
304+
assert payload_converter.from_payload(completed) == response
305+
assert visitor.system_envelope_count == 1
306+
307+
308+
async def test_unknown_system_nexus_payload_raises_application_error():
309+
data_converter = temporalio.converter.default()
310+
payload_converter = nexus_system._get_payload_converter(
311+
data_converter.payload_converter,
312+
data_converter.failure_converter,
313+
)
314+
system_payload = payload_converter.to_payload(
315+
workflowservice_pb2.StartWorkflowExecutionResponse(run_id="test-run-id")
316+
)
317+
assert system_payload is not None
318+
completion = WorkflowActivationCompletion(
319+
run_id="3",
320+
successful=Success(
321+
commands=[
322+
WorkflowCommand(
323+
update_response=UpdateResponse(completed=system_payload),
324+
)
325+
]
326+
),
327+
)
328+
329+
with pytest.raises(temporalio.exceptions.ApplicationError) as err:
330+
await PayloadVisitor().visit(Visitor(), completion)
331+
332+
assert (
333+
err.value.message == "Unknown Temporal system payload: "
334+
"temporal.api.workflowservice.v1.StartWorkflowExecutionResponse"
335+
)
336+
assert not err.value.non_retryable
337+
338+
269339
async def test_concurrent_throughput():
270340
"""Demonstrate that concurrent visitation is faster than serialized for I/O-bound codecs."""
271341
N_CMDS = 10

0 commit comments

Comments
 (0)