Skip to content
Open
Show file tree
Hide file tree
Changes from 7 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
2 changes: 2 additions & 0 deletions temporalio/bridge/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -82,6 +82,8 @@ class ClientConfig:
http_connect_proxy_config: ClientHttpConnectProxyConfig | None
dns_load_balancing_config: ClientDnsLoadBalancingConfig | None
grpc_compression: str
payloads_warn_size: int
memo_warn_size: int


@dataclass
Expand Down
2 changes: 1 addition & 1 deletion temporalio/bridge/sdk-core
Submodule sdk-core updated 32 files
+6 −5 .github/workflows/heavy.yml
+54 −44 .github/workflows/per-pr.yml
+25 −0 CHANGELOG.md
+157 −1 crates/client/src/grpc.rs
+49 −4 crates/client/src/lib.rs
+26 −0 crates/client/src/options_structs.rs
+11 −0 crates/client/src/request_extensions.rs
+393 −116 crates/client/src/schedules.rs
+10 −6 crates/client/src/workflow_handle.rs
+23 −1 crates/common/build.rs
+95 −73 crates/common/src/payload_limits.rs
+127 −0 crates/common/src/telemetry/log_export.rs
+36 −9 crates/macros/src/lib.rs
+18 −0 crates/sdk-core-c-bridge/include/temporal-sdk-core-c-bridge.h
+14 −0 crates/sdk-core-c-bridge/src/client.rs
+2 −0 crates/sdk-core-c-bridge/src/tests/context.rs
+6 −0 crates/sdk-core-c-bridge/src/worker.rs
+125 −0 crates/sdk-core/src/core_tests/activity_tasks.rs
+2 −0 crates/sdk-core/src/telemetry/metrics.rs
+78 −21 crates/sdk-core/src/worker/activities.rs
+33 −2 crates/sdk-core/src/worker/activities/activity_heartbeat_manager.rs
+32 −18 crates/sdk-core/src/worker/client.rs
+26 −3 crates/sdk-core/src/worker/mod.rs
+55 −1 crates/sdk-core/src/worker/nexus.rs
+49 −8 crates/sdk-core/src/worker/tuner/resource_based.rs
+77 −21 crates/sdk-core/src/worker/workflow/mod.rs
+35 −3 crates/sdk-core/tests/cloud_tests.rs
+36 −44 crates/sdk-core/tests/integ_tests/schedule_tests.rs
+379 −5 crates/sdk-core/tests/integ_tests/worker_tests.rs
+1 −0 crates/sdk-core/tests/shared_tests/mod.rs
+28 −0 crates/sdk/src/lib.rs
+2 −0 mise.toml
6 changes: 6 additions & 0 deletions temporalio/bridge/src/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,8 @@ pub struct ClientConfig {
http_connect_proxy_config: Option<ClientHttpConnectProxyConfig>,
dns_load_balancing_config: Option<ClientDnsLoadBalancingConfig>,
grpc_compression: String,
payloads_warn_size: u64,
memo_warn_size: u64,
}

#[derive(FromPyObject)]
Expand Down Expand Up @@ -268,6 +270,10 @@ impl ClientConfig {
.maybe_http_connect_proxy(self.http_connect_proxy_config.map(Into::into))
.dns_load_balancing(dns_load_balancing)
.grpc_compression(grpc_compression_from_str(&self.grpc_compression)?)
.payload_limits(temporalio_client::PayloadLimitsOptions {
payloads_warn_size: self.payloads_warn_size,
memo_warn_size: self.memo_warn_size,
})
.headers(ascii_headers)
.binary_headers(binary_headers)
.maybe_api_key(self.api_key)
Expand Down
2 changes: 2 additions & 0 deletions temporalio/bridge/src/worker.rs
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,7 @@ pub struct WorkerConfig {
nexus_task_poller_behavior: PollerBehavior,
plugins: Vec<String>,
storage_drivers: HashSet<String>,
disable_payload_error_limit: bool,
}

#[derive(FromPyObject)]
Expand Down Expand Up @@ -770,6 +771,7 @@ fn convert_worker_config(
.map(|r#type| StorageDriverInfo { r#type })
.collect::<HashSet<_>>(),
)
.disable_payload_error_limit(conf.disable_payload_error_limit)
.build()
.map_err(|err| PyValueError::new_err(format!("Invalid worker config: {err}")))
}
Expand Down
28 changes: 3 additions & 25 deletions temporalio/bridge/worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,7 @@ class WorkerConfig:
nexus_task_poller_behavior: PollerBehavior
plugins: Sequence[str]
storage_drivers: set[str]
disable_payload_error_limit: bool


@dataclass
Expand Down Expand Up @@ -284,10 +285,8 @@ class _Visitor(VisitorFunctions):
def __init__(
self,
f: Callable[[Sequence[Payload]], Awaitable[list[Payload]]],
visit_system_nexus_envelope: Callable[[Payload], Awaitable[None]] | None = None,
):
self._f = f
self._visit_system_nexus_envelope = visit_system_nexus_envelope

async def visit_payload(self, payload: Payload) -> None:
new_payload = (await self._f([payload]))[0]
Expand All @@ -303,10 +302,6 @@ async def visit_payloads(self, payloads: PayloadSequence) -> None:
del payloads[:]
payloads.extend(new_payloads)

async def visit_system_nexus_envelope(self, payload: Payload) -> None:
if self._visit_system_nexus_envelope is not None:
await self._visit_system_nexus_envelope(payload)


async def decode_activation(
activation: temporalio.bridge.proto.workflow_activation.WorkflowActivation,
Expand Down Expand Up @@ -348,39 +343,22 @@ async def encode_completion(
Returns:
Metrics from any external storage store operations that occurred.
"""

async def _validate_system_nexus_envelope(payload: Payload) -> None:
data_converter._validate_payload_limits([payload])

await CommandAwarePayloadVisitor(
skip_search_attributes=True,
skip_headers=not encode_headers,
).visit(
_Visitor(
data_converter._encode_payload_sequence,
visit_system_nexus_envelope=_validate_system_nexus_envelope,
),
_Visitor(data_converter._encode_payload_sequence),
completion,
)

async def _store_and_validate(
payloads: Sequence[Payload],
) -> list[Payload]:
stored = await data_converter._external_store_payload_sequence(payloads)
data_converter._validate_payload_limits(stored)
return stored

metrics = temporalio.converter._extstore.StorageOperationMetrics()
with metrics.track():
await CommandAwarePayloadVisitor(
skip_search_attributes=True,
skip_headers=not encode_headers,
concurrency_limit=storage_concurrency_limit,
).visit(
_Visitor(
_store_and_validate,
visit_system_nexus_envelope=_validate_system_nexus_envelope,
),
_Visitor(data_converter._external_store_payload_sequence),
completion,
)

Expand Down
2 changes: 2 additions & 0 deletions temporalio/client/_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -232,6 +232,8 @@ async def connect(
http_connect_proxy_config=http_connect_proxy_config,
dns_load_balancing_config=dns_load_balancing_config,
grpc_compression=grpc_compression,
payloads_warn_size=data_converter.payload_limits.payload_size_warning,
memo_warn_size=data_converter.payload_limits.memo_size_warning,
)

def make_lambda(
Expand Down
61 changes: 0 additions & 61 deletions temporalio/converter/_data_converter.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@
from __future__ import annotations

import dataclasses
import warnings
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from logging import getLogger
Expand Down Expand Up @@ -32,9 +31,6 @@
)
from temporalio.converter._payload_limits import (
PayloadLimitsConfig,
PayloadSizeWarning,
_PayloadSizeError,
_ServerPayloadErrorLimits,
)
from temporalio.converter._serialization_context import (
SerializationContext,
Expand Down Expand Up @@ -99,9 +95,6 @@ class DataConverter(WithSerializationContext):
default: ClassVar[DataConverter]
"""Singleton default data converter."""

_payload_error_limits: _ServerPayloadErrorLimits | None = None
"""Server-reported limits for payloads."""

def __post_init__(self) -> None: # noqa: D105
object.__setattr__(self, "payload_converter", self.payload_converter_class())
object.__setattr__(self, "failure_converter", self.failure_converter_class())
Expand All @@ -124,7 +117,6 @@ async def encode(
payloads = self.payload_converter.to_payloads(values)
payloads = await self._encode_payload_sequence(payloads)
payloads = await self._external_store_payload_sequence(payloads)
self._validate_payload_limits(payloads)
return payloads

async def decode(
Expand Down Expand Up @@ -230,11 +222,6 @@ def _with_contexts(
"""Return an instance with both serialization and store contexts applied."""
return self.with_context(serialization_ctx)._with_store_context(store_ctx)

def _with_payload_error_limits(
self, limits: _ServerPayloadErrorLimits | None
) -> DataConverter:
return dataclasses.replace(self, _payload_error_limits=limits)

async def _decode_memo(
self,
source: temporalio.api.common.v1.Memo,
Expand Down Expand Up @@ -273,16 +260,6 @@ async def _encode_memo_existing(
if not isinstance(v, temporalio.api.common.v1.Payload):
payload = (await self.encode([v]))[0]
memo.fields[k].CopyFrom(payload)
# Memos have their field payloads validated all together in one unit
DataConverter._validate_limits(
list(memo.fields.values()),
self._payload_error_limits.memo_size_error
if self._payload_error_limits
else None,
"[TMPRL1103] Attempted to upload memo with size that exceeded the error limit.",
self.payload_limits.memo_size_warning,
"[TMPRL1103] Attempted to upload memo with size that exceeded the warning limit.",
)

async def _transform_outbound_payload(
self, payload: temporalio.api.common.v1.Payload
Expand All @@ -291,7 +268,6 @@ async def _transform_outbound_payload(
payload = (await self.payload_codec.encode([payload]))[0]
if self.external_storage:
payload = await self.external_storage._store_payload(payload)
self._validate_payload_limits([payload])
return payload

async def _transform_outbound_payloads(
Expand All @@ -301,7 +277,6 @@ async def _transform_outbound_payloads(
await self.payload_codec.encode_wrapper(payloads)
if self.external_storage:
await self.external_storage._store_payloads(payloads)
self._validate_payload_limits(payloads.payloads)

async def _transform_inbound_payload(
self, payload: temporalio.api.common.v1.Payload
Expand Down Expand Up @@ -376,42 +351,6 @@ async def _decode_payload_sequence(
def _decode_payload_has_effect(self) -> bool:
return self.payload_codec is not None or self.external_storage is not None

def _validate_payload_limits(
self,
payloads: Sequence[temporalio.api.common.v1.Payload],
):
DataConverter._validate_limits(
payloads,
self._payload_error_limits.payload_size_error
if self._payload_error_limits
else None,
"[TMPRL1103] Attempted to upload payloads with size that exceeded the error limit.",
self.payload_limits.payload_size_warning,
"[TMPRL1103] Attempted to upload payloads with size that exceeded the warning limit.",
)

@staticmethod
def _validate_limits(
payloads: Sequence[temporalio.api.common.v1.Payload],
error_limit: int | None,
error_message: str,
warning_limit: int,
warning_message: str,
):
total_size = sum(payload.ByteSize() for payload in payloads)

if error_limit and error_limit > 0 and total_size > error_limit:
raise _PayloadSizeError(
f"{error_message} Size: {total_size} bytes, Limit: {error_limit} bytes"
)

if warning_limit > 0 and total_size > warning_limit:
# TODO: Use a context aware logger to log extra information about workflow/activity/etc
warnings.warn(
f"{warning_message} Size: {total_size} bytes, Limit: {warning_limit} bytes",
PayloadSizeWarning,
)


def default() -> DataConverter:
"""Default data converter.
Expand Down
5 changes: 1 addition & 4 deletions temporalio/converter/_failure_converter.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,6 @@
import temporalio.api.failure.v1
import temporalio.exceptions
from temporalio.converter._payload_converter import PayloadConverter
from temporalio.converter._payload_limits import _PayloadSizeError

logger = getLogger("temporalio.converter")

Expand Down Expand Up @@ -108,9 +107,7 @@ def to_failure(
# Convert to failure error
failure_error = temporalio.exceptions.ApplicationError(
str(exception),
type="PayloadSizeError"
if isinstance(exception, _PayloadSizeError)
else exception.__class__.__name__,
type=exception.__class__.__name__,
)
failure_error.__traceback__ = exception.__traceback__
failure_error.__cause__ = exception.__cause__
Expand Down
36 changes: 7 additions & 29 deletions temporalio/converter/_payload_limits.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,44 +4,22 @@

from dataclasses import dataclass

import temporalio.exceptions


@dataclass(frozen=True)
class PayloadLimitsConfig:
"""Configuration for when payload sizes exceed limits."""

memo_size_warning: int = 2 * 1024
"""The limit (in bytes) at which a memo size warning is logged."""
"""The limit (in bytes) at which a memo size warning is logged. Set to 0 to disable."""

payload_size_warning: int = 512 * 1024
"""The limit (in bytes) at which a payload size warning is logged."""
"""The limit (in bytes) at which a payload size warning is logged. Set to 0 to disable."""


class PayloadSizeWarning(RuntimeWarning):
"""The size of payloads is above the warning limit."""


class _PayloadSizeError(temporalio.exceptions.TemporalError): # type:ignore[reportUnusedClass]
"""Error raised when payloads size exceeds payload size limits."""

def __init__(self, message: str):
"""Initialize a payloads size error."""
super().__init__(message)
self._message = message

@property
def message(self) -> str:
"""Message."""
return self._message


@dataclass(frozen=True)
class _ServerPayloadErrorLimits: # type:ignore[reportUnusedClass]
"""Error limits for payloads as described by the Temporal server."""

memo_size_error: int
"""The limit (in bytes) at which a memo size error is raised."""
"""The size of payloads is above the warning limit.

payload_size_error: int
"""The limit (in bytes) at which a payload size error is raised."""
.. deprecated::
Payload size warnings are no longer raised through the :mod:`warnings` module. This
symbol is retained for backwards compatibility and is no longer raised.
"""
1 change: 1 addition & 0 deletions temporalio/runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -180,6 +180,7 @@ def formatted(self) -> str:
# We intentionally aren't using __str__ or __format__ so they can keep
# their original dataclass impls
targets = [
"temporalio_common",
"temporalio_sdk_core",
"temporalio_client",
"temporalio_sdk",
Expand Down
4 changes: 4 additions & 0 deletions temporalio/service.py
Original file line number Diff line number Diff line change
Expand Up @@ -208,6 +208,8 @@ class ConnectConfig:
http_connect_proxy_config: HttpConnectProxyConfig | None = None
dns_load_balancing_config: DnsLoadBalancingConfig | None = None
grpc_compression: GrpcCompression = GrpcCompression.GZIP
payloads_warn_size: int = 512 * 1024
memo_warn_size: int = 2 * 1024

def __post_init__(self) -> None:
"""Set extra defaults on unset properties."""
Expand Down Expand Up @@ -271,6 +273,8 @@ def _to_bridge_config(self) -> temporalio.bridge.client.ClientConfig:
else None
),
grpc_compression=self.grpc_compression._to_bridge_config(),
payloads_warn_size=self.payloads_warn_size,
memo_warn_size=self.memo_warn_size,
)


Expand Down
Loading
Loading