Skip to content

Commit ff318a5

Browse files
committed
Adapt System Nexus tracing to generated request types
1 parent b26919f commit ff318a5

6 files changed

Lines changed: 195 additions & 35 deletions

File tree

‎temporalio/converter/_payload_converter.py‎

Lines changed: 13 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -613,14 +613,26 @@ def wrap(payload_converter: PayloadConverter) -> PayloadConverter:
613613
return _TemporalTransferTypePayloadConverter(payload_converter)
614614

615615
def to_payloads(
616-
self, values: Sequence[Any]
616+
self,
617+
values: Sequence[Any],
618+
headers: Mapping[str, temporalio.api.common.v1.Payload] | None = None,
617619
) -> list[temporalio.api.common.v1.Payload]:
618620
"""See base class."""
619621
transfer_type_values: list[Any] = []
620622
for value in values:
621623
converter = _get_transfer_type_converter(type(value))
622624
if converter is not None:
623625
value = converter.to_transfer_type(value)
626+
if (
627+
headers
628+
and isinstance(value, google.protobuf.message.Message)
629+
and "header" in value.DESCRIPTOR.fields_by_name
630+
):
631+
# System Nexus starts with generated models, so headers can only be
632+
# applied after conversion to a request protobuf and before encoding.
633+
temporalio.common._apply_headers(
634+
headers, getattr(value, "header").fields
635+
)
624636
transfer_type_values.append(value)
625637
return self._inner_payload_converter.to_payloads(transfer_type_values)
626638

‎temporalio/nexus/system/__init__.py‎

Lines changed: 13 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -12,8 +12,6 @@
1212
from dataclasses import dataclass
1313
from typing import Any
1414

15-
from google.protobuf.message import Message
16-
1715
import temporalio.api.common.v1
1816
import temporalio.common
1917
import temporalio.converter
@@ -93,27 +91,31 @@ class _SystemNexusPayloadConverter(temporalio.converter.PayloadConverter):
9391
"""Payload converter for system Nexus outer envelopes."""
9492

9593
_user_converters: _SystemNexusUserConverters
96-
_outer_payload_converter: temporalio.converter.PayloadConverter
94+
_outer_payload_converter: _TemporalTransferTypePayloadConverter
95+
_headers: Mapping[str, temporalio.api.common.v1.Payload] | None
9796

9897
def __init__(
9998
self,
10099
user_payload_converter: temporalio.converter.PayloadConverter,
101100
user_failure_converter: temporalio.converter.FailureConverter,
101+
headers: Mapping[str, temporalio.api.common.v1.Payload] | None = None,
102102
) -> None:
103103
"""Create a payload converter for system Nexus outer envelopes."""
104104
self._user_converters = _SystemNexusUserConverters(
105105
user_payload_converter, user_failure_converter
106106
)
107-
self._outer_payload_converter = _TemporalTransferTypePayloadConverter.wrap(
107+
108+
self._outer_payload_converter = _TemporalTransferTypePayloadConverter(
108109
_SystemNexusOuterPayloadConverter()
109110
)
111+
self._headers = headers
110112

111113
def to_payloads(
112114
self, values: Sequence[Any]
113115
) -> list[temporalio.api.common.v1.Payload]:
114116
"""See base class."""
115117
with _user_converter_context(self._user_converters):
116-
return self._outer_payload_converter.to_payloads(values)
118+
return self._outer_payload_converter.to_payloads(values, self._headers)
117119

118120
def from_payloads(
119121
self,
@@ -134,23 +136,13 @@ def is_system_endpoint(endpoint: str) -> bool:
134136
return endpoint == TEMPORAL_SYSTEM_ENDPOINT
135137

136138

137-
def _apply_headers_to_request(
138-
request: Message,
139-
headers: Mapping[str, temporalio.api.common.v1.Payload],
140-
) -> None:
141-
"""Apply headers to a system request when it supports Temporal headers."""
142-
if not headers or "header" not in request.DESCRIPTOR.fields_by_name:
143-
return
144-
request_header = getattr(request, "header")
145-
for key, payload in headers.items():
146-
request_header.fields[key].CopyFrom(payload)
147-
148-
149139
def _is_system_payload(payload: temporalio.api.common.v1.Payload) -> bool:
150140
return (
151141
payload.metadata.get(_SYSTEM_PAYLOAD_METADATA_KEY)
152142
== _SYSTEM_PAYLOAD_METADATA_VALUE
153143
)
144+
145+
154146
async def maybe_visit_payload(
155147
payload: temporalio.api.common.v1.Payload,
156148
visitor_functions: VisitorFunctions,
@@ -175,9 +167,12 @@ async def maybe_visit_payload(
175167
def _get_payload_converter( # pyright: ignore[reportUnusedFunction]
176168
user_payload_converter: temporalio.converter.PayloadConverter,
177169
user_failure_converter: temporalio.converter.FailureConverter,
170+
headers: Mapping[str, temporalio.api.common.v1.Payload] | None = None,
178171
) -> temporalio.converter.PayloadConverter:
179172
"""Return the fixed payload converter for system Nexus outer envelopes."""
180-
return _SystemNexusPayloadConverter(user_payload_converter, user_failure_converter)
173+
return _SystemNexusPayloadConverter(
174+
user_payload_converter, user_failure_converter, headers
175+
)
181176

182177

183178
def _get_serialization_context( # pyright: ignore[reportUnusedFunction]

‎temporalio/worker/_workflow_instance.py‎

Lines changed: 43 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -2189,22 +2189,51 @@ async def operation_handle_fn() -> OutputT:
21892189
async def _outbound_start_system_nexus_operation(
21902190
self, input: StartSystemNexusOperationInput[Any, OutputT]
21912191
) -> _NexusOperationHandle[OutputT]:
2192-
temporalio.nexus.system._apply_headers_to_request(input.input, input.headers)
2193-
return await self._outbound_start_nexus_operation(
2194-
StartNexusOperationInput(
2195-
endpoint=temporalio.nexus.system.TEMPORAL_SYSTEM_ENDPOINT,
2196-
service=input.service,
2197-
operation=input.operation,
2198-
input=input.input,
2199-
output_type=input.output_type,
2200-
schedule_to_close_timeout=input.schedule_to_close_timeout,
2201-
schedule_to_start_timeout=input.schedule_to_start_timeout,
2202-
start_to_close_timeout=input.start_to_close_timeout,
2203-
cancellation_type=input.cancellation_type,
2204-
headers=None,
2205-
summary=input.summary,
2192+
nexus_input = StartNexusOperationInput(
2193+
endpoint=temporalio.nexus.system.TEMPORAL_SYSTEM_ENDPOINT,
2194+
service=input.service,
2195+
operation=input.operation,
2196+
input=input.input,
2197+
output_type=input.output_type,
2198+
schedule_to_close_timeout=input.schedule_to_close_timeout,
2199+
schedule_to_start_timeout=input.schedule_to_start_timeout,
2200+
start_to_close_timeout=input.start_to_close_timeout,
2201+
cancellation_type=input.cancellation_type,
2202+
headers=None,
2203+
summary=input.summary,
2204+
)
2205+
handle: _NexusOperationHandle[OutputT]
2206+
2207+
async def operation_handle_fn() -> OutputT:
2208+
return cast(
2209+
OutputT,
2210+
await self._await_temporal_operation(
2211+
handle._result_fut,
2212+
lambda _err, command: handle._apply_cancel_command(command),
2213+
),
22062214
)
2215+
2216+
payload_converter = temporalio.nexus.system._get_payload_converter(
2217+
self._workflow_context_payload_converter,
2218+
self._workflow_context_failure_converter,
2219+
input.headers,
22072220
)
2221+
handle = _NexusOperationHandle(
2222+
self,
2223+
self._next_seq("nexus_operation"),
2224+
nexus_input,
2225+
operation_handle_fn(),
2226+
payload_converter,
2227+
)
2228+
handle._apply_schedule_command()
2229+
self._pending_nexus_operations[handle._seq] = handle
2230+
2231+
await self._await_temporal_operation(
2232+
handle._start_fut,
2233+
lambda _err, command: handle._apply_cancel_command(command),
2234+
reraise_on_workflow_cancellation=True,
2235+
)
2236+
return handle
22082237

22092238
#### Miscellaneous helpers ####
22102239
# These are in alphabetical order.

‎tests/contrib/opentelemetry/test_opentelemetry.py‎

Lines changed: 60 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -88,6 +88,34 @@ class TracingWorkflowActionActivity:
8888
fail_on_non_replay_before_complete: bool = False
8989

9090

91+
@workflow.defn
92+
class LegacySignalWithStartHeaderWorkflow:
93+
def __init__(self) -> None:
94+
self._signaled = False
95+
96+
@workflow.run
97+
async def run(self) -> bool:
98+
await workflow.wait_condition(lambda: self._signaled)
99+
return "_tracer-data" in workflow.info().headers
100+
101+
@workflow.signal
102+
def notify(self) -> None:
103+
self._signaled = True
104+
105+
106+
@workflow.defn
107+
class LegacySignalWithStartCallerWorkflow:
108+
@workflow.run
109+
async def run(self, target_id: str, task_queue: str) -> str:
110+
handle = await workflow.signal_with_start_workflow(
111+
LegacySignalWithStartHeaderWorkflow.run,
112+
id=target_id,
113+
task_queue=task_queue,
114+
signal=LegacySignalWithStartHeaderWorkflow.notify,
115+
)
116+
return handle.id
117+
118+
91119
@dataclass
92120
class TracingWorkflowActionContinueAsNew:
93121
param: TracingWorkflowParam
@@ -229,6 +257,38 @@ def update_validator(self) -> None:
229257
pass
230258

231259

260+
async def test_legacy_otel_workflow_signal_with_start_propagates_trace_headers(
261+
client: Client, env: WorkflowEnvironment
262+
):
263+
if env.supports_time_skipping:
264+
pytest.skip("Nexus tests don't work with the Java test server")
265+
provider = TracerProvider()
266+
tracer = provider.get_tracer(__name__)
267+
config = client.config()
268+
config["interceptors"] = [TracingInterceptor(tracer)]
269+
client = Client(**config)
270+
271+
async with Worker(
272+
client,
273+
task_queue=f"signal-with-start-{uuid.uuid4()}",
274+
workflows=[
275+
LegacySignalWithStartCallerWorkflow,
276+
LegacySignalWithStartHeaderWorkflow,
277+
],
278+
workflow_runner=UnsandboxedWorkflowRunner(),
279+
) as worker:
280+
target_id = f"signal-with-start-target-{uuid.uuid4()}"
281+
with tracer.start_as_current_span("signal-with-start"):
282+
caller = await client.start_workflow(
283+
LegacySignalWithStartCallerWorkflow.run,
284+
args=[target_id, worker.task_queue],
285+
id=f"signal-with-start-caller-{uuid.uuid4()}",
286+
task_queue=worker.task_queue,
287+
)
288+
assert await caller.result() == target_id
289+
assert await client.get_workflow_handle(target_id).result() is True
290+
291+
232292
async def test_opentelemetry_tracing(client: Client, env: WorkflowEnvironment):
233293
# TODO(cretz): Fix
234294
if env.supports_time_skipping:

‎tests/contrib/opentelemetry/test_opentelemetry_plugin.py‎

Lines changed: 57 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -122,6 +122,34 @@ async def run(self):
122122
return
123123

124124

125+
@workflow.defn
126+
class SignalWithStartHeaderWorkflow:
127+
def __init__(self) -> None:
128+
self._signaled = False
129+
130+
@workflow.run
131+
async def run(self) -> bool:
132+
await workflow.wait_condition(lambda: self._signaled)
133+
return "_tracer-data" in workflow.info().headers
134+
135+
@workflow.signal
136+
def notify(self) -> None:
137+
self._signaled = True
138+
139+
140+
@workflow.defn
141+
class SignalWithStartCallerWorkflow:
142+
@workflow.run
143+
async def run(self, target_id: str, task_queue: str) -> str:
144+
handle = await workflow.signal_with_start_workflow(
145+
SignalWithStartHeaderWorkflow.run,
146+
id=target_id,
147+
task_queue=task_queue,
148+
signal=SignalWithStartHeaderWorkflow.notify,
149+
)
150+
return handle.id
151+
152+
125153
async def test_otel_tracing_basic(client: Client, reset_otel_tracer_provider: Any): # type: ignore[reportUnusedParameter]
126154
exporter = InMemorySpanExporter()
127155
provider = create_tracer_provider()
@@ -169,6 +197,35 @@ async def test_otel_tracing_basic(client: Client, reset_otel_tracer_provider: An
169197
)
170198

171199

200+
async def test_otel_workflow_signal_with_start_propagates_trace_headers(
201+
client: Client,
202+
env: WorkflowEnvironment,
203+
reset_otel_tracer_provider: Any, # type: ignore[reportUnusedParameter]
204+
):
205+
if env.supports_time_skipping:
206+
pytest.skip("Nexus tests don't work with the Java test server")
207+
provider = create_tracer_provider()
208+
opentelemetry.trace.set_tracer_provider(provider)
209+
config = client.config()
210+
config["plugins"] = [OpenTelemetryPlugin()]
211+
client = Client(**config)
212+
213+
async with new_worker(
214+
client, SignalWithStartCallerWorkflow, SignalWithStartHeaderWorkflow
215+
) as worker:
216+
target_id = f"signal-with-start-target-{uuid.uuid4()}"
217+
with get_tracer(__name__).start_as_current_span("signal-with-start"):
218+
caller = await client.start_workflow(
219+
SignalWithStartCallerWorkflow.run,
220+
args=[target_id, worker.task_queue],
221+
id=f"signal-with-start-caller-{uuid.uuid4()}",
222+
task_queue=worker.task_queue,
223+
execution_timeout=timedelta(seconds=3),
224+
)
225+
assert await caller.result() == target_id
226+
assert await client.get_workflow_handle(target_id).result() is True
227+
228+
172229
@workflow.defn
173230
class ComprehensiveWorkflow:
174231
def __init__(self) -> None:

‎tests/nexus/test_temporal_system_nexus.py‎

Lines changed: 9 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -36,6 +36,7 @@
3636
from temporalio.worker import (
3737
Interceptor,
3838
StartNexusOperationInput,
39+
StartSystemNexusOperationInput,
3940
Worker,
4041
WorkflowInboundInterceptor,
4142
WorkflowInterceptorClassInput,
@@ -253,6 +254,12 @@ async def start_nexus_operation(
253254
interceptor_traces.append(("workflow.start_nexus_operation", input))
254255
return await super().start_nexus_operation(input)
255256

257+
async def start_system_nexus_operation(
258+
self, input: StartSystemNexusOperationInput[Any, Any]
259+
) -> workflow.NexusOperationHandle[Any]:
260+
interceptor_traces.append(("workflow.start_system_nexus_operation", input))
261+
return await super().start_system_nexus_operation(input)
262+
256263

257264
def _assert_stored_payloads_include(
258265
driver: InMemoryTestDriver, expected_payload_data: set[bytes]
@@ -269,8 +276,8 @@ def _assert_stored_payloads_include(
269276
def _assert_start_nexus_operation_interceptor_trace() -> None:
270277
assert len(interceptor_traces) == 1
271278
trace_name, trace_value = interceptor_traces.pop()
272-
assert trace_name == "workflow.start_nexus_operation"
273-
trace_input = cast(StartNexusOperationInput[Any, Any], trace_value)
279+
assert trace_name == "workflow.start_system_nexus_operation"
280+
trace_input = cast(StartSystemNexusOperationInput[Any, Any], trace_value)
274281
request = trace_input.input
275282
assert request.id == "system-nexus-workflow-id"
276283
assert request.signal == "test-signal"

0 commit comments

Comments
 (0)