From b61dad7e03230f12444b125daddf5e6d4caf1747 Mon Sep 17 00:00:00 2001 From: Tim Conley Date: Thu, 16 Jul 2026 11:53:50 -0700 Subject: [PATCH 1/3] Wrap payload converters for temporal intermediate models --- temporalio/activity.py | 12 +++- temporalio/converter/__init__.py | 2 + temporalio/converter/_data_converter.py | 7 +- temporalio/converter/_payload_converter.py | 77 ++++++++++++++++++++++ temporalio/worker/_workflow_instance.py | 6 +- tests/test_converter.py | 54 +++++++++++++++ 6 files changed, 154 insertions(+), 4 deletions(-) diff --git a/temporalio/activity.py b/temporalio/activity.py index 4e632701e..1c80438e9 100644 --- a/temporalio/activity.py +++ b/temporalio/activity.py @@ -238,9 +238,17 @@ def payload_converter(self) -> temporalio.converter.PayloadConverter: self.payload_converter_class_or_instance, temporalio.converter.PayloadConverter, ): - self._payload_converter = self.payload_converter_class_or_instance + self._payload_converter = ( + temporalio.converter.TemporalIntermediatePayloadConverter.wrap( + self.payload_converter_class_or_instance + ) + ) else: - self._payload_converter = self.payload_converter_class_or_instance() + self._payload_converter = ( + temporalio.converter.TemporalIntermediatePayloadConverter.wrap( + self.payload_converter_class_or_instance() + ) + ) return self._payload_converter @property diff --git a/temporalio/converter/__init__.py b/temporalio/converter/__init__.py index 3821cbd68..7cdec2a76 100644 --- a/temporalio/converter/__init__.py +++ b/temporalio/converter/__init__.py @@ -33,6 +33,7 @@ JSONTypeConverter, JSONTypeConverterUnhandled, PayloadConverter, + TemporalIntermediatePayloadConverter, value_to_type, ) from temporalio.converter._payload_limits import ( @@ -83,6 +84,7 @@ "PayloadLimitsConfig", "PayloadSizeWarning", "SerializationContext", + "TemporalIntermediatePayloadConverter", "WithSerializationContext", "WorkflowSerializationContext", "decode_search_attributes", diff --git a/temporalio/converter/_data_converter.py b/temporalio/converter/_data_converter.py index 13b48e695..b0c87b926 100644 --- a/temporalio/converter/_data_converter.py +++ b/temporalio/converter/_data_converter.py @@ -29,6 +29,7 @@ ) from temporalio.converter._payload_converter import ( PayloadConverter, + TemporalIntermediatePayloadConverter, ) from temporalio.converter._payload_limits import ( PayloadLimitsConfig, @@ -103,7 +104,11 @@ class DataConverter(WithSerializationContext): """Server-reported limits for payloads.""" def __post_init__(self) -> None: # noqa: D105 - object.__setattr__(self, "payload_converter", self.payload_converter_class()) + object.__setattr__( + self, + "payload_converter", + TemporalIntermediatePayloadConverter.wrap(self.payload_converter_class()), + ) object.__setattr__(self, "failure_converter", self.failure_converter_class()) async def encode( diff --git a/temporalio/converter/_payload_converter.py b/temporalio/converter/_payload_converter.py index 8ee85ef72..261ad1033 100644 --- a/temporalio/converter/_payload_converter.py +++ b/temporalio/converter/_payload_converter.py @@ -514,6 +514,83 @@ def from_payload( raise RuntimeError("Failed parsing") from err +class TemporalIntermediatePayloadConverter(PayloadConverter, WithSerializationContext): + """Payload converter wrapper for generated Temporal intermediate hooks. + + Values with a ``_temporal_to_intermediate`` method are first converted to + their intermediate value, then encoded by the wrapped payload converter. When + decoding to a type with ``_temporal_from_intermediate``, the wrapped + converter first decodes the payload to the intermediate value and this + wrapper constructs the requested user-facing type from it. + """ + + _inner_payload_converter: PayloadConverter + + def __init__(self, inner_payload_converter: PayloadConverter) -> None: + """Create a Temporal intermediate payload converter.""" + self._inner_payload_converter = inner_payload_converter + + @staticmethod + def wrap(payload_converter: PayloadConverter) -> PayloadConverter: + """Wrap a payload converter unless it is already wrapped.""" + if isinstance(payload_converter, TemporalIntermediatePayloadConverter): + return payload_converter + return TemporalIntermediatePayloadConverter(payload_converter) + + def to_payloads( + self, values: Sequence[Any] + ) -> list[temporalio.api.common.v1.Payload]: + """See base class.""" + intermediate_values: list[Any] = [] + for value in values: + to_intermediate = getattr(value, "_temporal_to_intermediate", None) + if to_intermediate is not None: + value = to_intermediate(payload_converter=self._inner_payload_converter) + intermediate_values.append(value) + return self._inner_payload_converter.to_payloads(intermediate_values) + + def from_payloads( + self, + payloads: Sequence[temporalio.api.common.v1.Payload], + type_hints: list[type] | None = None, + ) -> list[Any]: + """See base class.""" + if type_hints is None: + return self._inner_payload_converter.from_payloads(payloads, None) + normalized_type_hints: list[type | None] = list(type_hints) + if len(normalized_type_hints) < len(payloads): + normalized_type_hints.extend([None] * (len(payloads) - len(type_hints))) + inner_type_hints = [ + None + if getattr(type_hint, "_temporal_from_intermediate", None) is not None + else type_hint + for type_hint in normalized_type_hints + ] + values = self._inner_payload_converter.from_payloads( + payloads, typing.cast("list[type]", inner_type_hints) + ) + return [ + from_intermediate(value, payload_converter=self._inner_payload_converter) + if ( + from_intermediate := getattr( + type_hint, "_temporal_from_intermediate", None + ) + ) + is not None + else value + for value, type_hint in zip(values, normalized_type_hints) + ] + + def with_context(self, context: SerializationContext) -> Self: + """Return a new instance with context set on the inner converter.""" + if not isinstance(self._inner_payload_converter, WithSerializationContext): + return self + inner_payload_converter = self._inner_payload_converter.with_context(context) + if inner_payload_converter is self._inner_payload_converter: + return self + return type(self)(inner_payload_converter) + + class AdvancedJSONEncoder(json.JSONEncoder): """Advanced JSON encoder. diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index 74edc66b7..4d329059b 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -246,7 +246,11 @@ def __init__(self, det: WorkflowInstanceDetails) -> None: self._defn = det.defn self._workflow_input: ExecuteWorkflowInput | None = None self._info = det.info - self._context_free_payload_converter = det.payload_converter_class() + self._context_free_payload_converter = ( + temporalio.converter.TemporalIntermediatePayloadConverter.wrap( + det.payload_converter_class() + ) + ) self._context_free_failure_converter = det.failure_converter_class() workflow_context = temporalio.converter.WorkflowSerializationContext( namespace=det.info.namespace, diff --git a/tests/test_converter.py b/tests/test_converter.py index 10365f9c1..2281513f2 100644 --- a/tests/test_converter.py +++ b/tests/test_converter.py @@ -44,6 +44,8 @@ JSONTypeConverter, JSONTypeConverterUnhandled, PayloadCodec, + PayloadConverter, + TemporalIntermediatePayloadConverter, decode_search_attributes, encode_search_attribute_values, value_to_type, @@ -254,6 +256,58 @@ def test_binary_proto(): assert decoded == proto +@dataclass +class TemporalIntermediateModel: + value: str + + @classmethod + def _temporal_from_intermediate( + cls, + intermediate: temporalio.api.common.v1.WorkflowExecution, + *, + payload_converter: PayloadConverter | None = None, + ) -> TemporalIntermediateModel: + assert payload_converter is not None + return cls(value=intermediate.workflow_id) + + def _temporal_to_intermediate( + self, *, payload_converter: PayloadConverter | None = None + ) -> temporalio.api.common.v1.WorkflowExecution: + assert payload_converter is not None + return temporalio.api.common.v1.WorkflowExecution( + workflow_id=self.value, + run_id="run-id", + ) + + +class CustomDefaultPayloadConverter(DefaultPayloadConverter): + pass + + +def test_temporal_intermediate_payload_converter_wraps_user_converter(): + data_converter = DataConverter( + payload_converter_class=CustomDefaultPayloadConverter + ) + converter = data_converter.payload_converter + assert isinstance(converter, TemporalIntermediatePayloadConverter) + value = TemporalIntermediateModel("workflow-id") + + payload = converter.to_payload(value) + + assert payload.metadata["encoding"] == b"json/protobuf" + assert ( + payload.metadata["messageType"] == b"temporal.api.common.v1.WorkflowExecution" + ) + assert all("temporal-wire" not in key for key in payload.metadata) + assert all(b"temporal-wire" not in value for value in payload.metadata.values()) + assert converter.from_payload(payload, TemporalIntermediateModel) == value + + plain_proto_payload = converter.to_payload( + temporalio.api.common.v1.WorkflowExecution(workflow_id="id1", run_id="id2") + ) + assert plain_proto_payload.metadata["encoding"] == b"json/protobuf" + + def test_encode_search_attribute_values(): with pytest.raises(TypeError, match="of type tuple not one of"): encode_search_attribute_values([("bad type",)]) # type: ignore[arg-type] From 585bc27d62e9dc5f299685289ef239c32bcb215a Mon Sep 17 00:00:00 2001 From: Tim Conley Date: Thu, 16 Jul 2026 13:04:40 -0700 Subject: [PATCH 2/3] Support intermediate hooks in system Nexus conversion --- temporalio/converter/_payload_converter.py | 4 +- temporalio/nexus/system/__init__.py | 74 ++++++++++++++++++++-- temporalio/worker/_workflow_instance.py | 4 +- tests/nexus/test_temporal_system_nexus.py | 12 +++- tests/test_converter.py | 8 +-- tests/worker/test_visitor.py | 4 +- 6 files changed, 87 insertions(+), 19 deletions(-) diff --git a/temporalio/converter/_payload_converter.py b/temporalio/converter/_payload_converter.py index 261ad1033..8cdcc7abe 100644 --- a/temporalio/converter/_payload_converter.py +++ b/temporalio/converter/_payload_converter.py @@ -545,7 +545,7 @@ def to_payloads( for value in values: to_intermediate = getattr(value, "_temporal_to_intermediate", None) if to_intermediate is not None: - value = to_intermediate(payload_converter=self._inner_payload_converter) + value = to_intermediate() intermediate_values.append(value) return self._inner_payload_converter.to_payloads(intermediate_values) @@ -570,7 +570,7 @@ def from_payloads( payloads, typing.cast("list[type]", inner_type_hints) ) return [ - from_intermediate(value, payload_converter=self._inner_payload_converter) + from_intermediate(value) if ( from_intermediate := getattr( type_hint, "_temporal_from_intermediate", None diff --git a/temporalio/nexus/system/__init__.py b/temporalio/nexus/system/__init__.py index 21c5a1408..e76756cd7 100644 --- a/temporalio/nexus/system/__init__.py +++ b/temporalio/nexus/system/__init__.py @@ -2,22 +2,82 @@ from __future__ import annotations +import contextlib +import contextvars +from collections.abc import Iterator, Sequence +from typing import Any + import temporalio.api.common.v1 import temporalio.converter from temporalio.bridge._visitor_functions import VisitorFunctions from temporalio.converter import BinaryProtoPayloadConverter, CompositePayloadConverter TEMPORAL_SYSTEM_ENDPOINT = "__temporal_system" +_user_payload_converter: contextvars.ContextVar[ + temporalio.converter.PayloadConverter | None +] = contextvars.ContextVar("temporal-system-nexus-user-payload-converter", default=None) -class SystemNexusPayloadConverter(CompositePayloadConverter): - """Payload converter for system Nexus outer envelopes.""" +@contextlib.contextmanager +def user_payload_converter_context( + payload_converter: temporalio.converter.PayloadConverter, +) -> Iterator[None]: + """Set the user payload converter for system Nexus model conversion.""" + token = _user_payload_converter.set(payload_converter) + try: + yield + finally: + _user_payload_converter.reset(token) + + +def current_user_payload_converter() -> temporalio.converter.PayloadConverter: + """Return the active user payload converter for system Nexus model conversion.""" + payload_converter = _user_payload_converter.get() + if payload_converter is None: + raise RuntimeError("System Nexus user payload converter context is not active") + return payload_converter + + +class _SystemNexusOuterPayloadConverter(CompositePayloadConverter): + """Payload converter for system Nexus outer proto envelopes.""" def __init__(self) -> None: """Create a payload converter for system Nexus outer envelopes.""" super().__init__(BinaryProtoPayloadConverter()) +class SystemNexusPayloadConverter(temporalio.converter.PayloadConverter): + """Payload converter for system Nexus outer envelopes.""" + + _user_payload_converter: temporalio.converter.PayloadConverter + _outer_payload_converter: temporalio.converter.PayloadConverter + + def __init__(self, user_payload_converter: temporalio.converter.PayloadConverter) -> None: + """Create a payload converter for system Nexus outer envelopes.""" + self._user_payload_converter = user_payload_converter + self._outer_payload_converter = ( + temporalio.converter.TemporalIntermediatePayloadConverter.wrap( + _SystemNexusOuterPayloadConverter() + ) + ) + + def to_payloads( + self, values: Sequence[Any] + ) -> list[temporalio.api.common.v1.Payload]: + """See base class.""" + with user_payload_converter_context(self._user_payload_converter): + return self._outer_payload_converter.to_payloads(values) + + def from_payloads( + self, + payloads: Sequence[temporalio.api.common.v1.Payload], + type_hints: list[type] | None = None, + ) -> list[Any]: + """See base class.""" + with user_payload_converter_context(self._user_payload_converter): + return self._outer_payload_converter.from_payloads(payloads, type_hints) + + def is_system_endpoint(endpoint: str) -> bool: """Return whether a Nexus endpoint is the Temporal system endpoint.""" return endpoint == TEMPORAL_SYSTEM_ENDPOINT @@ -33,7 +93,7 @@ async def maybe_visit_payload( if not is_system_endpoint(endpoint): return None - payload_converter = get_payload_converter() + payload_converter = _SystemNexusOuterPayloadConverter() value = payload_converter.from_payload(payload) from ._payload_visitor import PayloadVisitor @@ -43,15 +103,19 @@ async def maybe_visit_payload( return payload_converter.to_payload(value) -def get_payload_converter() -> temporalio.converter.PayloadConverter: +def get_payload_converter( + user_payload_converter: temporalio.converter.PayloadConverter, +) -> temporalio.converter.PayloadConverter: """Return the fixed payload converter for system Nexus outer envelopes.""" - return SystemNexusPayloadConverter() + return SystemNexusPayloadConverter(user_payload_converter) __all__ = [ "TEMPORAL_SYSTEM_ENDPOINT", + "current_user_payload_converter", "get_payload_converter", "is_system_endpoint", "maybe_visit_payload", "SystemNexusPayloadConverter", + "user_payload_converter_context", ] diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index 4d329059b..dd7f4571a 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -2093,7 +2093,9 @@ async def operation_handle_fn() -> OutputT: t.uncancel() # type: ignore[union-attr] payload_converter = ( - temporalio.nexus.system.get_payload_converter() + temporalio.nexus.system.get_payload_converter( + self._workflow_context_payload_converter + ) if temporalio.nexus.system.is_system_endpoint(input.endpoint) else self._context_free_payload_converter ) diff --git a/tests/nexus/test_temporal_system_nexus.py b/tests/nexus/test_temporal_system_nexus.py index b689ee8d9..4629d230b 100644 --- a/tests/nexus/test_temporal_system_nexus.py +++ b/tests/nexus/test_temporal_system_nexus.py @@ -186,7 +186,9 @@ def _new_system_nexus_request_payload() -> temporalio.api.common.v1.Payload: assert nested_payload is not None request = workflowservice_pb2.SignalWithStartWorkflowExecutionRequest() request.input.payloads.add().CopyFrom(nested_payload) - payload = nexus_system.get_payload_converter().to_payload(request) + payload = nexus_system.get_payload_converter( + temporalio.converter.PayloadConverter.default + ).to_payload(request) assert payload is not None return payload @@ -201,7 +203,9 @@ async def test_schedule_system_nexus_endpoint_ignores_operation_registry() -> No await PayloadVisitor().visit(visitor, completion) schedule = completion.successful.commands[0].schedule_nexus_operation - decoded = nexus_system.get_payload_converter().from_payload(schedule.input) + decoded = nexus_system.get_payload_converter( + temporalio.converter.PayloadConverter.default + ).from_payload(schedule.input) assert isinstance( decoded, workflowservice_pb2.SignalWithStartWorkflowExecutionRequest ) @@ -339,7 +343,9 @@ def _field_is_repeated(field: FieldDescriptor) -> bool: ], ) def test_system_nexus_proto_roundtrip(message_type: type[Message]) -> None: - payload_converter = nexus_system.get_payload_converter() + payload_converter = nexus_system.get_payload_converter( + temporalio.converter.PayloadConverter.default + ) proto_value = _build_proto_sample(message_type) payload = payload_converter.to_payload(proto_value) assert payload is not None diff --git a/tests/test_converter.py b/tests/test_converter.py index 2281513f2..0ae2a2464 100644 --- a/tests/test_converter.py +++ b/tests/test_converter.py @@ -264,16 +264,10 @@ class TemporalIntermediateModel: def _temporal_from_intermediate( cls, intermediate: temporalio.api.common.v1.WorkflowExecution, - *, - payload_converter: PayloadConverter | None = None, ) -> TemporalIntermediateModel: - assert payload_converter is not None return cls(value=intermediate.workflow_id) - def _temporal_to_intermediate( - self, *, payload_converter: PayloadConverter | None = None - ) -> temporalio.api.common.v1.WorkflowExecution: - assert payload_converter is not None + def _temporal_to_intermediate(self) -> temporalio.api.common.v1.WorkflowExecution: return temporalio.api.common.v1.WorkflowExecution( workflow_id=self.value, run_id="run-id", diff --git a/tests/worker/test_visitor.py b/tests/worker/test_visitor.py index bd4004625..aa9a931e1 100644 --- a/tests/worker/test_visitor.py +++ b/tests/worker/test_visitor.py @@ -357,7 +357,9 @@ async def _visit(self) -> None: finally: active_visits -= 1 - payload_converter = nexus_system.get_payload_converter() + payload_converter = nexus_system.get_payload_converter( + temporalio.converter.PayloadConverter.default + ) system_request = workflowservice_pb2.SignalWithStartWorkflowExecutionRequest( input=Payloads(payloads=[Payload(data=b"workflow-input")]), signal_input=Payloads(payloads=[Payload(data=b"signal-input")]), From 3d36e32d4365eb5d647c519bba6cc2443fdb2ccd Mon Sep 17 00:00:00 2001 From: Tim Conley Date: Thu, 16 Jul 2026 13:06:50 -0700 Subject: [PATCH 3/3] Regenerate system Nexus APIs for intermediate models --- scripts/nex_gen_support.py | 75 +++++-- .../nexus/system/workflow_service/__init__.py | 4 +- .../workflow_service/_resources/__init__.py | 5 - .../_support/nex_gen_support.py | 65 +++++- .../nexus/system/workflow_service/models.py | 133 ++++++++++- .../operations/signal_with_start_workflow.py | 207 +++++++++--------- .../{service.py => services.py} | 9 +- tests/nexus/test_temporal_system_nexus.py | 9 +- 8 files changed, 360 insertions(+), 147 deletions(-) delete mode 100644 temporalio/nexus/system/workflow_service/_resources/__init__.py rename temporalio/nexus/system/workflow_service/{service.py => services.py} (63%) diff --git a/scripts/nex_gen_support.py b/scripts/nex_gen_support.py index 58b51e263..687ec692b 100644 --- a/scripts/nex_gen_support.py +++ b/scripts/nex_gen_support.py @@ -1,15 +1,20 @@ import collections.abc -import typing from datetime import timedelta +import typing import google.protobuf.duration_pb2 - import temporalio.api.common.v1.message_pb2 as common_pb2 import temporalio.api.enums.v1.workflow_pb2 as workflow_enums_pb2 import temporalio.api.taskqueue.v1.message_pb2 as taskqueue_pb2 import temporalio.api.workflow.v1 -import temporalio.common import temporalio.converter +import temporalio.common +import temporalio.nexus.system + + +def _current_payload_converter( +) -> temporalio.converter.PayloadConverter: + return temporalio.nexus.system.current_user_payload_converter() def retry_policy_from_proto( @@ -38,9 +43,7 @@ def workflow_function_name( def signal_function_to_proto( value: str | collections.abc.Callable[..., typing.Any], ) -> str: - from temporalio.workflow import ( - _SignalDefinition, # pyright: ignore[reportPrivateUsage] - ) + from temporalio.workflow import _SignalDefinition # pyright: ignore[reportPrivateUsage] return _SignalDefinition.must_name_from_fn_or_str(value) # pyright: ignore[reportUnknownMemberType] @@ -52,6 +55,12 @@ def workflow_type_to_proto( return common_pb2.WorkflowType(name=workflow_function_name(workflow_type)) +def workflow_type_from_proto( + proto: common_pb2.WorkflowType, +) -> str: + return proto.name + + def task_queue_from_proto( proto: taskqueue_pb2.TaskQueue, ) -> str: @@ -73,9 +82,13 @@ def workflow_namespace() -> str: def payloads_to_proto( values: collections.abc.Sequence[typing.Any], ) -> common_pb2.Payloads: - from temporalio.workflow import payload_converter + return _current_payload_converter().to_payloads_wrapper(values) + - return payload_converter().to_payloads_wrapper(values) +def payloads_from_proto( + proto: common_pb2.Payloads, +) -> list[object]: + return list(_current_payload_converter().from_payloads_wrapper(proto)) def _clone_payload(payload: common_pb2.Payload) -> common_pb2.Payload: @@ -84,23 +97,25 @@ def _clone_payload(payload: common_pb2.Payload) -> common_pb2.Payload: return clone -def _value_to_payload(value: object | common_pb2.Payload) -> common_pb2.Payload: +def _value_to_payload( + value: object | common_pb2.Payload, +) -> common_pb2.Payload: if isinstance(value, common_pb2.Payload): return _clone_payload(value) - from temporalio.workflow import payload_converter - payloads = payload_converter().to_payloads_wrapper([value]) + payloads = _current_payload_converter().to_payloads_wrapper([value]) return _clone_payload(payloads.payloads[0]) -def _payload_to_value(payload: common_pb2.Payload) -> object: +def _payload_to_value( + payload: common_pb2.Payload, +) -> object: wrapper = common_pb2.Payloads() wrapper.payloads.add().CopyFrom(payload) - from temporalio.workflow import payload_converter return typing.cast( object, - payload_converter().from_payloads_wrapper(wrapper)[0], + _current_payload_converter().from_payloads_wrapper(wrapper)[0], ) @@ -131,7 +146,9 @@ def memo_to_proto( return message -def duration_from_proto(proto: google.protobuf.duration_pb2.Duration) -> timedelta: +def duration_from_proto( + proto: google.protobuf.duration_pb2.Duration, +) -> timedelta: return proto.ToTimedelta() @@ -177,6 +194,12 @@ def search_attributes_to_proto( return proto +def search_attributes_from_proto( + proto: common_pb2.SearchAttributes, +) -> temporalio.common.TypedSearchAttributes: + return temporalio.converter.decode_typed_search_attributes(proto) + + def priority_from_proto( proto: common_pb2.Priority, ) -> temporalio.common.Priority: @@ -193,3 +216,25 @@ def versioning_override_to_proto( versioning_override: temporalio.common.VersioningOverride, ) -> temporalio.api.workflow.v1.VersioningOverride: return versioning_override._to_proto() # pyright: ignore[reportPrivateUsage] + + +def versioning_override_from_proto( + proto: temporalio.api.workflow.v1.VersioningOverride, +) -> temporalio.common.VersioningOverride: + if proto.HasField("pinned") and proto.pinned.HasField("version"): + version = proto.pinned.version + return temporalio.common.PinnedVersioningOverride( + temporalio.common.WorkerDeploymentVersion( + deployment_name=version.deployment_name, + build_id=version.build_id, + ) + ) + if proto.pinned_version: + return temporalio.common.PinnedVersioningOverride( + temporalio.common.WorkerDeploymentVersion.from_canonical_string( + proto.pinned_version + ) + ) + if proto.auto_upgrade: + return temporalio.common.AutoUpgradeVersioningOverride() + raise ValueError("unknown versioning override proto shape") diff --git a/temporalio/nexus/system/workflow_service/__init__.py b/temporalio/nexus/system/workflow_service/__init__.py index 7c24fa125..4a92138ac 100644 --- a/temporalio/nexus/system/workflow_service/__init__.py +++ b/temporalio/nexus/system/workflow_service/__init__.py @@ -2,7 +2,7 @@ from __future__ import annotations -from . import service as _service +from . import services as _services from .operations.signal_with_start_workflow import signal_with_start_workflow __all__ = [ @@ -14,5 +14,5 @@ ( "temporal.api.workflowservice.v1.WorkflowService", "SignalWithStartWorkflowExecution", - ): _service.WorkflowService.signal_with_start_workflow, + ): _services.WorkflowService.signal_with_start_workflow, } diff --git a/temporalio/nexus/system/workflow_service/_resources/__init__.py b/temporalio/nexus/system/workflow_service/_resources/__init__.py deleted file mode 100644 index 373efbd33..000000000 --- a/temporalio/nexus/system/workflow_service/_resources/__init__.py +++ /dev/null @@ -1,5 +0,0 @@ -# Generated by nex-gen. DO NOT EDIT! - -from __future__ import annotations - -__all__ = [] diff --git a/temporalio/nexus/system/workflow_service/_support/nex_gen_support.py b/temporalio/nexus/system/workflow_service/_support/nex_gen_support.py index 58b51e263..331c36a60 100644 --- a/temporalio/nexus/system/workflow_service/_support/nex_gen_support.py +++ b/temporalio/nexus/system/workflow_service/_support/nex_gen_support.py @@ -10,6 +10,11 @@ import temporalio.api.workflow.v1 import temporalio.common import temporalio.converter +import temporalio.nexus.system + + +def _current_payload_converter() -> temporalio.converter.PayloadConverter: + return temporalio.nexus.system.current_user_payload_converter() def retry_policy_from_proto( @@ -52,6 +57,12 @@ def workflow_type_to_proto( return common_pb2.WorkflowType(name=workflow_function_name(workflow_type)) +def workflow_type_from_proto( + proto: common_pb2.WorkflowType, +) -> str: + return proto.name + + def task_queue_from_proto( proto: taskqueue_pb2.TaskQueue, ) -> str: @@ -73,9 +84,13 @@ def workflow_namespace() -> str: def payloads_to_proto( values: collections.abc.Sequence[typing.Any], ) -> common_pb2.Payloads: - from temporalio.workflow import payload_converter + return _current_payload_converter().to_payloads_wrapper(values) + - return payload_converter().to_payloads_wrapper(values) +def payloads_from_proto( + proto: common_pb2.Payloads, +) -> list[object]: + return list(_current_payload_converter().from_payloads_wrapper(proto)) def _clone_payload(payload: common_pb2.Payload) -> common_pb2.Payload: @@ -84,23 +99,25 @@ def _clone_payload(payload: common_pb2.Payload) -> common_pb2.Payload: return clone -def _value_to_payload(value: object | common_pb2.Payload) -> common_pb2.Payload: +def _value_to_payload( + value: object | common_pb2.Payload, +) -> common_pb2.Payload: if isinstance(value, common_pb2.Payload): return _clone_payload(value) - from temporalio.workflow import payload_converter - payloads = payload_converter().to_payloads_wrapper([value]) + payloads = _current_payload_converter().to_payloads_wrapper([value]) return _clone_payload(payloads.payloads[0]) -def _payload_to_value(payload: common_pb2.Payload) -> object: +def _payload_to_value( + payload: common_pb2.Payload, +) -> object: wrapper = common_pb2.Payloads() wrapper.payloads.add().CopyFrom(payload) - from temporalio.workflow import payload_converter return typing.cast( object, - payload_converter().from_payloads_wrapper(wrapper)[0], + _current_payload_converter().from_payloads_wrapper(wrapper)[0], ) @@ -131,7 +148,9 @@ def memo_to_proto( return message -def duration_from_proto(proto: google.protobuf.duration_pb2.Duration) -> timedelta: +def duration_from_proto( + proto: google.protobuf.duration_pb2.Duration, +) -> timedelta: return proto.ToTimedelta() @@ -177,6 +196,12 @@ def search_attributes_to_proto( return proto +def search_attributes_from_proto( + proto: common_pb2.SearchAttributes, +) -> temporalio.common.TypedSearchAttributes: + return temporalio.converter.decode_typed_search_attributes(proto) + + def priority_from_proto( proto: common_pb2.Priority, ) -> temporalio.common.Priority: @@ -193,3 +218,25 @@ def versioning_override_to_proto( versioning_override: temporalio.common.VersioningOverride, ) -> temporalio.api.workflow.v1.VersioningOverride: return versioning_override._to_proto() # pyright: ignore[reportPrivateUsage] + + +def versioning_override_from_proto( + proto: temporalio.api.workflow.v1.VersioningOverride, +) -> temporalio.common.VersioningOverride: + if proto.HasField("pinned") and proto.pinned.HasField("version"): + version = proto.pinned.version + return temporalio.common.PinnedVersioningOverride( + temporalio.common.WorkerDeploymentVersion( + deployment_name=version.deployment_name, + build_id=version.build_id, + ) + ) + if proto.pinned_version: + return temporalio.common.PinnedVersioningOverride( + temporalio.common.WorkerDeploymentVersion.from_canonical_string( + proto.pinned_version + ) + ) + if proto.auto_upgrade: + return temporalio.common.AutoUpgradeVersioningOverride() + raise ValueError("unknown versioning override proto shape") diff --git a/temporalio/nexus/system/workflow_service/models.py b/temporalio/nexus/system/workflow_service/models.py index 05e1e3088..4a0dc74e2 100644 --- a/temporalio/nexus/system/workflow_service/models.py +++ b/temporalio/nexus/system/workflow_service/models.py @@ -2,6 +2,7 @@ from __future__ import annotations +# pyright: reportPrivateUsage=false import collections.abc import dataclasses import datetime @@ -12,20 +13,31 @@ import temporalio.common from ._support import ( + duration_from_proto, duration_to_proto, + memo_from_proto, memo_to_proto, payload_from_proto, payload_to_proto, + payloads_from_proto, payloads_to_proto, + priority_from_proto, priority_to_proto, + retry_policy_from_proto, retry_policy_to_proto, + search_attributes_from_proto, search_attributes_to_proto, signal_function_to_proto, + task_queue_from_proto, task_queue_to_proto, + versioning_override_from_proto, versioning_override_to_proto, + workflow_id_conflict_policy_from_proto, workflow_id_conflict_policy_to_proto, + workflow_id_reuse_policy_from_proto, workflow_id_reuse_policy_to_proto, workflow_namespace, + workflow_type_from_proto, workflow_type_to_proto, ) @@ -59,8 +71,84 @@ class SignalWithStartWorkflowRequest: versioning_override: temporalio.common.VersioningOverride | None = None start_delay: datetime.timedelta | None = None user_metadata: UserMetadata | None = None + namespace: str = dataclasses.field(default_factory=workflow_namespace) - def to_proto( + @classmethod + def _temporal_from_intermediate( + cls, + proto: temporalio.api.workflowservice.v1.request_response_pb2.SignalWithStartWorkflowExecutionRequest, + ) -> SignalWithStartWorkflowRequest: + if not proto.HasField("workflow_type"): + raise ValueError( + "missing required field SignalWithStartWorkflowRequest.workflow" + ) + workflow = workflow_type_from_proto(proto.workflow_type) + if not proto.workflow_id: + raise ValueError("missing required field SignalWithStartWorkflowRequest.id") + id = proto.workflow_id + if not proto.HasField("task_queue"): + raise ValueError( + "missing required field SignalWithStartWorkflowRequest.task_queue" + ) + task_queue = task_queue_from_proto(proto.task_queue) + if not proto.signal_name: + raise ValueError( + "missing required field SignalWithStartWorkflowRequest.signal" + ) + signal = proto.signal_name + return cls( + workflow=workflow, + args=payloads_from_proto(proto.input) if proto.HasField("input") else None, + id=id, + task_queue=task_queue, + signal=signal, + signal_args=payloads_from_proto(proto.signal_input) + if proto.HasField("signal_input") + else None, + execution_timeout=duration_from_proto(proto.workflow_execution_timeout) + if proto.HasField("workflow_execution_timeout") + else None, + run_timeout=duration_from_proto(proto.workflow_run_timeout) + if proto.HasField("workflow_run_timeout") + else None, + task_timeout=duration_from_proto(proto.workflow_task_timeout) + if proto.HasField("workflow_task_timeout") + else None, + request_id=proto.request_id if bool(proto.request_id) else None, + id_reuse_policy=workflow_id_reuse_policy_from_proto( + proto.workflow_id_reuse_policy + ), + id_conflict_policy=workflow_id_conflict_policy_from_proto( + proto.workflow_id_conflict_policy + ) + if proto.workflow_id_conflict_policy != 0 + else None, + retry_policy=retry_policy_from_proto(proto.retry_policy) + if proto.HasField("retry_policy") + else None, + cron_schedule=proto.cron_schedule if bool(proto.cron_schedule) else None, + memo=memo_from_proto(proto.memo) if proto.HasField("memo") else None, + search_attributes=search_attributes_from_proto(proto.search_attributes) + if proto.HasField("search_attributes") + else None, + priority=priority_from_proto(proto.priority) + if proto.HasField("priority") + else None, + versioning_override=versioning_override_from_proto( + proto.versioning_override + ) + if proto.HasField("versioning_override") + else None, + start_delay=duration_from_proto(proto.workflow_start_delay) + if proto.HasField("workflow_start_delay") + else None, + user_metadata=UserMetadata._temporal_from_intermediate(proto.user_metadata) + if proto.HasField("user_metadata") + else None, + namespace=proto.namespace, + ) + + def _temporal_to_intermediate( self, ) -> temporalio.api.workflowservice.v1.request_response_pb2.SignalWithStartWorkflowExecutionRequest: message = temporalio.api.workflowservice.v1.request_response_pb2.SignalWithStartWorkflowExecutionRequest() @@ -108,8 +196,10 @@ def to_proto( if self.start_delay is not None: message.workflow_start_delay.CopyFrom(duration_to_proto(self.start_delay)) if self.user_metadata is not None: - message.user_metadata.CopyFrom(self.user_metadata.to_proto()) - message.namespace = workflow_namespace() + message.user_metadata.CopyFrom( + self.user_metadata._temporal_to_intermediate() + ) + message.namespace = self.namespace return message @@ -119,7 +209,7 @@ class UserMetadata: static_details: typing.Any | None = None @classmethod - def from_proto( + def _temporal_from_intermediate( cls, proto: temporalio.api.sdk.v1.user_metadata_pb2.UserMetadata, ) -> UserMetadata: @@ -132,10 +222,43 @@ def from_proto( else None, ) - def to_proto(self) -> temporalio.api.sdk.v1.user_metadata_pb2.UserMetadata: + def _temporal_to_intermediate( + self, + ) -> temporalio.api.sdk.v1.user_metadata_pb2.UserMetadata: message = temporalio.api.sdk.v1.user_metadata_pb2.UserMetadata() if self.static_summary is not None: message.summary.CopyFrom(payload_to_proto(self.static_summary)) if self.static_details is not None: message.details.CopyFrom(payload_to_proto(self.static_details)) return message + + +@dataclasses.dataclass(slots=True) +class SignalWithStartWorkflowResponse: + """ + .. warning:: + This API is experimental and subject to change. + """ + + run_id: str | None = None + started: bool | None = None + + @classmethod + def _temporal_from_intermediate( + cls, + proto: temporalio.api.workflowservice.v1.request_response_pb2.SignalWithStartWorkflowExecutionResponse, + ) -> SignalWithStartWorkflowResponse: + return cls( + run_id=proto.run_id if bool(proto.run_id) else None, + started=proto.started if bool(proto.started) else None, + ) + + def _temporal_to_intermediate( + self, + ) -> temporalio.api.workflowservice.v1.request_response_pb2.SignalWithStartWorkflowExecutionResponse: + message = temporalio.api.workflowservice.v1.request_response_pb2.SignalWithStartWorkflowExecutionResponse() + if self.run_id is not None: + message.run_id = self.run_id + if self.started is not None: + message.started = self.started + return message diff --git a/temporalio/nexus/system/workflow_service/operations/signal_with_start_workflow.py b/temporalio/nexus/system/workflow_service/operations/signal_with_start_workflow.py index 0865e5a88..64aa88c62 100644 --- a/temporalio/nexus/system/workflow_service/operations/signal_with_start_workflow.py +++ b/temporalio/nexus/system/workflow_service/operations/signal_with_start_workflow.py @@ -8,7 +8,6 @@ import typing_extensions -import temporalio.api.workflowservice.v1.request_response_pb2 import temporalio.common if typing.TYPE_CHECKING: @@ -16,6 +15,7 @@ from ..models import ( SignalWithStartWorkflowRequest, + SignalWithStartWorkflowResponse, UserMetadata, ) @@ -33,15 +33,14 @@ async def _signal_with_start_workflow( get_external_workflow_handle, ) - request_proto = request.to_proto() nexus_client = create_nexus_client( service="temporal.api.workflowservice.v1.WorkflowService", endpoint="__temporal_system", ) handle = await nexus_client.start_operation( operation="SignalWithStartWorkflowExecution", - input=request_proto, - output_type=temporalio.api.workflowservice.v1.request_response_pb2.SignalWithStartWorkflowExecutionResponse, + input=request, + output_type=SignalWithStartWorkflowResponse, ) result = await handle return get_external_workflow_handle(request.id, run_id=result.run_id) @@ -77,17 +76,17 @@ async def signal_with_start_workflow( # Overload case: -# - workflow name with optional list-form workflow arguments -# - signal name with optional list-form signal arguments +# - workflow name with positional workflow arguments +# - signal method callable with no signal arguments @typing.overload async def signal_with_start_workflow( workflow: str, - *, - args: list[typing.Any] | None = ..., + *args: object, id: str, task_queue: str, - signal: str, - signal_args: list[typing.Any] | None = ..., + signal: collections.abc.Callable[ + [SelfType], None | collections.abc.Awaitable[None] + ], execution_timeout: datetime.timedelta | None = ..., run_timeout: datetime.timedelta | None = ..., task_timeout: datetime.timedelta | None = ..., @@ -103,23 +102,22 @@ async def signal_with_start_workflow( start_delay: datetime.timedelta | None = ..., static_summary: str | None = ..., static_details: str | None = ..., -) -> ExternalWorkflowHandle[object]: ... +) -> ExternalWorkflowHandle[SelfType]: ... # Overload case: -# - workflow method callable with typed positional workflow arguments -# - signal name with optional list-form signal arguments +# - workflow name with positional workflow arguments +# - signal method callable with a typed single signal arguments @typing.overload async def signal_with_start_workflow( - workflow: collections.abc.Callable[ - [SelfType, typing_extensions.Unpack[WorkflowArgs]], - collections.abc.Awaitable[WorkflowResult], - ], - *args: typing_extensions.Unpack[WorkflowArgs], + workflow: str, + *args: object, id: str, task_queue: str, - signal: str, - signal_args: list[typing.Any] | None = ..., + signal: collections.abc.Callable[ + [SelfType, SignalArg], None | collections.abc.Awaitable[None] + ], + signal_args: SignalArg, execution_timeout: datetime.timedelta | None = ..., run_timeout: datetime.timedelta | None = ..., task_timeout: datetime.timedelta | None = ..., @@ -139,20 +137,16 @@ async def signal_with_start_workflow( # Overload case: -# - workflow method callable with list-form workflow arguments -# - signal name with optional list-form signal arguments +# - workflow name with positional workflow arguments +# - signal callable with list-form signal arguments @typing.overload async def signal_with_start_workflow( - workflow: collections.abc.Callable[ - [SelfType, typing_extensions.Unpack[WorkflowArgs]], - collections.abc.Awaitable[WorkflowResult], - ], - *, - args: list[typing.Any], + workflow: str, + *args: object, id: str, task_queue: str, - signal: str, - signal_args: list[typing.Any] | None = ..., + signal: collections.abc.Callable[..., None | collections.abc.Awaitable[None]], + signal_args: list[typing.Any], execution_timeout: datetime.timedelta | None = ..., run_timeout: datetime.timedelta | None = ..., task_timeout: datetime.timedelta | None = ..., @@ -168,21 +162,21 @@ async def signal_with_start_workflow( start_delay: datetime.timedelta | None = ..., static_summary: str | None = ..., static_details: str | None = ..., -) -> ExternalWorkflowHandle[SelfType]: ... +) -> ExternalWorkflowHandle[object]: ... # Overload case: -# - workflow name with positional workflow arguments -# - signal method callable with no signal arguments +# - workflow name with optional list-form workflow arguments +# - signal name with optional list-form signal arguments @typing.overload async def signal_with_start_workflow( workflow: str, - *args: object, + *, + args: list[typing.Any] | None = ..., id: str, task_queue: str, - signal: collections.abc.Callable[ - [SelfType], None | collections.abc.Awaitable[None] - ], + signal: str, + signal_args: list[typing.Any] | None = ..., execution_timeout: datetime.timedelta | None = ..., run_timeout: datetime.timedelta | None = ..., task_timeout: datetime.timedelta | None = ..., @@ -198,7 +192,7 @@ async def signal_with_start_workflow( start_delay: datetime.timedelta | None = ..., static_summary: str | None = ..., static_details: str | None = ..., -) -> ExternalWorkflowHandle[SelfType]: ... +) -> ExternalWorkflowHandle[object]: ... # Overload case: @@ -233,20 +227,19 @@ async def signal_with_start_workflow( # Overload case: -# - workflow method callable with typed positional workflow arguments -# - signal method callable with no signal arguments +# - workflow name with optional list-form workflow arguments +# - signal method callable with a typed single signal arguments @typing.overload async def signal_with_start_workflow( - workflow: collections.abc.Callable[ - [SelfType, typing_extensions.Unpack[WorkflowArgs]], - collections.abc.Awaitable[WorkflowResult], - ], - *args: typing_extensions.Unpack[WorkflowArgs], + workflow: str, + *, + args: list[typing.Any] | None = ..., id: str, task_queue: str, signal: collections.abc.Callable[ - [SelfType], None | collections.abc.Awaitable[None] + [SelfType, SignalArg], None | collections.abc.Awaitable[None] ], + signal_args: SignalArg, execution_timeout: datetime.timedelta | None = ..., run_timeout: datetime.timedelta | None = ..., task_timeout: datetime.timedelta | None = ..., @@ -266,21 +259,17 @@ async def signal_with_start_workflow( # Overload case: -# - workflow method callable with list-form workflow arguments -# - signal method callable with no signal arguments +# - workflow name with optional list-form workflow arguments +# - signal callable with list-form signal arguments @typing.overload async def signal_with_start_workflow( - workflow: collections.abc.Callable[ - [SelfType, typing_extensions.Unpack[WorkflowArgs]], - collections.abc.Awaitable[WorkflowResult], - ], + workflow: str, *, - args: list[typing.Any], + args: list[typing.Any] | None = ..., id: str, task_queue: str, - signal: collections.abc.Callable[ - [SelfType], None | collections.abc.Awaitable[None] - ], + signal: collections.abc.Callable[..., None | collections.abc.Awaitable[None]], + signal_args: list[typing.Any], execution_timeout: datetime.timedelta | None = ..., run_timeout: datetime.timedelta | None = ..., task_timeout: datetime.timedelta | None = ..., @@ -296,22 +285,23 @@ async def signal_with_start_workflow( start_delay: datetime.timedelta | None = ..., static_summary: str | None = ..., static_details: str | None = ..., -) -> ExternalWorkflowHandle[SelfType]: ... +) -> ExternalWorkflowHandle[object]: ... # Overload case: -# - workflow name with positional workflow arguments -# - signal method callable with a typed single signal arguments +# - workflow method callable with typed positional workflow arguments +# - signal name with optional list-form signal arguments @typing.overload async def signal_with_start_workflow( - workflow: str, - *args: object, + workflow: collections.abc.Callable[ + [SelfType, typing_extensions.Unpack[WorkflowArgs]], + collections.abc.Awaitable[WorkflowResult], + ], + *args: typing_extensions.Unpack[WorkflowArgs], id: str, task_queue: str, - signal: collections.abc.Callable[ - [SelfType, SignalArg], None | collections.abc.Awaitable[None] - ], - signal_args: SignalArg, + signal: str, + signal_args: list[typing.Any] | None = ..., execution_timeout: datetime.timedelta | None = ..., run_timeout: datetime.timedelta | None = ..., task_timeout: datetime.timedelta | None = ..., @@ -331,19 +321,20 @@ async def signal_with_start_workflow( # Overload case: -# - workflow name with optional list-form workflow arguments -# - signal method callable with a typed single signal arguments +# - workflow method callable with typed positional workflow arguments +# - signal method callable with no signal arguments @typing.overload async def signal_with_start_workflow( - workflow: str, - *, - args: list[typing.Any] | None = ..., + workflow: collections.abc.Callable[ + [SelfType, typing_extensions.Unpack[WorkflowArgs]], + collections.abc.Awaitable[WorkflowResult], + ], + *args: typing_extensions.Unpack[WorkflowArgs], id: str, task_queue: str, signal: collections.abc.Callable[ - [SelfType, SignalArg], None | collections.abc.Awaitable[None] + [SelfType], None | collections.abc.Awaitable[None] ], - signal_args: SignalArg, execution_timeout: datetime.timedelta | None = ..., run_timeout: datetime.timedelta | None = ..., task_timeout: datetime.timedelta | None = ..., @@ -397,22 +388,19 @@ async def signal_with_start_workflow( # Overload case: -# - workflow method callable with list-form workflow arguments -# - signal method callable with a typed single signal arguments +# - workflow method callable with typed positional workflow arguments +# - signal callable with list-form signal arguments @typing.overload async def signal_with_start_workflow( workflow: collections.abc.Callable[ [SelfType, typing_extensions.Unpack[WorkflowArgs]], collections.abc.Awaitable[WorkflowResult], ], - *, - args: list[typing.Any], + *args: typing_extensions.Unpack[WorkflowArgs], id: str, task_queue: str, - signal: collections.abc.Callable[ - [SelfType, SignalArg], None | collections.abc.Awaitable[None] - ], - signal_args: SignalArg, + signal: collections.abc.Callable[..., None | collections.abc.Awaitable[None]], + signal_args: list[typing.Any], execution_timeout: datetime.timedelta | None = ..., run_timeout: datetime.timedelta | None = ..., task_timeout: datetime.timedelta | None = ..., @@ -432,16 +420,20 @@ async def signal_with_start_workflow( # Overload case: -# - workflow name with positional workflow arguments -# - signal callable with list-form signal arguments +# - workflow method callable with list-form workflow arguments +# - signal name with optional list-form signal arguments @typing.overload async def signal_with_start_workflow( - workflow: str, - *args: object, + workflow: collections.abc.Callable[ + [SelfType, typing_extensions.Unpack[WorkflowArgs]], + collections.abc.Awaitable[WorkflowResult], + ], + *, + args: list[typing.Any], id: str, task_queue: str, - signal: collections.abc.Callable[..., None | collections.abc.Awaitable[None]], - signal_args: list[typing.Any], + signal: str, + signal_args: list[typing.Any] | None = ..., execution_timeout: datetime.timedelta | None = ..., run_timeout: datetime.timedelta | None = ..., task_timeout: datetime.timedelta | None = ..., @@ -457,21 +449,25 @@ async def signal_with_start_workflow( start_delay: datetime.timedelta | None = ..., static_summary: str | None = ..., static_details: str | None = ..., -) -> ExternalWorkflowHandle[object]: ... +) -> ExternalWorkflowHandle[SelfType]: ... # Overload case: -# - workflow name with optional list-form workflow arguments -# - signal callable with list-form signal arguments +# - workflow method callable with list-form workflow arguments +# - signal method callable with no signal arguments @typing.overload async def signal_with_start_workflow( - workflow: str, + workflow: collections.abc.Callable[ + [SelfType, typing_extensions.Unpack[WorkflowArgs]], + collections.abc.Awaitable[WorkflowResult], + ], *, - args: list[typing.Any] | None = ..., + args: list[typing.Any], id: str, task_queue: str, - signal: collections.abc.Callable[..., None | collections.abc.Awaitable[None]], - signal_args: list[typing.Any], + signal: collections.abc.Callable[ + [SelfType], None | collections.abc.Awaitable[None] + ], execution_timeout: datetime.timedelta | None = ..., run_timeout: datetime.timedelta | None = ..., task_timeout: datetime.timedelta | None = ..., @@ -487,23 +483,26 @@ async def signal_with_start_workflow( start_delay: datetime.timedelta | None = ..., static_summary: str | None = ..., static_details: str | None = ..., -) -> ExternalWorkflowHandle[object]: ... +) -> ExternalWorkflowHandle[SelfType]: ... # Overload case: -# - workflow method callable with typed positional workflow arguments -# - signal callable with list-form signal arguments +# - workflow method callable with list-form workflow arguments +# - signal method callable with a typed single signal arguments @typing.overload async def signal_with_start_workflow( workflow: collections.abc.Callable[ [SelfType, typing_extensions.Unpack[WorkflowArgs]], collections.abc.Awaitable[WorkflowResult], ], - *args: typing_extensions.Unpack[WorkflowArgs], + *, + args: list[typing.Any], id: str, task_queue: str, - signal: collections.abc.Callable[..., None | collections.abc.Awaitable[None]], - signal_args: list[typing.Any], + signal: collections.abc.Callable[ + [SelfType, SignalArg], None | collections.abc.Awaitable[None] + ], + signal_args: SignalArg, execution_timeout: datetime.timedelta | None = ..., run_timeout: datetime.timedelta | None = ..., task_timeout: datetime.timedelta | None = ..., @@ -629,6 +628,11 @@ async def signal_with_start_workflow( Returns: A workflow handle to the started workflow. """ + if positional_args and args is not None: + raise TypeError("cannot specify both positional arguments and args") + normalized_args: list[typing.Any] | None = ( + list(positional_args) if positional_args else args + ) normalized_signal_args: list[typing.Any] | None if signal_args is None: normalized_signal_args = None @@ -636,11 +640,6 @@ async def signal_with_start_workflow( normalized_signal_args = typing.cast(list[typing.Any], signal_args) else: normalized_signal_args = [signal_args] - if positional_args and args is not None: - raise TypeError("cannot specify both positional arguments and args") - normalized_args: list[typing.Any] | None = ( - list(positional_args) if positional_args else args - ) user_metadata = ( None if static_summary is None and static_details is None diff --git a/temporalio/nexus/system/workflow_service/service.py b/temporalio/nexus/system/workflow_service/services.py similarity index 63% rename from temporalio/nexus/system/workflow_service/service.py rename to temporalio/nexus/system/workflow_service/services.py index 7ce5849ca..e2e901157 100644 --- a/temporalio/nexus/system/workflow_service/service.py +++ b/temporalio/nexus/system/workflow_service/services.py @@ -4,7 +4,10 @@ from nexusrpc import Operation, service -import temporalio.api.workflowservice.v1.request_response_pb2 +from .models import ( + SignalWithStartWorkflowRequest, + SignalWithStartWorkflowResponse, +) @service(name="temporal.api.workflowservice.v1.WorkflowService") @@ -16,6 +19,6 @@ class WorkflowService: # .. warning:: This API is experimental and subject to change. signal_with_start_workflow: Operation[ - temporalio.api.workflowservice.v1.request_response_pb2.SignalWithStartWorkflowExecutionRequest, - temporalio.api.workflowservice.v1.request_response_pb2.SignalWithStartWorkflowExecutionResponse, + SignalWithStartWorkflowRequest, + SignalWithStartWorkflowResponse, ] = Operation(name="SignalWithStartWorkflowExecution") diff --git a/tests/nexus/test_temporal_system_nexus.py b/tests/nexus/test_temporal_system_nexus.py index 4629d230b..4cef932e3 100644 --- a/tests/nexus/test_temporal_system_nexus.py +++ b/tests/nexus/test_temporal_system_nexus.py @@ -14,6 +14,7 @@ import temporalio.api.workflowservice.v1.request_response_pb2 as workflowservice_pb2 import temporalio.converter import temporalio.nexus.system as nexus_system +import temporalio.nexus.system.workflow_service.models as workflow_service_models from temporalio import workflow from temporalio.bridge._visitor import PayloadVisitor from temporalio.bridge.proto.workflow_completion.workflow_completion_pb2 import ( @@ -135,12 +136,12 @@ def _assert_start_nexus_operation_interceptor_trace() -> None: assert trace_name == "workflow.start_nexus_operation" trace_input = cast(StartNexusOperationInput[Any, Any], trace_value) request = cast( - workflowservice_pb2.SignalWithStartWorkflowExecutionRequest, + workflow_service_models.SignalWithStartWorkflowRequest, trace_input.input, ) - assert request.workflow_id == "system-nexus-workflow-id" - assert request.signal_name == "test-signal" - assert request.workflow_type.name == "test-workflow" + assert request.id == "system-nexus-workflow-id" + assert request.signal == "test-signal" + assert request.workflow == "test-workflow" class _MarkingPayloadVisitor: