Skip to content
Closed
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
75 changes: 60 additions & 15 deletions scripts/nex_gen_support.py
Original file line number Diff line number Diff line change
@@ -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(
Expand Down Expand Up @@ -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]

Expand All @@ -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:
Expand All @@ -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:
Expand All @@ -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],
)


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


Expand Down Expand Up @@ -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:
Expand All @@ -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")
12 changes: 10 additions & 2 deletions temporalio/activity.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 2 additions & 0 deletions temporalio/converter/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@
JSONTypeConverter,
JSONTypeConverterUnhandled,
PayloadConverter,
TemporalIntermediatePayloadConverter,
value_to_type,
)
from temporalio.converter._payload_limits import (
Expand Down Expand Up @@ -83,6 +84,7 @@
"PayloadLimitsConfig",
"PayloadSizeWarning",
"SerializationContext",
"TemporalIntermediatePayloadConverter",
"WithSerializationContext",
"WorkflowSerializationContext",
"decode_search_attributes",
Expand Down
7 changes: 6 additions & 1 deletion temporalio/converter/_data_converter.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
)
from temporalio.converter._payload_converter import (
PayloadConverter,
TemporalIntermediatePayloadConverter,
)
from temporalio.converter._payload_limits import (
PayloadLimitsConfig,
Expand Down Expand Up @@ -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(
Expand Down
77 changes: 77 additions & 0 deletions temporalio/converter/_payload_converter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
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)
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.

Expand Down
Loading
Loading