diff --git a/CHANGELOG.md b/CHANGELOG.md index 33366c1d..ff166885 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,12 @@ adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). ## Unreleased +FIXED + +- Asynchronous Azure Blob payload uploads and downloads no longer run gzip +compression or decompression on the calling event loop, keeping concurrent +async operations responsive during large transfers. + ## v1.10.1 FIXED diff --git a/azure-functions-durable/CHANGELOG.md b/azure-functions-durable/CHANGELOG.md index b2b91818..206c9acc 100644 --- a/azure-functions-durable/CHANGELOG.md +++ b/azure-functions-durable/CHANGELOG.md @@ -7,6 +7,18 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## Unreleased +ADDED + +- Added `DFApp.configure_large_payloads(payload_store=...)` to externalize large +durable payloads to Azure Blob Storage or a custom payload store. +- Configured clients automatically hydrate stored payloads, including +orchestration history and entity operation inputs and results. + +FIXED + +- Preserved the application's source directory in `context.function_directory` +for decorated activities and durable-client functions. + ## v2.0.0rc1 CHANGED diff --git a/azure-functions-durable/README.md b/azure-functions-durable/README.md index 1cd626c4..e441e75d 100644 --- a/azure-functions-durable/README.md +++ b/azure-functions-durable/README.md @@ -32,6 +32,129 @@ Key capabilities include durable orchestrations and sub-orchestrations, durable timers, external events, durable entities, retries, versioning, durable HTTP calls (`context.call_http(...)`), recurring scheduled tasks, and history export. +## Large payloads + +Configure a `durabletask.payload.PayloadStore` once at app startup to store large +serialized payloads outside orchestration history. For Azure Blob Storage, install +the optional dependencies: + +```bash +pip install azure-functions-durable "durabletask[azure-blob-payloads]" aiohttp +``` + +In your Function app, configure the root `DFApp` before any invocations: + +```python +import os + +import azure.durable_functions as df +from durabletask.extensions.azure_blob_payloads import ( + BlobPayloadStore, + BlobPayloadStoreOptions, +) + +app = df.DFApp() +app.configure_large_payloads( + payload_store=BlobPayloadStore(BlobPayloadStoreOptions( + connection_string=os.environ["PAYLOAD_STORAGE_CONNECTION_STRING"], + container_name="durable-payloads", + threshold_bytes=256 * 1024, + )) +) +``` + +Set `PAYLOAD_STORAGE_CONNECTION_STRING` in your Function app settings (or in +`local.settings.json` for local development). The store automatically uploads +serialized payloads above the threshold and downloads their contents when the +SDK consumes them. Orchestration and activity inputs and outputs, custom status, +external events, and entity inputs, results, and state use the configured store. +Sub-orchestrations and continue-as-new use it as well. The default maximum stored +payload size is 10 MiB; `max_stored_payload_bytes` can configure this limit. + +Configuration applies to both synchronous and asynchronous durable clients and +all registered blueprints, including blueprints imported before configuration. +There is one store per Python worker process. Registering the same store object +again is allowed; registering a different object raises `ValueError`. Configure +every scaled-out worker with access to the same backing storage and retain that +access across deployments. Keep the store open for the process lifetime. + +> [!WARNING] +> Keep payload blobs for as long as any retained orchestration history or entity +> state references them, including histories needed for replay. Purging an +> orchestration does not delete its payload blobs; manage retention separately. + +This is SDK-managed storage, separate from the Azure Storage backend's automatic +large-message handling. Without configuration, the SDK keeps payloads inline. +Use the configured Python clients to retrieve hydrated payloads. Host management +HTTP endpoints and other consumers that do not use this configuration can expose +reference strings instead. Applications exchanging externalized payloads must +agree on the store and reference encoding; Functions references are JSON strings. + +> [!WARNING] +> Storage failures can fail durable invocations, including orchestrations. +> Storage transport retries are separate from durable activity retry policies. +> This SDK does not add an activity retry policy or guarantee that the Functions +> host abandons and redelivers a work item after a storage failure. A transient +> storage error can therefore become a terminal orchestration failure. + +Registered orchestration and entity handlers await the store's async methods +before and after execution. Orchestrators remain synchronous generators, and +orchestration replay and entity code run on execution threads with their +invocation logging context preserved. Their payload downloads and uploads do +not occupy those threads, and serialization does not access storage during +replay. Custom stores must implement genuinely nonblocking async methods to +benefit from this behavior. + +For orchestration and entity execution, the SDK reuses the Functions runtime's +thread pool when the runtime exposes it; otherwise it uses a process-wide SDK +pool. Both honor `PYTHON_THREADPOOL_THREAD_COUNT`. + +Activities retain their synchronous or asynchronous calling convention. +Synchronous activities use synchronous storage inside the host-managed execution +thread and remain directly callable without `await`; async activities await +async storage. Direct calls to decorated activities return ordinary Python values +without accessing payload storage; transport processing applies only to host +binding invocations. Binding converters perform no storage I/O. Synchronous functions +still receive the synchronous durable client, and synchronous client APIs use +synchronous storage. Direct `Orchestrator.handle()` and `Orchestrator.create()` +adapters also remain synchronous. + +Both client history APIs hydrate entity operation inputs and results, including +values nested in the host's entity protocol envelopes. During orchestration +replay, nested entity results are hydrated, but historical nested request inputs +are not downloaded because replay only needs their correlation metadata. +Historical scheduled activity inputs are also not downloaded during replay; +explicit history retrieval continues to hydrate those inputs. + +> [!WARNING] +> With payload storage configured, whole payload strings recognized by the +> store's `is_known_token()` are reserved references, not literal application +> data. For `BlobPayloadStore`, this includes strings of the form +> `blob:v1::`, with nonempty container and blob names. +> Functions recognizes both raw and JSON-quoted references. A matching string +> is treated as already externalized on output and downloaded on input, even +> below the size threshold. Missing or inaccessible references raise errors; +> they do not fall back to literal strings. + +To pass a reference as application data for later retrieval, wrap it in an +object, for example `{"reference": "blob:v1:container:blob"}`. Reference detection +does not recursively inspect strings inside application JSON objects. The +wrapper preserves the literal reference whether the object stays inline or is +itself externalized. Keep the wrapper whenever passing that value across a +durable payload boundary; passing its string field alone opts back into reference +interpretation. Custom payload stores define their own reserved token syntax. + +> [!WARNING] +> Recognized references are trusted transport inputs, not authorization +> boundaries. `BlobPayloadStore` reads from the container named in the token +> using its configured credentials; `container_name` selects the upload +> container and does not restrict downloads. An account-wide connection string +> can therefore allow reads outside that container. Use least-privilege +> credentials scoped to the intended payload storage, and explicitly decide +> whether external callers may supply references. Reject untrusted references +> or validate their allowed storage locations before passing them into durable +> APIs; token recognition alone does not authorize a read. + ## Unit testing entities Use `execute_entity()` to run one entity operation in-process without a diff --git a/azure-functions-durable/azure/durable_functions/client.py b/azure-functions-durable/azure/durable_functions/client.py index 751c27e1..5b637b55 100644 --- a/azure-functions-durable/azure/durable_functions/client.py +++ b/azure-functions-durable/azure/durable_functions/client.py @@ -8,11 +8,12 @@ import threading from datetime import datetime, timedelta -from typing import Any, Mapping, Optional, Union, cast +from typing import Any, Mapping, Optional, Union, cast, override from warnings import deprecated import azure.functions as func from urllib.parse import urlparse, quote +from durabletask import history from durabletask.client import ( AsyncTaskHubGrpcClient, OrchestrationQuery, @@ -26,6 +27,11 @@ AzureFunctionsDefaultClientInterceptorImpl, ) from .internal.serialization import DEFAULT_FUNCTIONS_DATA_CONVERTER +from .internal.payloads import ( + get_transport_payload_store, + hydrate_entity_history, + hydrate_entity_history_async, +) from .http.http_management_payload import HttpManagementPayload, replace_url_origin from .internal.compat.durable_orchestration_status import DurableOrchestrationStatus from .internal.compat.entity_state_response import EntityStateResponse @@ -178,6 +184,7 @@ def __init__(self, client_as_string: str): interceptors=interceptors, channel_options=channel_options, data_converter=DEFAULT_FUNCTIONS_DATA_CONVERTER, + payload_store=get_transport_payload_store(), emit_trace_spans=False, logger=_LOGGER) @@ -193,6 +200,16 @@ def __init__(self, client_as_string: str): self._creation_loop = None self._close_scheduled = False + @override + async def get_orchestration_history( + self, instance_id: str, *, execution_id: str | None = None, + for_work_item_processing: bool = False) -> list[history.HistoryEvent]: + events = await super().get_orchestration_history( + instance_id, execution_id=execution_id, + for_work_item_processing=for_work_item_processing) + await hydrate_entity_history_async(events, self._payload_store, instance_id) + return events + def schedule_close(self) -> None: """Schedule the underlying gRPC channel to close after the invocation. @@ -664,9 +681,20 @@ def __init__(self, client_as_string: str): interceptors=interceptors, channel_options=channel_options, data_converter=DEFAULT_FUNCTIONS_DATA_CONVERTER, + payload_store=get_transport_payload_store(), emit_trace_spans=False, logger=_LOGGER) + @override + def get_orchestration_history( + self, instance_id: str, *, execution_id: str | None = None, + for_work_item_processing: bool = False) -> list[history.HistoryEvent]: + events = super().get_orchestration_history( + instance_id, execution_id=execution_id, + for_work_item_processing=for_work_item_processing) + hydrate_entity_history(events, self._payload_store, instance_id) + return events + @classmethod def get_cached(cls, client_as_string: str) -> "SyncDurableFunctionsClient": """Get the process-wide client for a durable-client binding configuration. diff --git a/azure-functions-durable/azure/durable_functions/decorators/durable_app.py b/azure-functions-durable/azure/durable_functions/decorators/durable_app.py index 6747186b..70ce34c5 100644 --- a/azure-functions-durable/azure/durable_functions/decorators/durable_app.py +++ b/azure-functions-durable/azure/durable_functions/decorators/durable_app.py @@ -10,6 +10,7 @@ from azure.functions.decorators.function_app import DecoratorApi, FunctionBuilder from durabletask import task +from durabletask.payload import PayloadStore from .metadata import OrchestrationTrigger, ActivityTrigger, EntityTrigger, \ DurableClient @@ -19,9 +20,9 @@ builtin_http_activity, builtin_http_poll_orchestrator, ) -from ..internal.compat.activity import wrap_activity +from ..internal.compat.activity import wrap_activity, wrap_activity_payloads +from ..internal.invocation import preserve_function_source, wrap_invocation from ..worker import DurableFunctionsWorker -from ..orchestrator import Orchestrator class Blueprint(TriggerApi, BindingApi): @@ -166,6 +167,7 @@ def configure_history_export(self, writer: Any) -> None: def _configure_orchestrator_callable( self, wrap: Callable[[Callable[..., Any]], FunctionBuilder], + context_name: str, input_type: Optional[type] = None ) -> Callable[[task.Orchestrator[Any, Any]], FunctionBuilder]: """Obtain decorator to construct an Orchestrator class from a user-defined Function. @@ -193,7 +195,13 @@ def decorator(orchestrator_func: task.Orchestrator[Any, Any]) -> FunctionBuilder # feed it to a v1-style ``context.get_input()``. orchestrator_func._df_input_type = input_type # type: ignore[attr-defined] # noqa: E501 - handle = Orchestrator.create(orchestrator_func) + worker = DurableFunctionsWorker() + + async def handle(context: func.OrchestrationContext) -> str: + return await worker.execute_orchestration_request_async(orchestrator_func, context) + + handle.orchestrator_function = orchestrator_func # pyright: ignore[reportFunctionMemberAccess] + handle = wrap_invocation(handle, "context", registered_trigger_name=context_name) # invoke next decorator, with the Orchestrator as input handle.__name__ = orchestrator_func.__name__ @@ -204,6 +212,7 @@ def decorator(orchestrator_func: task.Orchestrator[Any, Any]) -> FunctionBuilder def _configure_entity_callable( self, wrap: Callable[[Callable[..., Any]], FunctionBuilder], + context_name: str, entity_name: Optional[str] = None ) -> Callable[[task.Entity[Any, Any]], FunctionBuilder]: """Obtain decorator to construct an Entity class from a user-defined Function. @@ -235,18 +244,16 @@ def decorator(entity_func: task.Entity[Any, Any]) -> FunctionBuilder: # Construct an orchestrator based on the end-user code worker = DurableFunctionsWorker() - # TODO: Because this handle method is the one actually exposed to the Functions SDK decorator, - # the parameter name will always be "context" here, even if the user specified a different name. - # We need to find a way to allow custom context names (like "ctx"). # The generated handle is what the Azure Functions host registers, # so its ``context`` parameter must be annotated with # ``azure.functions.EntityContext`` for the host's entityTrigger # binding converter to accept it; at runtime the host passes that # transport context (exposing ``.body``). - def handle(context: func.EntityContext) -> str: - return worker.execute_entity_batch_request(entity_func, context) + async def handle(context: func.EntityContext) -> str: + return await worker.execute_entity_batch_request_async(entity_func, context) handle.entity_function = entity_func # pyright: ignore[reportFunctionMemberAccess] + handle = wrap_invocation(handle, "context", registered_trigger_name=context_name) # invoke next decorator, with the Entity as input handle.__name__ = entity_func.__name__ @@ -294,14 +301,15 @@ def orchestration_trigger(self, context_name: str, def wrap(fb: FunctionBuilder) -> FunctionBuilder: def decorator() -> FunctionBuilder: + registered = fb._function._func # pyright: ignore[reportPrivateUsage] fb.add_trigger( - trigger=OrchestrationTrigger(name=context_name, + trigger=OrchestrationTrigger(name=getattr(registered, "_df_trigger_name", context_name), orchestration=orchestration)) return fb return decorator() - return self._configure_orchestrator_callable(wrap, input_type=input_type) + return self._configure_orchestrator_callable(wrap, context_name, input_type=input_type) def activity_trigger(self, input_name: str, activity: Optional[str] = None @@ -326,7 +334,13 @@ def decorator(user_fn: Callable[..., Any]) -> FunctionBuilder: # Adapt a durabletask-native two-argument activity ((ctx, input)) # to the host's single-input convention; one-argument activities # pass through unchanged. - return wrap(wrap_activity(user_fn, input_name)) + function = (user_fn._function._func # pyright: ignore[reportPrivateUsage] + if isinstance(user_fn, FunctionBuilder) else user_fn) + registered = wrap_activity_payloads(wrap_activity(function, input_name), input_name) + if isinstance(user_fn, FunctionBuilder): + user_fn._function._func = registered # pyright: ignore[reportPrivateUsage] + return wrap(user_fn) + return wrap(registered) return decorator @@ -347,14 +361,15 @@ def entity_trigger(self, @self._build_function def wrap(fb: FunctionBuilder) -> FunctionBuilder: def decorator() -> FunctionBuilder: + registered = fb._function._func # pyright: ignore[reportPrivateUsage] fb.add_trigger( - trigger=EntityTrigger(name=context_name, + trigger=EntityTrigger(name=getattr(registered, "_df_trigger_name", context_name), entity_name=entity_name)) return fb return decorator() - return self._configure_entity_callable(wrap, entity_name) + return self._configure_entity_callable(wrap, context_name, entity_name) def durable_client_input(self, client_name: str, @@ -439,6 +454,7 @@ def set_client_metadata(client_bound: Callable[..., Any]) -> None: annotations[client_name] = str client_bound.__annotations__ = annotations setattr(client_bound, "client_function", function) + preserve_function_source(client_bound, function) if is_async_function: @wraps(function) @@ -477,6 +493,19 @@ class DFApp(Blueprint, FunctionRegister): Exports the decorators required to declare and index DF Function-types. """ + def configure_large_payloads(self, *, payload_store: PayloadStore) -> None: + """Enable payload externalization for this app and its blueprints. + + Call once at app startup in every worker process, before invocations. + The store is shared by all durable clients, orchestrations, entities, + and activities in the process. All scaled-out workers must have access + to the same backing storage. Registering a different store in the same + process raises ValueError. Re-registering the same object is allowed. + """ + from ..internal.payloads import configure_payload_store + + configure_payload_store(payload_store) + def register_functions(self, function_container: DecoratorApi) -> None: """Register the functions of a blueprint into this app. diff --git a/azure-functions-durable/azure/durable_functions/internal/compat/activity.py b/azure-functions-durable/azure/durable_functions/internal/compat/activity.py index 05e3e391..ee046c88 100644 --- a/azure-functions-durable/azure/durable_functions/internal/compat/activity.py +++ b/azure-functions-durable/azure/durable_functions/internal/compat/activity.py @@ -16,8 +16,10 @@ from __future__ import annotations import inspect +import json import keyword import typing +from functools import wraps from collections.abc import ( Mapping, MutableMapping, @@ -28,6 +30,17 @@ ) from typing import Any, Callable, cast +from azure.functions._durable_functions import df_loads + +from ..converters import ActivityTriggerConverter +from ..invocation import preserve_function_source +from ..payloads import ( + ActivityPayload, + deexternalize_payload, + deexternalize_payload_async, + externalize_activity_output, + externalize_activity_output_async, +) from .orchestration_context import accepts_two_positional_args @@ -125,9 +138,11 @@ def wrap_activity(fn: Callable[..., Any], input_name: str) -> Callable[..., Any] # and the original is invoked positionally (its second-parameter name is # irrelevant). namespace: dict[str, Any] = {"_fn": fn, "_ctx": _NO_ACTIVITY_CONTEXT} + async_prefix = "async " if inspect.iscoroutinefunction(fn) else "" + await_prefix = "await " if async_prefix else "" exec( # noqa: S102 - input_name is validated to be a bare identifier above - f"def _activity_adapter({input_name}):\n" - f" return _fn(_ctx, {input_name})\n", + f"{async_prefix}def _activity_adapter({input_name}):\n" + f" return {await_prefix}_fn(_ctx, {input_name})\n", namespace, ) adapter = cast("Callable[..., Any]", namespace["_activity_adapter"]) @@ -159,4 +174,55 @@ def wrap_activity(fn: Callable[..., Any], input_name: str) -> Callable[..., Any] if ret_ann is not inspect.Parameter.empty: annotations["return"] = ret_ann adapter.__annotations__ = annotations + preserve_function_source(adapter, fn) return adapter + + +def wrap_activity_payloads(fn: Callable[..., Any], input_name: str) -> Callable[..., Any]: + """Run payload I/O inside the invocation, preserving sync/async dispatch.""" + signature = inspect.signature(fn) + + def decode(value: str) -> Any: + try: + return df_loads(value) + except json.JSONDecodeError: + return value + except Exception as error: + raise ValueError('activity trigger input must be a string or a ' + f'valid json serializable ({value})') from error + + wrapper: Callable[..., Any] + if inspect.iscoroutinefunction(fn): + @wraps(fn) + async def async_wrapper(*args: Any, **kwargs: Any) -> Any: + bound = signature.bind(*args, **kwargs) + value = bound.arguments.get(input_name) + if not isinstance(value, ActivityPayload): + return await fn(*args, **kwargs) + bound.arguments[input_name] = decode(await deexternalize_payload_async(value.value)) + result = await fn(*bound.args, **bound.kwargs) + if result is None: + return None + encoded = ActivityTriggerConverter.encode(result, expected_type=None).value + return ActivityPayload(await externalize_activity_output_async(encoded)) + + wrapper = async_wrapper + else: + @wraps(fn) + def sync_wrapper(*args: Any, **kwargs: Any) -> Any: + bound = signature.bind(*args, **kwargs) + value = bound.arguments.get(input_name) + if not isinstance(value, ActivityPayload): + return fn(*args, **kwargs) + bound.arguments[input_name] = decode(deexternalize_payload(value.value)) + result = fn(*bound.args, **bound.kwargs) + if result is None: + return None + encoded = ActivityTriggerConverter.encode(result, expected_type=None).value + return ActivityPayload(externalize_activity_output(encoded)) + + wrapper = sync_wrapper + + setattr(wrapper, "__signature__", signature) + preserve_function_source(wrapper, fn) + return wrapper diff --git a/azure-functions-durable/azure/durable_functions/internal/converters.py b/azure-functions-durable/azure/durable_functions/internal/converters.py index c0892f4c..34eadd38 100644 --- a/azure-functions-durable/azure/durable_functions/internal/converters.py +++ b/azure-functions-durable/azure/durable_functions/internal/converters.py @@ -42,6 +42,7 @@ ENTITY_TRIGGER, ORCHESTRATION_TRIGGER, ) +from .payloads import ActivityPayload, get_transport_payload_store _TriggerMetadata = Optional[Mapping[str, meta.Datum]] @@ -136,11 +137,15 @@ def decode(cls, data: meta.Datum, *, # carrying a custom-object envelope surfaces as TypeError below and is # re-raised as ValueError. if data_type in ['string', 'json']: + value = data.value + store = get_transport_payload_store() + if store is not None: + return ActivityPayload(value) try: - result = df_loads(data.value) + result = df_loads(value) except json.JSONDecodeError: # String failover if the content is not json serializable - result = data.value + result = value except Exception as e: raise ValueError( 'activity trigger input must be a string or a ' @@ -154,6 +159,8 @@ def decode(cls, data: meta.Datum, *, @classmethod def encode(cls, obj: Any, *, expected_type: Optional[type]) -> meta.Datum: + if isinstance(obj, ActivityPayload): + return meta.Datum(type='json', value=obj.value) try: result = df_dumps(obj) except TypeError as e: diff --git a/azure-functions-durable/azure/durable_functions/internal/invocation.py b/azure-functions-durable/azure/durable_functions/internal/invocation.py new file mode 100644 index 00000000..32011d05 --- /dev/null +++ b/azure-functions-durable/azure/durable_functions/internal/invocation.py @@ -0,0 +1,125 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Async host invocation context and offloading of synchronous durable code.""" + +import asyncio +import inspect +import logging +import os +import sys +from concurrent.futures import ThreadPoolExecutor +from contextvars import ContextVar, copy_context +from functools import lru_cache, wraps +from types import FunctionType +from typing import Any, Callable, ParamSpec, TypeVar + +import azure.functions as func + +_invocation_context: ContextVar[func.Context | None] = ContextVar("durable_invocation_context", default=None) +_Parameters = ParamSpec("_Parameters") +_Result = TypeVar("_Result") + + +def preserve_function_source(wrapper: Callable[..., Any], original: Callable[..., Any]) -> None: + """Preserve the filename used by the Functions worker to index the app directory.""" + if isinstance(wrapper, FunctionType) and isinstance(original, FunctionType): + wrapper.__code__ = wrapper.__code__.replace(co_filename=original.__code__.co_filename) + + +@lru_cache(maxsize=1) +def _fallback_executor() -> ThreadPoolExecutor: + setting = os.environ.get("PYTHON_THREADPOOL_THREAD_COUNT") + max_workers = None + if setting is not None: + try: + max_workers = int(setting) + if not 1 <= max_workers <= sys.maxsize: + raise ValueError("out of range") + except ValueError: + logging.getLogger(__name__).warning("Invalid PYTHON_THREADPOOL_THREAD_COUNT; using the default thread count") + max_workers = None + return ThreadPoolExecutor(max_workers=max_workers, thread_name_prefix="durable-functions") + + +def _executor() -> ThreadPoolExecutor: + runtime = sys.modules.get("azure_functions_runtime") + get_executor = getattr(runtime, "get_threadpool_executor", None) + if callable(get_executor): + executor = get_executor() + if isinstance(executor, ThreadPoolExecutor): + return executor + return _fallback_executor() + + +async def run_sync(function: Callable[_Parameters, _Result], *args: _Parameters.args, + **kwargs: _Parameters.kwargs) -> _Result: + context = _invocation_context.get() + copied_context = copy_context() + + def invoke() -> _Result: + storage = context.thread_local_storage if context is not None else None + previous_id = getattr(storage, "invocation_id", None) + runtime = sys.modules.get("azure_functions_runtime") + invocation_id: ContextVar[str | None] | None = getattr(runtime, "invocation_id_cv", None) + invocation_token = None + if context is not None: + setattr(storage, "invocation_id", context.invocation_id) + if isinstance(invocation_id, ContextVar): + invocation_token = invocation_id.set(context.invocation_id) + try: + return function(*args, **kwargs) + finally: + if invocation_id is not None and invocation_token is not None: + invocation_id.reset(invocation_token) + if storage is not None: + setattr(storage, "invocation_id", previous_id) + + return await asyncio.get_running_loop().run_in_executor(_executor(), copied_context.run, invoke) + + +def wrap_invocation(function: Callable[..., Any], trigger_name: str, + *, registered_trigger_name: str | None = None) -> Callable[..., Any]: + signature = inspect.signature(function) + registered_name = registered_trigger_name or trigger_name + if registered_name == "context": + registered_name = "_durable_input" + while registered_name in signature.parameters: + registered_name = "_" + registered_name + parameters = [ + parameter.replace(name=registered_name) if parameter.name == trigger_name else parameter + for parameter in signature.parameters.values() + ] + inject_context = "context" not in signature.parameters or trigger_name == "context" + if inject_context: + context_parameter = inspect.Parameter( + "context", inspect.Parameter.KEYWORD_ONLY, default=None, annotation=func.Context) + insertion = next((index for index, parameter in enumerate(parameters) + if parameter.kind == inspect.Parameter.VAR_KEYWORD), len(parameters)) + parameters.insert(insertion, context_parameter) + host_signature = signature.replace(parameters=parameters) + + @wraps(function) + async def wrapper(*args: Any, **kwargs: Any) -> Any: + bound = host_signature.bind(*args, **kwargs) + context = bound.arguments.get("context") + if inject_context: + bound.arguments.pop("context", None) + if registered_name != trigger_name: + bound.arguments[trigger_name] = bound.arguments.pop(registered_name) + token = _invocation_context.set(context) + try: + original = inspect.BoundArguments(signature, bound.arguments) + return await function(*original.args, **original.kwargs) + finally: + _invocation_context.reset(token) + + annotations = dict(function.__annotations__) + if registered_name != trigger_name and trigger_name in annotations: + annotations[registered_name] = annotations.pop(trigger_name) + if inject_context: + annotations["context"] = func.Context + wrapper.__annotations__ = annotations + setattr(wrapper, "__signature__", host_signature) + setattr(wrapper, "_df_trigger_name", registered_name) + return wrapper diff --git a/azure-functions-durable/azure/durable_functions/internal/payloads.py b/azure-functions-durable/azure/durable_functions/internal/payloads.py new file mode 100644 index 00000000..46e2ee07 --- /dev/null +++ b/azure-functions-durable/azure/durable_functions/internal/payloads.py @@ -0,0 +1,251 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Process-wide payload storage configured at Function app startup.""" + +import json +from collections.abc import Iterator, Sequence +from dataclasses import dataclass +from itertools import chain +from typing import Any, cast, override +from uuid import UUID + +from google.protobuf.wrappers_pb2 import StringValue + +from durabletask import history +from durabletask.entities import EntityInstanceId +from durabletask.internal.orchestrator_service_pb2 import ( + ActivityRequest, ActivityResponse, HistoryEvent, OrchestratorRequest, +) +from durabletask.payload import ( + LargePayloadStorageOptions, + PayloadStore, + deexternalize_payloads, + deexternalize_payloads_async, + externalize_payloads, + externalize_payloads_async, +) + +_payload_store: PayloadStore | None = None + + +def configure_payload_store(payload_store: object) -> None: + """Register one store per worker process, allowing identical registrations.""" + global _payload_store + if not isinstance(payload_store, PayloadStore): + raise TypeError("payload_store must be a PayloadStore") + if _payload_store is not None and _payload_store is not payload_store: + raise ValueError("A different payload store is already configured in this process") + _payload_store = payload_store + + +def get_payload_store() -> PayloadStore | None: + """Return the app's store, or None when externalization is disabled.""" + return _payload_store + + +class _FunctionsPayloadStore(PayloadStore): + """Keep references valid JSON for the Functions host's payload readers.""" + + def __init__(self, store: PayloadStore) -> None: + self._store = store + + @property + @override + def options(self) -> LargePayloadStorageOptions: + return self._store.options + + @override + def upload(self, data: bytes, *, instance_id: str | None = None) -> str: + return json.dumps(self._store.upload(data, instance_id=instance_id)) + + @override + async def upload_async(self, data: bytes, *, instance_id: str | None = None) -> str: + return json.dumps(await self._store.upload_async(data, instance_id=instance_id)) + + def _unwrap(self, token: str) -> str: + if self._store.is_known_token(token): + return token + if not token.lstrip(" \t\r\n").startswith('"'): + return token + try: + value = json.loads(token) + except (ValueError, TypeError): + return token + return value if isinstance(value, str) else token + + @override + def is_known_token(self, value: str) -> bool: + return self._store.is_known_token(self._unwrap(value)) + + @override + def download(self, token: str) -> bytes: + return self._store.download(self._unwrap(token)) + + @override + async def download_async(self, token: str) -> bytes: + return await self._store.download_async(self._unwrap(token)) + + +def get_transport_payload_store() -> PayloadStore | None: + """Resolve the configured store with Functions-compatible reference encoding.""" + store = get_payload_store() + return _FunctionsPayloadStore(store) if store is not None else None + + +def deexternalize_payload(value: str) -> str: + """Resolve a reference before the Functions JSON codec reads the payload.""" + store = get_transport_payload_store() + if store is None: + return value + request = ActivityRequest(input=StringValue(value=value)) + deexternalize_payloads(request, store) + return request.input.value + + +@dataclass(frozen=True, slots=True) +class ActivityPayload: + value: str + + +async def deexternalize_payload_async(value: str) -> str: + store = get_transport_payload_store() + if store is None: + return value + request = ActivityRequest(input=StringValue(value=value)) + await deexternalize_payloads_async(request, store) + return request.input.value + + +async def externalize_activity_output_async(value: str) -> str: + store = get_transport_payload_store() + if store is None: + return value + response = ActivityResponse(result=StringValue(value=value)) + await externalize_payloads_async(response, store) + return response.result.value + + +def externalize_activity_output(value: str) -> str: + """Apply the core payload policy to a serialized activity output.""" + store = get_transport_payload_store() + if store is None: + return value + response = ActivityResponse(result=StringValue(value=value)) + externalize_payloads(response, store) + return response.result.value + + +def _entity_payload_fields( + events: Sequence[history.HistoryEvent], instance_id: str, *, include_inputs: bool, +) -> Iterator[tuple[history.EventSentEvent | history.EventRaisedEvent, dict[str, Any], str]]: + pending: set[str] = set() + for event in events: + if isinstance(event, history.EntityOperationCalledEvent): + pending.add(event.request_id) + continue + if not isinstance(event, (history.EventSentEvent, history.EventRaisedEvent)): + continue + if isinstance(event, history.EventRaisedEvent): + if event.name not in pending: + continue + pending.remove(event.name) + field = "result" + else: + if event.name != "op" and not event.name.startswith("op@"): + continue + try: + EntityInstanceId.parse(event.instance_id) + except ValueError: + continue + field = "input" + try: + parsed = json.loads(event.input or "") + except ValueError: + continue + if not isinstance(parsed, dict): + continue + envelope = cast(dict[str, Any], parsed) + if isinstance(event, history.EventSentEvent): + if not isinstance(envelope.get("op"), str) or not envelope["op"]: + continue + request_id = envelope.get("id") + if not isinstance(request_id, str): + continue + try: + UUID(request_id) + except ValueError: + continue + if envelope.get("signal", False) is False: + if envelope.get("parent") != instance_id: + continue + pending.add(request_id) + elif envelope.get("signal") is not True: + continue + if (include_inputs or field == "result") and isinstance(envelope.get(field), str): + yield event, envelope, field + + +def hydrate_entity_history( + events: list[history.HistoryEvent], store: PayloadStore | None, instance_id: str, + *, include_inputs: bool = True, +) -> None: + """Hydrate serialized entity protocol fields, never arbitrary object members.""" + if store is None: + return + for event, envelope, field in _entity_payload_fields(events, instance_id, include_inputs=include_inputs): + request = ActivityRequest(input=StringValue(value=envelope[field])) + deexternalize_payloads(request, store) + if request.input.value != envelope[field]: + envelope[field] = request.input.value + event.input = json.dumps(envelope) + + +async def hydrate_entity_history_async( + events: list[history.HistoryEvent], store: PayloadStore | None, instance_id: str, + *, include_inputs: bool = True, +) -> None: + """Hydrate entity history using the store's asynchronous download API.""" + if store is None: + return + for event, envelope, field in _entity_payload_fields(events, instance_id, include_inputs=include_inputs): + request = ActivityRequest(input=StringValue(value=envelope[field])) + await deexternalize_payloads_async(request, store) + if request.input.value != envelope[field]: + envelope[field] = request.input.value + event.input = json.dumps(envelope) + + +def _entity_request_events(request: OrchestratorRequest) -> list[tuple[HistoryEvent, history.HistoryEvent]]: + return [ + (event, history._from_protobuf(event)) # pyright: ignore[reportPrivateUsage] + for event in chain(request.pastEvents, request.newEvents) + if event.WhichOneof("eventType") in ("eventSent", "eventRaised", "entityOperationCalled") + ] + + +def discard_scheduled_activity_inputs(request: OrchestratorRequest) -> None: + """Omit activity inputs that replay never consumes before downloading references.""" + for event in chain(request.pastEvents, request.newEvents): + if event.HasField("taskScheduled"): + event.taskScheduled.ClearField("input") + + +def _update_entity_request_events(events: list[tuple[HistoryEvent, history.HistoryEvent]]) -> None: + for source, event in events: + if isinstance(event, history.EventRaisedEvent) and event.input is not None: + source.eventRaised.input.value = event.input + + +def hydrate_entity_request(request: OrchestratorRequest, store: PayloadStore) -> None: + """Hydrate entity replies for replay, leaving unused historical inputs alone.""" + events = _entity_request_events(request) + hydrate_entity_history([event for _, event in events], store, request.instanceId, include_inputs=False) + _update_entity_request_events(events) + + +async def hydrate_entity_request_async(request: OrchestratorRequest, store: PayloadStore) -> None: + """Await entity reply payloads before replay, without downloading historical inputs.""" + events = _entity_request_events(request) + await hydrate_entity_history_async([event for _, event in events], store, request.instanceId, include_inputs=False) + _update_entity_request_events(events) diff --git a/azure-functions-durable/azure/durable_functions/internal/serialization.py b/azure-functions-durable/azure/durable_functions/internal/serialization.py index 425c434b..b8f296c0 100644 --- a/azure-functions-durable/azure/durable_functions/internal/serialization.py +++ b/azure-functions-durable/azure/durable_functions/internal/serialization.py @@ -12,7 +12,7 @@ from __future__ import annotations from azure.functions._durable_functions import df_dumps, df_loads -from typing import Any +from typing import Any, override from durabletask.serialization import JsonDataConverter @@ -34,21 +34,25 @@ class FunctionsDataConverter(JsonDataConverter): dataclass / ``from_json`` policy. """ + @override def serialize(self, value: Any) -> str | None: if value is None: return None return df_dumps(value) + @override def deserialize(self, data: str | None, target_type: type | None = None) -> Any: if data is None or data == "": return None return df_loads(data, expected_type=target_type) + @override def coerce(self, value: Any, target_type: type | None = None) -> Any: if value is None or target_type is None: return value return self.deserialize(self.serialize(value), target_type) + @override def can_reconstruct(self, target_type: Any) -> bool: return True diff --git a/azure-functions-durable/azure/durable_functions/worker.py b/azure-functions-durable/azure/durable_functions/worker.py index d15ecbf3..3a569d93 100644 --- a/azure-functions-durable/azure/durable_functions/worker.py +++ b/azure-functions-durable/azure/durable_functions/worker.py @@ -7,6 +7,10 @@ from typing import Any, Optional from durabletask import task +from durabletask.payload import ( + deexternalize_payloads, deexternalize_payloads_async, + externalize_payloads, externalize_payloads_async, +) from durabletask.internal.orchestrator_service_pb2 import ( EntityBatchRequest, EntityBatchResult, @@ -16,8 +20,13 @@ ) from durabletask.worker import TaskHubGrpcWorker from .internal.azurefunctions_null_stub import AzureFunctionsNullStub +from .internal.invocation import run_sync from .internal.compat.entity_context import wrap_entity from .internal.compat.orchestration_context import wrap_orchestrator +from .internal.payloads import ( + discard_scheduled_activity_inputs, get_transport_payload_store, + hydrate_entity_request, hydrate_entity_request_async, +) from .internal.serialization import DEFAULT_FUNCTIONS_DATA_CONVERTER _LOGGER = logging.getLogger(__name__) @@ -84,12 +93,37 @@ def _register_entity_once(self, func: task.Entity[Any, Any]) -> None: self._registered_entity_functions[name] = func def execute_orchestration_request(self, func: task.Orchestrator[Any, Any], context: Any) -> str: + payload_store = get_transport_payload_store() context_body = getattr(context, "body", None) if context_body is None: context_body = context orchestration_context = context_body request = OrchestratorRequest() request.ParseFromString(base64.b64decode(orchestration_context)) + if payload_store is not None: + discard_scheduled_activity_inputs(request) + deexternalize_payloads(request, payload_store) + hydrate_entity_request(request, payload_store) + response = self._run_orchestration(func, request) + if payload_store is not None: + externalize_payloads(response, payload_store, instance_id=request.instanceId) + return base64.b64encode(response.SerializeToString()).decode("utf-8") + + async def execute_orchestration_request_async(self, func: task.Orchestrator[Any, Any], context: Any) -> str: + payload_store = get_transport_payload_store() + request = OrchestratorRequest() + request.ParseFromString(base64.b64decode(getattr(context, "body", None) or context)) + if payload_store is not None: + discard_scheduled_activity_inputs(request) + await deexternalize_payloads_async(request, payload_store) + await hydrate_entity_request_async(request, payload_store) + response = await run_sync(self._run_orchestration, func, request) + if payload_store is not None: + await externalize_payloads_async(response, payload_store, instance_id=request.instanceId) + return base64.b64encode(response.SerializeToString()).decode("utf-8") + + def _run_orchestration( + self, func: task.Orchestrator[Any, Any], request: OrchestratorRequest) -> OrchestratorResponse: stub: Any = AzureFunctionsNullStub() response: Optional[OrchestratorResponse] = None @@ -113,17 +147,36 @@ def stub_complete(stub_response: OrchestratorResponse) -> None: if response is None: raise RuntimeError("Orchestrator execution did not produce a response.") - # Return the protobuf response serialized and base64-encoded, the exact - # format the Durable Functions host expects. - return base64.b64encode(response.SerializeToString()).decode("utf-8") + return response def execute_entity_batch_request(self, func: task.Entity[Any, Any], context: Any) -> str: + payload_store = get_transport_payload_store() context_body = getattr(context, "body", None) if context_body is None: context_body = context orchestration_context = context_body request = EntityBatchRequest() request.ParseFromString(base64.b64decode(orchestration_context)) + if payload_store is not None: + deexternalize_payloads(request, payload_store) + response = self._run_entity_batch(func, request) + if payload_store is not None: + externalize_payloads(response, payload_store, instance_id=request.instanceId) + return base64.b64encode(response.SerializeToString()).decode("utf-8") + + async def execute_entity_batch_request_async(self, func: task.Entity[Any, Any], context: Any) -> str: + payload_store = get_transport_payload_store() + request = EntityBatchRequest() + request.ParseFromString(base64.b64decode(getattr(context, "body", None) or context)) + if payload_store is not None: + await deexternalize_payloads_async(request, payload_store) + response = await run_sync(self._run_entity_batch, func, request) + if payload_store is not None: + await externalize_payloads_async(response, payload_store, instance_id=request.instanceId) + return base64.b64encode(response.SerializeToString()).decode("utf-8") + + def _run_entity_batch( + self, func: task.Entity[Any, Any], request: EntityBatchRequest) -> EntityBatchResult: stub: Any = AzureFunctionsNullStub() response: Optional[EntityBatchResult] = None @@ -137,6 +190,4 @@ def stub_complete(stub_response: EntityBatchResult) -> None: if response is None: raise RuntimeError("Entity execution did not produce a response.") - # Return the protobuf response serialized and base64-encoded, the exact - # format the Durable Functions host expects. - return base64.b64encode(response.SerializeToString()).decode("utf-8") + return response diff --git a/durabletask-azuremanaged/CHANGELOG.md b/durabletask-azuremanaged/CHANGELOG.md index 9376577e..18a4c9ec 100644 --- a/durabletask-azuremanaged/CHANGELOG.md +++ b/durabletask-azuremanaged/CHANGELOG.md @@ -7,6 +7,11 @@ adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). ## Unreleased +FIXED + +- With the corresponding core SDK update, asynchronous Azure Blob payload +transfers no longer block the event loop during compression or decompression. + ## v1.10.1 CHANGED diff --git a/durabletask/extensions/azure_blob_payloads/blob_payload_store.py b/durabletask/extensions/azure_blob_payloads/blob_payload_store.py index be3839cf..bde74383 100644 --- a/durabletask/extensions/azure_blob_payloads/blob_payload_store.py +++ b/durabletask/extensions/azure_blob_payloads/blob_payload_store.py @@ -5,6 +5,7 @@ from __future__ import annotations +import asyncio import gzip import logging import uuid @@ -173,7 +174,7 @@ async def upload_async(self, data: bytes, *, instance_id: str | None = None) -> await self._ensure_container_async() if self._options.enable_compression: - data = gzip.compress(data) + data = await asyncio.to_thread(gzip.compress, data) blob_name = self._make_blob_name(instance_id) container_client: AsyncContainerClient = self._async_blob_service_client.get_container_client( @@ -194,7 +195,7 @@ async def download_async(self, token: str) -> bytes: blob_data = await stream.readall() if self._options.enable_compression: - blob_data = gzip.decompress(blob_data) + blob_data = await asyncio.to_thread(gzip.decompress, blob_data) logger.debug("Downloaded %d bytes <- %s", len(blob_data), token) return blob_data diff --git a/noxfile.py b/noxfile.py index 99a3e1e5..2e3ce8a6 100644 --- a/noxfile.py +++ b/noxfile.py @@ -496,8 +496,10 @@ def functions_e2e(session: nox.Session) -> None: session.env[ "AzureFunctionsJobHost__extensions__durableTask__hubName" ] = _new_test_namespace("nox") + session.env["E2E_PAYLOAD_CONTAINER"] = _new_test_namespace("payloads") session.install("-r", "requirements.txt") _install_packages(session, editable=True) + session.install("-e", f"{REPO_ROOT}[azure-blob-payloads]", "aiohttp") session.install("pytest", "opentelemetry-exporter-otlp-proto-grpc") for app in E2E_APPS: _link_app_venv(session, E2E_APPS_DIR / app) diff --git a/tests/azure-functions-durable/conftest.py b/tests/azure-functions-durable/conftest.py new file mode 100644 index 00000000..172a3c15 --- /dev/null +++ b/tests/azure-functions-durable/conftest.py @@ -0,0 +1,47 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Fixtures for the independently runnable Functions provider test suite.""" + +import pytest + +from durabletask.payload import LargePayloadStorageOptions, PayloadStore + + +class FakePayloadStore(PayloadStore): + """In-memory storage with recognizable references and configurable limits.""" + + def __init__(self, threshold_bytes: int = 100, + max_stored_payload_bytes: int = 10 * 1024 * 1024) -> None: + self._options = LargePayloadStorageOptions( + threshold_bytes=threshold_bytes, + max_stored_payload_bytes=max_stored_payload_bytes, + enable_compression=False, + ) + self._blobs: dict[str, bytes] = {} + + @property + def options(self) -> LargePayloadStorageOptions: + return self._options + + def upload(self, data: bytes, *, instance_id: str | None = None) -> str: + token = f"blob:v1:test-container:blob-{len(self._blobs)}" + self._blobs[token] = data + return token + + async def upload_async(self, data: bytes, *, instance_id: str | None = None) -> str: + return self.upload(data, instance_id=instance_id) + + def download(self, token: str) -> bytes: + return self._blobs[token] + + async def download_async(self, token: str) -> bytes: + return self.download(token) + + def is_known_token(self, value: str) -> bool: + return value.startswith("blob:v1:test-container:") + + +@pytest.fixture +def payload_store_factory() -> type[FakePayloadStore]: + return FakePayloadStore diff --git a/tests/azure-functions-durable/e2e/apps/dtask_style/function_app.py b/tests/azure-functions-durable/e2e/apps/dtask_style/function_app.py index 65c6563c..32eb368e 100644 --- a/tests/azure-functions-durable/e2e/apps/dtask_style/function_app.py +++ b/tests/azure-functions-durable/e2e/apps/dtask_style/function_app.py @@ -14,6 +14,8 @@ compatibility layer supports, end-to-end against a real Functions host. """ +import os + import azure.functions as func import azure.durable_functions as df @@ -22,15 +24,22 @@ import client_routes import entities import history_export_routes +import large_payloads import orchestrators +from durabletask.extensions.azure_blob_payloads import BlobPayloadStore, BlobPayloadStoreOptions app = df.DFApp(http_auth_level=func.AuthLevel.ANONYMOUS) +app.configure_large_payloads(payload_store=BlobPayloadStore(BlobPayloadStoreOptions( + connection_string=os.environ.get("AzureWebJobsStorage", "UseDevelopmentStorage=true"), + container_name=os.environ.get("E2E_PAYLOAD_CONTAINER", "functions-e2e-payloads"), +))) app.register_functions(activities.bp) app.register_functions(entities.bp) app.register_functions(orchestrators.bp) app.register_functions(client_routes.bp) app.register_functions(history_export_routes.bp) +app.register_functions(large_payloads.bp) # Opt in to durabletask scheduled tasks: registers the schedule entity and # operation orchestrator so schedules can be managed via ScheduledTaskClient. diff --git a/tests/azure-functions-durable/e2e/apps/dtask_style/large_payloads.py b/tests/azure-functions-durable/e2e/apps/dtask_style/large_payloads.py new file mode 100644 index 00000000..afbf8dec --- /dev/null +++ b/tests/azure-functions-durable/e2e/apps/dtask_style/large_payloads.py @@ -0,0 +1,104 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Blob-backed payload round trips through the real Functions bindings.""" + +import json +from pathlib import Path +from typing import Any + +import azure.functions as func +import azure.durable_functions as df +from durabletask import history, task +from durabletask.entities import EntityInstanceId + +bp = df.Blueprint() + + +@bp.durable_client_input(client_name="client") +@bp.activity_trigger(input_name="payload") +def payload_echo(payload: dict, client: df.SyncDurableFunctionsClient, context: func.Context) -> dict: + assert isinstance(client, df.SyncDurableFunctionsClient) + assert context.thread_local_storage.invocation_id == context.invocation_id + assert json.loads((Path(context.function_directory) / "host.json").read_text())["version"] == "2.0" + return {"data": payload["data"], "stages": [*payload["stages"], "activity"]} + + +@bp.orchestration_trigger(context_name="context") +def payload_roundtrip(ctx: task.OrchestrationContext, payload: dict[str, Any]): + first = yield ctx.call_activity("payload_echo", input=payload) + second = yield ctx.call_activity("payload_echo_async", input=first) + ctx.set_custom_status(second) + return second + + +@bp.activity_trigger(input_name="payload") +@bp.durable_client_input(client_name="client") +async def payload_echo_async(payload: dict, client: df.DurableFunctionsClient, context: func.Context) -> dict: + assert isinstance(client, df.DurableFunctionsClient) + assert json.loads((Path(context.function_directory) / "host.json").read_text())["version"] == "2.0" + return {"data": payload["data"], "stages": [*payload["stages"], "activity"]} + + +@bp.orchestration_trigger(context_name="context") +def payload_entity_roundtrip(ctx: task.OrchestrationContext, payload: dict[str, Any]): + entity_id = EntityInstanceId("probe", ctx.instance_id) + written = yield ctx.call_entity(entity_id, "set", input=payload) + restored = yield ctx.call_entity(entity_id, "get") + assert written == payload, f"Unexpected entity result: {str(written)[:120]}" + assert restored == payload, f"Unexpected entity state: {str(restored)[:120]}" + yield ctx.call_entity(entity_id, "delete") + return restored + + +@bp.orchestration_trigger(context_name="context") +def payload_event_roundtrip(ctx: task.OrchestrationContext, payload: Any): + return (yield ctx.wait_for_external_event("payload")) + + +@bp.orchestration_trigger(context_name="context") +def payload_continue_roundtrip(ctx: task.OrchestrationContext, payload: dict[str, Any]): + if not payload["stages"]: + ctx.continue_as_new({"data": payload["data"], "stages": ["continued"]}) + return + return (yield ctx.call_sub_orchestrator("payload_roundtrip", input=payload)) + + +@bp.route(route="payload-start-sync", methods=["POST"]) +@bp.durable_client_input(client_name="client") +def payload_start_sync( + req: func.HttpRequest, client: df.SyncDurableFunctionsClient) -> func.HttpResponse: + instance_id = client.schedule_new_orchestration("payload_roundtrip", input=req.get_json()) + return func.HttpResponse(json.dumps({"id": instance_id}), status_code=202) + + +@bp.route(route="payload-status-sync/{id}", methods=["GET"]) +@bp.durable_client_input(client_name="client") +def payload_status_sync( + req: func.HttpRequest, client: df.SyncDurableFunctionsClient) -> func.HttpResponse: + state = client.get_orchestration_state(req.route_params["id"], fetch_payloads=True) + assert state is not None + return func.HttpResponse(json.dumps({ + "input": json.loads(state.serialized_input or "null"), + "output": json.loads(state.serialized_output or "null"), + }), mimetype="application/json") + + +def _entity_history_response(events: list[history.HistoryEvent]) -> func.HttpResponse: + inputs = [json.loads(event.input or "null") for event in events + if isinstance(event, (history.EventSentEvent, history.EventRaisedEvent))] + return func.HttpResponse(json.dumps(inputs), mimetype="application/json") + + +@bp.route(route="payload-history-sync/{id}", methods=["GET"]) +@bp.durable_client_input(client_name="client") +def payload_history_sync( + req: func.HttpRequest, client: df.SyncDurableFunctionsClient) -> func.HttpResponse: + return _entity_history_response(client.get_orchestration_history(req.route_params["id"])) + + +@bp.route(route="payload-history-async/{id}", methods=["GET"]) +@bp.durable_client_input(client_name="client") +async def payload_history_async( + req: func.HttpRequest, client: df.DurableFunctionsClient) -> func.HttpResponse: + return _entity_history_response(await client.get_orchestration_history(req.route_params["id"])) diff --git a/tests/azure-functions-durable/e2e/apps/dtask_style/requirements.txt b/tests/azure-functions-durable/e2e/apps/dtask_style/requirements.txt index 20177ffe..be12e5d4 100644 --- a/tests/azure-functions-durable/e2e/apps/dtask_style/requirements.txt +++ b/tests/azure-functions-durable/e2e/apps/dtask_style/requirements.txt @@ -4,3 +4,5 @@ # a real app rather than being pip-installed by the host at start time. azure-functions azure-functions-durable +durabletask[azure-blob-payloads] +aiohttp diff --git a/tests/azure-functions-durable/e2e/test_dtask_large_payloads_e2e.py b/tests/azure-functions-durable/e2e/test_dtask_large_payloads_e2e.py new file mode 100644 index 00000000..e78156c5 --- /dev/null +++ b/tests/azure-functions-durable/e2e/test_dtask_large_payloads_e2e.py @@ -0,0 +1,72 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Large payloads across Functions clients, replay, and activity bindings.""" + +import json +import os + +import pytest + +from ._harness import http_request + +pytestmark = pytest.mark.functions_e2e + + +@pytest.mark.parametrize("size", [32, 300_000]) +@pytest.mark.parametrize("sync_start", [False, True]) +def test_payload_roundtrip(dtask_app, size, sync_start): + payload = {"data": "x" * size, "stages": []} + if sync_start: + response = http_request( + "POST", f"{dtask_app.base_url}/api/payload-start-sync", data=payload) + assert response.status == 202, response.body + instance_id = response.json()["id"] + else: + instance_id = dtask_app.start_orchestration("payload_roundtrip", payload) + + status = dtask_app.wait_for_completion(instance_id) + assert status["runtimeStatus"] == "COMPLETED", status + expected = {"data": payload["data"], "stages": ["activity", "activity"]} + assert status["output"] == expected + assert status["customStatus"] == expected + response = http_request( + "GET", f"{dtask_app.base_url}/api/payload-status-sync/{instance_id}") + assert response.status == 200, response.body + assert response.json() == {"input": payload, "output": expected} + + if size > 262_144: + from azure.storage.blob import BlobServiceClient + + with BlobServiceClient.from_connection_string("UseDevelopmentStorage=true") as storage: + container = storage.get_container_client(os.environ["E2E_PAYLOAD_CONTAINER"]) + blobs = list(container.list_blobs(name_starts_with=f"{instance_id}/")) + assert len(blobs) >= 4 + + +@pytest.mark.parametrize("orchestrator", [ + "payload_entity_roundtrip", "payload_event_roundtrip", "payload_continue_roundtrip", +]) +def test_large_payload_durable_operations(dtask_app, orchestrator): + payload = {"data": "x" * 300_000, "stages": []} + instance_id = dtask_app.start_orchestration(orchestrator, payload) + if orchestrator == "payload_event_roundtrip": + dtask_app.raise_event(instance_id, "payload", payload) + status = dtask_app.wait_for_completion(instance_id) + assert status["runtimeStatus"] == "COMPLETED", status + expected = dict(payload) + if orchestrator == "payload_continue_roundtrip": + expected["stages"] = ["continued", "activity", "activity"] + assert status["output"] == expected + if orchestrator == "payload_entity_roundtrip": + for mode in ("sync", "async"): + response = http_request( + "GET", f"{dtask_app.base_url}/api/payload-history-{mode}/{instance_id}") + assert response.status == 200, response.body + envelopes = response.json() + results = [json.loads(envelope["result"]) for envelope in envelopes + if isinstance(envelope, dict) and envelope.get("result")] + assert results.count(payload) == 2 + inputs = [json.loads(envelope["input"]) for envelope in envelopes + if isinstance(envelope, dict) and envelope.get("op") == "set"] + assert inputs == [payload] diff --git a/tests/azure-functions-durable/test_activity_adapter_compat.py b/tests/azure-functions-durable/test_activity_adapter_compat.py index 5f8b5f04..067042ea 100644 --- a/tests/azure-functions-durable/test_activity_adapter_compat.py +++ b/tests/azure-functions-durable/test_activity_adapter_compat.py @@ -6,12 +6,23 @@ from __future__ import annotations import inspect +import asyncio +import json +import sys +import threading from collections.abc import Mapping +from concurrent.futures import ThreadPoolExecutor +from contextvars import ContextVar +from types import ModuleType, SimpleNamespace from typing import Any +from unittest.mock import AsyncMock, Mock import pytest -from azure.durable_functions.internal.compat.activity import wrap_activity +from azure.functions import meta +from azure.durable_functions.internal import invocation, payloads +from azure.durable_functions.internal.compat.activity import wrap_activity, wrap_activity_payloads +from azure.durable_functions.internal.converters import ActivityTriggerConverter def test_one_param_activity_passes_through_unchanged(): @@ -21,6 +32,34 @@ def act(x): assert wrap_activity(act, "x") is act +@pytest.mark.parametrize("native", [False, True]) +@pytest.mark.parametrize("user_async", [False, True]) +async def test_activity_wrappers_preserve_indexed_source_directory(monkeypatch, native, user_async): + monkeypatch.setattr(payloads, "_payload_store", None) + + def activity(payload): + return payload + + async def async_activity(payload): + return payload + + def native_activity(context, payload): + return payload + + async def async_native_activity(context, payload): + return payload + + original = (async_native_activity if user_async else native_activity) if native else ( + async_activity if user_async else activity) + adapted = wrap_activity(original, "payload") + wrapped = wrap_activity_payloads(adapted, "payload") + assert inspect.getfile(adapted) == inspect.getfile(original) + assert inspect.getfile(wrapped) == inspect.getfile(original) + assert inspect.iscoroutinefunction(wrapped) == user_async + result = await wrapped("value") if user_async else wrapped("value") + assert result == "value" + + def test_two_param_activity_is_adapted_to_single_input(): def act(ctx, payload): return (ctx, payload) @@ -96,3 +135,217 @@ def act(ctx, payload): wrap_activity(act, "not an identifier") with pytest.raises(ValueError, match="valid Python identifier"): wrap_activity(act, "class") # a keyword + + +def test_activity_converters_do_not_access_storage(monkeypatch, payload_store_factory): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + monkeypatch.setattr(store, "download", Mock(side_effect=AssertionError("converter download"))) + monkeypatch.setattr(store, "upload", Mock(side_effect=AssertionError("converter upload"))) + value = json.dumps("blob:v1:test-container:missing") + decoded = ActivityTriggerConverter.decode(meta.Datum(type="json", value=value), trigger_metadata=None) + assert isinstance(decoded, payloads.ActivityPayload) + assert decoded.value == value + encoded = ActivityTriggerConverter.encode({"large": "x" * 200}, expected_type=None) + assert json.loads(encoded.value) == {"large": "x" * 200} + assert ActivityTriggerConverter.encode(decoded, expected_type=None).value == value + + +@pytest.mark.asyncio +@pytest.mark.parametrize("native", [False, True]) +async def test_async_activity_awaits_storage_and_preserves_signature( + monkeypatch, payload_store_factory, native): + store = payload_store_factory() + entered = asyncio.Event() + release = asyncio.Event() + value = {"large": "x" * 200} + token = store.upload(json.dumps(value).encode()) + upload = store.upload + + async def download_async(reference): + entered.set() + await release.wait() + return store._blobs[reference] + + async def upload_async(data, *, instance_id=None): + return upload(data, instance_id=instance_id) + + monkeypatch.setattr(store, "download_async", AsyncMock(side_effect=download_async)) + monkeypatch.setattr(store, "upload_async", AsyncMock(side_effect=upload_async)) + monkeypatch.setattr(store, "download", Mock(side_effect=AssertionError("sync download"))) + monkeypatch.setattr(store, "upload", Mock(side_effect=AssertionError("sync upload"))) + monkeypatch.setattr(payloads, "_payload_store", None) + + async def activity(payload, client): + assert client == "extra binding" + assert payload == value + return payload + + async def native_activity(context, payload): + assert payload == value + return payload + + adapted = wrap_activity(native_activity if native else activity, "payload") + wrapper = wrap_activity_payloads(adapted, "payload") + assert inspect.iscoroutinefunction(wrapper) + assert inspect.signature(wrapper) == inspect.signature(adapted) + monkeypatch.setattr(payloads, "_payload_store", store) + decoded = ActivityTriggerConverter.decode( + meta.Datum(type="json", value=json.dumps(token)), trigger_metadata=None) + kwargs = {} if native else {"client": "extra binding"} + async with asyncio.timeout(5): + pending = asyncio.create_task(wrapper(payload=decoded, **kwargs)) + await entered.wait() + assert not pending.done() + release.set() + result = await pending + encoded = ActivityTriggerConverter.encode(result, expected_type=None) + assert json.loads(store._blobs[json.loads(encoded.value)]) == value + store.download_async.assert_awaited_once() + store.upload_async.assert_awaited_once() + + +@pytest.mark.parametrize("user_async", [False, True]) +@pytest.mark.parametrize("value", [None, "small", {"small": True}]) +async def test_host_activity_externalizes_large_output_from_inline_input( + monkeypatch, payload_store_factory, user_async, value): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + output = {"large": "x" * 200} + + def activity(payload): + assert payload == value + return output + + async def async_activity(payload): + return activity(payload) + + wrapper = wrap_activity_payloads(async_activity if user_async else activity, "payload") + decoded = ActivityTriggerConverter.decode( + meta.Datum(type="json", value=json.dumps(value)), trigger_metadata=None) + result = await wrapper(decoded) if user_async else wrapper(decoded) + encoded = ActivityTriggerConverter.encode(result, expected_type=None) + assert json.loads(store.download(json.loads(encoded.value))) == output + + +def test_sync_activity_uses_sync_storage_on_the_calling_thread(monkeypatch, payload_store_factory): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + value = {"large": "x" * 200} + token = store.upload(json.dumps(value).encode()) + thread_ids = [] + original_download = store.download + original_upload = store.upload + + def download(reference): + thread_ids.append(threading.get_ident()) + return original_download(reference) + + def upload(data, *, instance_id=None): + thread_ids.append(threading.get_ident()) + return original_upload(data, instance_id=instance_id) + + def activity(payload): + thread_ids.append(threading.get_ident()) + assert payload == value + return payload + + monkeypatch.setattr(store, "download", Mock(side_effect=download)) + monkeypatch.setattr(store, "upload", Mock(side_effect=upload)) + monkeypatch.setattr(store, "download_async", AsyncMock(side_effect=AssertionError("async download"))) + monkeypatch.setattr(store, "upload_async", AsyncMock(side_effect=AssertionError("async upload"))) + wrapper = wrap_activity_payloads(activity, "payload") + assert not inspect.iscoroutinefunction(wrapper) + decoded = ActivityTriggerConverter.decode( + meta.Datum(type="string", value=token), trigger_metadata=None) + result = wrapper(decoded) + encoded = ActivityTriggerConverter.encode(result, expected_type=None) + assert json.loads(original_download(json.loads(encoded.value))) == value + assert thread_ids == [threading.get_ident()] * 3 + store.download.assert_called_once_with(token) + store.upload.assert_called_once() + store.download_async.assert_not_called() + store.upload_async.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("use_async", [False, True]) +@pytest.mark.parametrize("operation", ["download", "upload"]) +async def test_activity_storage_errors_propagate_without_retry( + monkeypatch, payload_store_factory, use_async, operation): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + token = store.upload(json.dumps("x" * 200).encode()) + error = OSError("payload storage unavailable") + failure = AsyncMock(side_effect=error) if use_async else Mock(side_effect=error) + monkeypatch.setattr(store, operation + ("_async" if use_async else ""), failure) + called = [] + + def activity(payload): + called.append(True) + return payload + + async def async_activity(payload): + return activity(payload) + + wrapper = wrap_activity_payloads(async_activity if use_async else activity, "payload") + decoded = ActivityTriggerConverter.decode( + meta.Datum(type="string", value=token), trigger_metadata=None) + with pytest.raises(OSError) as raised: + if use_async: + await wrapper(decoded) + else: + wrapper(decoded) + assert raised.value is error + assert called == ([True] if operation == "upload" else []) + if use_async: + failure.assert_awaited_once() + else: + failure.assert_called_once() + + +@pytest.mark.parametrize("setting, expected", [(None, None), ("2", 2), ("invalid", None), ("0", None)]) +def test_sync_executor_honors_functions_thread_count(monkeypatch, setting, expected): + if setting is None: + monkeypatch.delenv("PYTHON_THREADPOOL_THREAD_COUNT", raising=False) + else: + monkeypatch.setenv("PYTHON_THREADPOOL_THREAD_COUNT", setting) + factory = Mock() + monkeypatch.setattr(invocation, "ThreadPoolExecutor", factory) + invocation._fallback_executor.__wrapped__() + factory.assert_called_once_with(max_workers=expected, thread_name_prefix="durable-functions") + + +async def test_sync_execution_reuses_runtime_pool_and_resets_invocation_context(monkeypatch): + runtime = ModuleType("azure_functions_runtime") + host_invocation_id = ContextVar("host_invocation_id", default=None) + setattr(runtime, "invocation_id_cv", host_invocation_id) + monkeypatch.setitem(sys.modules, "azure_functions_runtime", runtime) + storage = threading.local() + loop = asyncio.get_running_loop() + with ThreadPoolExecutor(max_workers=1) as executor: + setattr(runtime, "get_threadpool_executor", lambda: executor) + expected_thread = await loop.run_in_executor(executor, threading.get_ident) + await loop.run_in_executor(executor, setattr, storage, "invocation_id", "previous") + + def execute(payload): + assert threading.get_ident() == expected_thread + assert storage.invocation_id == host_invocation_id.get() == payload + if payload == "failure": + raise ValueError("user failure") + return payload + + async def handle(payload): + return await invocation.run_sync(execute, payload) + + wrapper = invocation.wrap_invocation(handle, "payload") + for identifier in ("first", "failure", "second"): + context = SimpleNamespace(invocation_id=identifier, thread_local_storage=storage) + if identifier == "failure": + with pytest.raises(ValueError, match="user failure"): + await wrapper(identifier, context=context) + else: + assert await wrapper(identifier, context=context) == identifier + assert await loop.run_in_executor(executor, host_invocation_id.get) is None + assert await loop.run_in_executor(executor, getattr, storage, "invocation_id") == "previous" + assert invocation._invocation_context.get() is None diff --git a/tests/azure-functions-durable/test_client_compat.py b/tests/azure-functions-durable/test_client_compat.py index 962ad8eb..8ed5ca20 100644 --- a/tests/azure-functions-durable/test_client_compat.py +++ b/tests/azure-functions-durable/test_client_compat.py @@ -4,7 +4,7 @@ import json from datetime import datetime, timedelta, timezone from types import SimpleNamespace -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock, Mock, patch import azure.functions as func import pytest @@ -22,7 +22,9 @@ from durabletask import history as dt_history, task as dt_task from durabletask.client import AsyncTaskHubGrpcClient, OrchestrationStatus from durabletask.entities import EntityInstanceId +from durabletask.internal import orchestrator_service_pb2 as pb from durabletask.task import RetryPolicy +from azure.durable_functions.internal import payloads _CLIENT_CONFIG = json.dumps({ @@ -117,6 +119,132 @@ def test_durable_clients_use_propagate_only_tracing(): assert init.call_args.kwargs["emit_trace_spans"] is False +@pytest.mark.asyncio +async def test_durable_clients_use_configured_payload_store(monkeypatch, payload_store_factory): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + sync_client = df.SyncDurableFunctionsClient(_CLIENT_CONFIG) + async_client = df.DurableFunctionsClient(_CLIENT_CONFIG) + try: + assert sync_client._payload_store._store is store + assert async_client._payload_store._store is store + finally: + sync_client.close() + await async_client.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("use_async", [False, True]) +@pytest.mark.parametrize("modern_request", [False, True]) +async def test_history_hydrates_only_correlated_entity_envelopes( + monkeypatch, payload_store_factory, use_async, modern_request): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + token = json.dumps(store.upload(b'{"data":"hydrated"}')) + download = store.download + monkeypatch.setattr(store, "download", Mock(wraps=download)) + monkeypatch.setattr(store, "download_async", AsyncMock(side_effect=download)) + request_id = "63c281d7-02d7-412c-9f66-1d6d26a83948" + request = pb.HistoryEvent(eventId=1) + if modern_request: + request.entityOperationCalled.requestId = request_id + request.entityOperationCalled.operation = "get" + request.entityOperationCalled.targetInstanceId.value = "@counter@one" + else: + request.eventSent.instanceId = "@counter@one" + request.eventSent.name = "op" + request.eventSent.input.value = json.dumps({ + "id": request_id, "op": "get", "parent": "instance", "input": token}) + reply = pb.HistoryEvent(eventId=2) + reply.eventRaised.name = request_id + reply.eventRaised.input.value = json.dumps({"result": token}) + ordinary = pb.HistoryEvent(eventId=3) + ordinary.eventRaised.name = "application-event" + ordinary.eventRaised.input.value = reply.eventRaised.input.value + scheduled = pb.HistoryEvent(eventId=4) + scheduled.taskScheduled.name = "echo" + scheduled.taskScheduled.input.value = token + chunks = [pb.HistoryChunk(events=[request]), pb.HistoryChunk(events=[reply, ordinary, scheduled])] + + async def stream(): + for chunk in chunks: + yield chunk + + client = (df.DurableFunctionsClient(_CLIENT_CONFIG) if use_async + else df.SyncDurableFunctionsClient(_CLIENT_CONFIG)) + stub = Mock() + stub.StreamInstanceHistory.return_value = stream() if use_async else iter(chunks) + try: + if use_async: + monkeypatch.setattr(client, "_get_stub", lambda: stub) + events = await client.get_orchestration_history("instance", execution_id="execution") + else: + monkeypatch.setattr(client, "_stub", stub) + events = client.get_orchestration_history("instance", execution_id="execution") + assert json.loads(json.loads(events[1].input)["result"]) == {"data": "hydrated"} + assert events[2].input == ordinary.eventRaised.input.value + assert json.loads(events[3].input) == {"data": "hydrated"} + assert json.loads(reply.eventRaised.input.value)["result"] == token + if not modern_request: + assert json.loads(json.loads(events[0].input)["input"]) == {"data": "hydrated"} + assert stub.StreamInstanceHistory.call_args.args[0].executionId.value == "execution" + expected_downloads = 2 if modern_request else 3 + assert store.download.call_count == (0 if use_async else expected_downloads) + assert store.download_async.await_count == (expected_downloads if use_async else 0) + finally: + if use_async: + await client.close() + else: + client.close() + + +@pytest.mark.parametrize("invalid", ["lock", "parent", "id", "target", "name", "json", "signal"]) +def test_history_preserves_unrelated_or_malformed_envelopes(monkeypatch, payload_store_factory, invalid): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + token = json.dumps("blob:v1:test-container:missing") + request_id = "63c281d7-02d7-412c-9f66-1d6d26a83948" + envelope = {"id": request_id, "op": "get", "parent": "instance", "input": token} + if invalid == "lock": + envelope.pop("op") + envelope["lockset"] = ["@counter@one"] + elif invalid == "parent": + envelope["parent"] = "another-instance" + elif invalid == "id": + envelope["id"] = "not-a-request-id" + elif invalid == "signal": + envelope["signal"] = "false" + request = dt_history.EventSentEvent( + event_id=1, timestamp=datetime.now(timezone.utc), + instance_id="ordinary-instance" if invalid == "target" else "@counter@one", + name="application-event" if invalid == "name" else "op", + input="not-json" if invalid == "json" else json.dumps(envelope)) + reply = dt_history.EventRaisedEvent( + event_id=2, timestamp=request.timestamp, name=request_id, + input=json.dumps({"result": token})) + original = [request.input, reply.input] + payloads.hydrate_entity_history([request, reply], payloads.get_transport_payload_store(), "instance") + assert [request.input, reply.input] == original + + +@pytest.mark.parametrize("name", ["op", "op@2026-09-01T00:00:00Z"]) +def test_history_hydrates_signals_without_correlating_replies(monkeypatch, payload_store_factory, name): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + token = json.dumps(store.upload(b'{"reference":"blob:v1:test-container:literal"}')) + request_id = "63c281d7-02d7-412c-9f66-1d6d26a83948" + request = dt_history.EventSentEvent( + event_id=1, timestamp=datetime.now(timezone.utc), instance_id="@counter@one", name=name, + input=json.dumps({"id": request_id, "op": "set", "signal": True, "input": token})) + reply = dt_history.EventRaisedEvent( + event_id=2, timestamp=request.timestamp, name=request_id, + input=json.dumps({"result": token})) + original_reply = reply.input + payloads.hydrate_entity_history([request, reply], payloads.get_transport_payload_store(), "instance") + assert json.loads(json.loads(request.input)["input"]) == {"reference": "blob:v1:test-container:literal"} + assert reply.input == original_reply + + def test_client_handles_all_config_fields_sent_as_null(): # Newer host extension bundles serialize the full client configuration and # can send any field explicitly as ``null``. Every field must collapse to diff --git a/tests/azure-functions-durable/test_converters.py b/tests/azure-functions-durable/test_converters.py index 69ff1b8d..b9c0e6c5 100644 --- a/tests/azure-functions-durable/test_converters.py +++ b/tests/azure-functions-durable/test_converters.py @@ -12,6 +12,7 @@ import json +import pytest from azure.functions import meta from azure.functions.meta import get_binding_registry @@ -29,6 +30,23 @@ OrchestrationTriggerConverter, register_durable_converters, ) +from azure.durable_functions.internal import payloads +from azure.durable_functions.internal.serialization import FunctionsDataConverter +from azure.durable_functions.internal.compat.activity import wrap_activity_payloads + + +def _encode_activity(value): + host_input = ActivityTriggerConverter.decode( + meta.Datum(type="json", value="null"), trigger_metadata=None) + result = wrap_activity_payloads(lambda payload: value, "payload")(host_input) + return ActivityTriggerConverter.encode(result, expected_type=None) + + +def _decode_activity(datum, **kwargs): + received = [] + wrapper = wrap_activity_payloads(lambda payload: received.append(payload), "payload") + wrapper(ActivityTriggerConverter.decode(datum, trigger_metadata=None)) + return received[0] # --------------------------------------------------------------------------- @@ -104,6 +122,127 @@ def test_activity_trigger_decode_falls_back_to_raw_string(): assert decoded == "not-json" +@pytest.mark.parametrize("data_type", ["string", "json"]) +@pytest.mark.parametrize("value", [{"data": "x" * 200}, "x" * 200, ["x" * 200]]) +def test_activity_trigger_externalizes_and_hydrates(monkeypatch, payload_store_factory, data_type, value): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + encoded = _encode_activity(value) + token = json.loads(encoded.value) + assert store.is_known_token(token) + assert json.loads(store.download(token)) == value + decoded = _decode_activity( + meta.Datum(type=data_type, value=encoded.value), trigger_metadata=None) + assert decoded == value + assert _decode_activity( + meta.Datum(type=data_type, value=token), trigger_metadata=None) == value + + +def test_activity_trigger_keeps_small_payload_inline(monkeypatch, payload_store_factory): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + encoded = _encode_activity({"small": True}) + assert json.loads(encoded.value) == {"small": True} + assert not store._blobs + + +def test_activity_trigger_payload_size_limit(monkeypatch, payload_store_factory): + store = payload_store_factory(threshold_bytes=10, max_stored_payload_bytes=100) + monkeypatch.setattr(payloads, "_payload_store", store) + with pytest.raises(ValueError, match="exceeds the maximum"): + _encode_activity("x" * 200) + assert not store._blobs + + +@pytest.mark.parametrize("reference", [ + "blob:v1:test-container:missing", + json.dumps("blob:v1:test-container:missing"), +]) +def test_activity_trigger_does_not_swallow_download_failure(monkeypatch, payload_store_factory, reference): + monkeypatch.setattr(payloads, "_payload_store", payload_store_factory()) + with pytest.raises(KeyError): + _decode_activity( + meta.Datum(type="string", value=reference), + trigger_metadata=None) + + +def test_whole_payload_token_string_is_reserved(monkeypatch, payload_store_factory): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + stored_value = {"stored": "payload contents"} + reference = store.upload(json.dumps(stored_value).encode()) + + encoded = _encode_activity(reference) + + assert encoded.value == json.dumps(reference) + assert len(store._blobs) == 1 + assert _decode_activity(encoded) == stored_value + assert FunctionsDataConverter().deserialize(payloads.deexternalize_payload(encoded.value)) == stored_value + + +@pytest.mark.parametrize("threshold_bytes", [10, 1024]) +@pytest.mark.parametrize("reference_exists", [False, True]) +def test_object_wrapped_reference_round_trips_as_data( + monkeypatch, payload_store_factory, threshold_bytes, reference_exists): + store = payload_store_factory(threshold_bytes=threshold_bytes) + monkeypatch.setattr(payloads, "_payload_store", store) + reference = "blob:v1:test-container:missing" + if reference_exists: + reference = store.upload(b'{"stored":"not the literal reference"}') + value = {"reference": reference} + + encoded = _encode_activity(value) + + assert len(store._blobs) == int(reference_exists) + int(threshold_bytes == 10) + assert _decode_activity(encoded) == value + assert FunctionsDataConverter().deserialize(payloads.deexternalize_payload(encoded.value)) == value + + +def test_codec_does_not_access_storage(monkeypatch, payload_store_factory): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + value = {"data": "x" * 200} + token = store.upload(json.dumps(value).encode()) + + def unexpected_download(token): + raise AssertionError("Serialization must not access storage") + monkeypatch.setattr(store, "download", unexpected_download) + converter = FunctionsDataConverter() + response = converter.deserialize(json.dumps({"result": json.dumps(token)})) + assert converter.deserialize(response["result"], str) == token + + +@pytest.mark.asyncio +async def test_transport_store_async_references(monkeypatch, payload_store_factory): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + transport = payloads.get_transport_payload_store() + value = b'{"data":"example"}' + token = await transport.upload_async(value, instance_id="test-instance") + assert transport.is_known_token(token) + assert store.is_known_token(json.loads(token)) + assert await transport.download_async(token) == value + assert await transport.download_async(json.loads(token)) == value + assert not transport.is_known_token('{"ordinary":"object"}') + assert not transport.is_known_token('"ordinary string"') + + +def test_reference_detection_skips_non_string_json(monkeypatch, payload_store_factory): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + transport = payloads.get_transport_payload_store() + token = store.upload(b'"value"') + assert transport.is_known_token(" \r\n\t" + json.dumps(token)) + assert transport.is_known_token('"\\u0062lob:v1:test-container:blob-0"') + + def unexpected_parse(value): + raise AssertionError("Non-string payload should not be parsed") + + monkeypatch.setattr(payloads.json, "loads", unexpected_parse) + for value in (' {"result":"example"}', '["example"]', 'null', 'true', '42', token): + assert transport.is_known_token(value) == (value == token) + + # --------------------------------------------------------------------------- # Durable client # --------------------------------------------------------------------------- diff --git a/tests/azure-functions-durable/test_decorator_compat.py b/tests/azure-functions-durable/test_decorator_compat.py index f0f1ab2c..3983e4ef 100644 --- a/tests/azure-functions-durable/test_decorator_compat.py +++ b/tests/azure-functions-durable/test_decorator_compat.py @@ -3,8 +3,12 @@ import azure.durable_functions as df import inspect +import threading +from contextvars import ContextVar +from types import SimpleNamespace import pytest -from unittest.mock import MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, Mock, patch +from azure.durable_functions.internal import payloads from azure.durable_functions.constants import ( ACTIVITY_TRIGGER, DURABLE_CLIENT, @@ -31,7 +35,7 @@ def my_orchestrator(context): context_name="context", orchestration="MyOrchestrator")(my_orchestrator) trigger = _trigger(fb) assert trigger.get_binding_name() == ORCHESTRATION_TRIGGER - assert trigger.name == "context" + assert trigger.name == "_durable_input" assert trigger.orchestration == "MyOrchestrator" @@ -84,9 +88,43 @@ def my_activity(ctx, payload): registered = fb._function._func assert list(inspect.signature(registered).parameters) == ["payload"] assert registered.__name__ == "my_activity" + assert not inspect.iscoroutinefunction(registered) assert registered("hello") == {"echo": "hello"} +@pytest.mark.parametrize("configured", [False, True]) +@pytest.mark.parametrize("user_async", [False, True]) +@pytest.mark.parametrize("value", [None, "hello", {"large": "x" * 200}]) +async def test_direct_activity_calls_skip_storage(monkeypatch, payload_store_factory, configured, user_async, value): + monkeypatch.setattr(payloads, "_payload_store", None) + app = df.DFApp() + + def activity(payload): + return payload + + async def async_activity(payload): + return payload + + registered = app.activity_trigger(input_name="payload")(async_activity if user_async else activity) + store = payload_store_factory() + monkeypatch.setattr(store, "download", Mock(side_effect=AssertionError("inline download"))) + monkeypatch.setattr(store, "upload", Mock(side_effect=AssertionError("inline upload"))) + monkeypatch.setattr(store, "download_async", AsyncMock(side_effect=AssertionError("inline download"))) + monkeypatch.setattr(store, "upload_async", AsyncMock(side_effect=AssertionError("inline upload"))) + if configured: + app.configure_large_payloads(payload_store=store) + function = registered.build().get_user_function() + assert inspect.iscoroutinefunction(function) == user_async + assert list(inspect.signature(function).parameters) == ["payload"] + result = await registered(value) if user_async else registered(value) + assert not inspect.isawaitable(result) + assert result is value + store.download.assert_not_called() + store.upload.assert_not_called() + store.download_async.assert_not_called() + store.upload_async.assert_not_called() + + # --------------------------------------------------------------------------- # entity_trigger # --------------------------------------------------------------------------- @@ -101,11 +139,11 @@ def my_entity(context): context_name="context", entity_name="MyEntity")(my_entity) trigger = _trigger(fb) assert trigger.get_binding_name() == ENTITY_TRIGGER - assert trigger.name == "context" + assert trigger.name == "_durable_input" assert trigger.entity_name == "MyEntity" -def test_entity_trigger_reuses_worker_across_invocations(): +async def test_entity_trigger_reuses_worker_across_invocations(): app = df.DFApp() def my_entity(context): @@ -115,17 +153,17 @@ def my_entity(context): "azure.durable_functions.decorators.durable_app.DurableFunctionsWorker" ) as worker_cls: worker = worker_cls.return_value - worker.execute_entity_batch_request.return_value = "encoded" + worker.execute_entity_batch_request_async = AsyncMock(return_value="encoded") fb = app.entity_trigger(context_name="context")(my_entity) handle = fb._function._func first_context = MagicMock() second_context = MagicMock() - first_result = handle(first_context) - second_result = handle(second_context) + first_result = await handle(first_context) + second_result = await handle(second_context) assert first_result == second_result == "encoded" worker_cls.assert_called_once_with() - assert worker.execute_entity_batch_request.call_args_list == [ + assert worker.execute_entity_batch_request_async.call_args_list == [ ((my_entity, first_context),), ((my_entity, second_context),), ] @@ -210,6 +248,53 @@ async def starter(client: int): assert starter.__annotations__["client"] is int +@pytest.mark.parametrize("client_outer", [False, True]) +@pytest.mark.parametrize("user_async", [False, True]) +async def test_activity_client_binding_preserves_user_type_and_invocation_context(client_outer, user_async): + app = df.DFApp() + host_thread = threading.get_ident() + invocation = SimpleNamespace(invocation_id="test-invocation", thread_local_storage=threading.local()) + invocation.thread_local_storage.invocation_id = invocation.invocation_id + trace = ContextVar("test_trace", default=None) + trace.set("trace-value") + + def activity(payload, client, context): + assert context is invocation + assert context.thread_local_storage.invocation_id == context.invocation_id + assert trace.get() == "trace-value" + assert threading.get_ident() == host_thread + assert isinstance(client, df.DurableFunctionsClient if user_async else df.SyncDurableFunctionsClient) + return payload + + async def async_activity(payload, client, context): + return activity(payload, client, context) + + function = async_activity if user_async else activity + activity_decorator = app.activity_trigger(input_name="payload") + client_decorator = app.durable_client_input(client_name="client") + registered = (client_decorator(activity_decorator(function)) if client_outer + else activity_decorator(client_decorator(function))) + original_filename = inspect.getfile(function) + function = registered.build().get_user_function() + assert inspect.getfile(function) == original_filename + assert inspect.iscoroutinefunction(function) == user_async + result = function(payload="result", client="{}", context=invocation) + assert (await result if user_async else result) == "result" + + +def test_activity_context_trigger_preserves_its_name_and_signature(): + app = df.DFApp() + + def activity(context): + return context + + registered = app.activity_trigger(input_name="context")(activity) + assert _trigger(registered).name == "context" + function = registered.build().get_user_function() + assert list(inspect.signature(function).parameters) == ["context"] + assert function(context="input") == "input" + + # --------------------------------------------------------------------------- # All decorators register a function builder # --------------------------------------------------------------------------- diff --git a/tests/azure-functions-durable/test_worker_compat.py b/tests/azure-functions-durable/test_worker_compat.py index ff82da7d..dbf7d29c 100644 --- a/tests/azure-functions-durable/test_worker_compat.py +++ b/tests/azure-functions-durable/test_worker_compat.py @@ -10,11 +10,14 @@ that path end-to-end without a sidecar or gRPC channel. """ +import asyncio import base64 import json +import threading from concurrent.futures import ThreadPoolExecutor from datetime import datetime from types import SimpleNamespace +from unittest.mock import AsyncMock, Mock import pytest @@ -22,17 +25,376 @@ import durabletask.internal.orchestrator_service_pb2 as pb import azure.durable_functions as df +from azure.durable_functions.internal import invocation, payloads from azure.durable_functions.worker import DurableFunctionsWorker +from durabletask.entities import EntityInstanceId +from durabletask.payload import PayloadStore TEST_INSTANCE_ID = "inst-123" +def test_configure_large_payloads_reaches_workers(monkeypatch): + monkeypatch.setattr(payloads, "_payload_store", None) + assert DurableFunctionsWorker()._payload_store is None + store = Mock(spec=PayloadStore) + app = df.DFApp() + app.configure_large_payloads(payload_store=store) + app.configure_large_payloads(payload_store=store) + assert DurableFunctionsWorker()._payload_store is None + assert payloads.get_transport_payload_store()._store is store + with pytest.raises(ValueError, match="different payload store"): + app.configure_large_payloads(payload_store=Mock(spec=PayloadStore)) + assert payloads.get_payload_store() is store + with pytest.raises(TypeError, match="must be a PayloadStore"): + app.configure_large_payloads(payload_store=None) + + def test_worker_uses_propagate_only_tracing(): worker = DurableFunctionsWorker() assert worker.emit_trace_spans is False +def test_worker_created_before_configuration_hydrates_and_externalizes(monkeypatch, payload_store_factory): + monkeypatch.setattr(payloads, "_payload_store", None) + worker = DurableFunctionsWorker() + store = payload_store_factory() + df.DFApp().configure_large_payloads(payload_store=store) + value = {"data": "x" * 200} + token = store.upload(json.dumps(value).encode()) + + def orchestrator(context): + assert context.get_input() == value + return value + + encoded = _encode_orchestrator_request("payload-orch", encoded_input=json.dumps(token)) + response = _decode_orchestrator_response( + worker.execute_orchestration_request(orchestrator, encoded)) + completion = _get_completion_action(response) + assert completion.orchestrationStatus == pb.ORCHESTRATION_STATUS_COMPLETED + result_token = json.loads(completion.result.value) + assert store.is_known_token(result_token) + assert json.loads(store.download(result_token)) == value + + +@pytest.mark.parametrize("entity", [False, True]) +@pytest.mark.parametrize("storage_failure", [False, True]) +@pytest.mark.parametrize("use_async", [False, True]) +async def test_worker_preserves_output_error(monkeypatch, payload_store_factory, entity, storage_failure, use_async): + store = payload_store_factory(max_stored_payload_bytes=150) + error = OSError("payload storage unavailable") + if storage_failure: + store = payload_store_factory() + monkeypatch.setattr(store, "upload", Mock(side_effect=error)) + monkeypatch.setattr(store, "upload_async", AsyncMock(side_effect=error)) + monkeypatch.setattr(payloads, "_payload_store", store) + + def orchestrator(context): + return "x" * 200 + + def counter(context): + context.set_state("x" * 200) + + with pytest.raises(OSError if storage_failure else ValueError) as raised: + worker = DurableFunctionsWorker() + if entity: + encoded = _encode_entity_batch_request("@counter@key", "set") + if use_async: + await worker.execute_entity_batch_request_async(counter, encoded) + else: + worker.execute_entity_batch_request(counter, encoded) + else: + encoded = _encode_orchestrator_request("oversized") + if use_async: + await worker.execute_orchestration_request_async(orchestrator, encoded) + else: + worker.execute_orchestration_request(orchestrator, encoded) + if storage_failure: + assert raised.value is error + else: + assert "202 bytes" in str(raised.value) + assert "150 bytes" in str(raised.value) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("entity", [False, True]) +@pytest.mark.parametrize("registered", [False, True]) +async def test_async_worker_uses_only_async_storage(monkeypatch, payload_store_factory, entity, registered): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + value = {"data": "x" * 200} + token = store.upload(json.dumps(value).encode()) + download = store.download + monkeypatch.setattr(store, "download_async", AsyncMock(side_effect=download)) + monkeypatch.setattr(store, "upload_async", AsyncMock(side_effect=store.upload)) + monkeypatch.setattr(store, "download", Mock(side_effect=AssertionError("sync download"))) + monkeypatch.setattr(store, "upload", Mock(side_effect=AssertionError("sync upload"))) + host_thread = threading.get_ident() + invocation = SimpleNamespace(invocation_id="worker-id", thread_local_storage=threading.local()) + + def check_execution_thread(): + assert threading.get_ident() != host_thread + if registered: + assert invocation.thread_local_storage.invocation_id == invocation.invocation_id + + def orchestrator(context): + check_execution_thread() + assert context.get_input() == value + return value + + def counter(context): + check_execution_thread() + assert context.get_input() == value + context.set_state(value) + + worker = DurableFunctionsWorker() + if registered: + app = df.DFApp() + builder = (app.entity_trigger(context_name="ctx")(counter) if entity + else app.orchestration_trigger(context_name="ctx")(orchestrator)) + encoded = (_encode_entity_batch_request("@counter@key", "set", json.dumps(token)) if entity + else _encode_orchestrator_request("async-payload", json.dumps(token))) + result = await builder._function._func(ctx=encoded, context=invocation) + if entity: + if not registered: + result = await worker.execute_entity_batch_request_async( + counter, _encode_entity_batch_request("@counter@key", "set", json.dumps(token))) + output = _decode_entity_response(result).entityState.value + else: + if not registered: + result = await worker.execute_orchestration_request_async( + orchestrator, _encode_orchestrator_request("async-payload", json.dumps(token))) + output = _get_completion_action(_decode_orchestrator_response(result)).result.value + assert json.loads(download(json.loads(output))) == value + store.download_async.assert_awaited_once() + store.upload_async.assert_awaited_once() + + +@pytest.mark.parametrize("entity", [False, True]) +async def test_storage_downloads_do_not_wait_for_execution_thread(monkeypatch, payload_store_factory, entity): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + token = json.dumps(store.upload(b'"input"')) + loop = asyncio.get_running_loop() + executing = asyncio.Event() + downloaded = asyncio.Event() + release = threading.Event() + downloads = [] + + async def download(reference): + downloads.append(reference) + if len(downloads) == 2: + downloaded.set() + return store._blobs[reference] + + def execute(context): + assert context.get_input() == "input" + loop.call_soon_threadsafe(executing.set) + assert release.wait(timeout=5) + return "done" + + monkeypatch.setattr(store, "download_async", download) + monkeypatch.setattr(store, "download", Mock(side_effect=AssertionError("sync download"))) + worker = DurableFunctionsWorker() + encoded = (_encode_entity_batch_request("@execute@key", "get", token) if entity + else _encode_orchestrator_request("execute", token)) + invoke = worker.execute_entity_batch_request_async if entity else worker.execute_orchestration_request_async + with ThreadPoolExecutor(max_workers=1) as executor: + monkeypatch.setattr(invocation, "_executor", lambda: executor) + first = asyncio.create_task(invoke(execute, encoded)) + second = None + try: + await asyncio.wait_for(executing.wait(), timeout=5) + second = asyncio.create_task(invoke(execute, encoded)) + await asyncio.wait_for(downloaded.wait(), timeout=5) + assert not first.done() + assert not second.done() + finally: + release.set() + await asyncio.gather(first, *([second] if second is not None else [])) + assert len(downloads) == 2 + + +@pytest.mark.parametrize("entity", [False, True]) +async def test_storage_uploads_do_not_hold_execution_thread(monkeypatch, payload_store_factory, entity): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + uploading = asyncio.Event() + release = asyncio.Event() + upload = store.upload + + async def upload_async(data, *, instance_id=None): + uploading.set() + await release.wait() + return upload(data, instance_id=instance_id) + + def orchestrator(context): + return "x" * 200 + + def counter(context): + context.set_state("x" * 200) + + monkeypatch.setattr(store, "upload_async", upload_async) + monkeypatch.setattr(store, "upload", Mock(side_effect=AssertionError("sync upload"))) + worker = DurableFunctionsWorker() + loop = asyncio.get_running_loop() + with ThreadPoolExecutor(max_workers=1) as executor: + monkeypatch.setattr(invocation, "_executor", lambda: executor) + execution = (worker.execute_entity_batch_request_async( + counter, _encode_entity_batch_request("@counter@key", "set")) if entity + else worker.execute_orchestration_request_async( + orchestrator, _encode_orchestrator_request("upload"))) + pending = asyncio.create_task(execution) + try: + await asyncio.wait_for(uploading.wait(), timeout=5) + assert not pending.done() + assert await asyncio.wait_for( + loop.run_in_executor(executor, lambda: "available"), timeout=5) == "available" + finally: + release.set() + result = await pending + if entity: + value = _decode_entity_response(result).entityState.value + else: + completion = _get_completion_action(_decode_orchestrator_response(result)) + assert completion.orchestrationStatus == pb.ORCHESTRATION_STATUS_COMPLETED + value = completion.result.value + assert json.loads(store.download(json.loads(value))) == "x" * 200 + + +@pytest.mark.parametrize("modern_request", [False, True]) +@pytest.mark.parametrize("storage_failure", [False, True]) +async def test_async_worker_hydrates_entity_envelopes_before_execution( + monkeypatch, payload_store_factory, modern_request, storage_failure): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + value = '{"data":"hydrated"}' + token = json.dumps(store.upload(value.encode())) + error = OSError("nested download failed") + monkeypatch.setattr(store, "download_async", AsyncMock(side_effect=error if storage_failure else store.download)) + monkeypatch.setattr(store, "download", Mock(side_effect=AssertionError("sync download"))) + request = pb.OrchestratorRequest(instanceId=TEST_INSTANCE_ID) + request_id = "63c281d7-02d7-412c-9f66-1d6d26a83948" + sent = request.pastEvents.add(eventId=1) + if modern_request: + sent.entityOperationCalled.requestId = request_id + else: + sent.eventSent.instanceId = "@counter@one" + sent.eventSent.name = "op" + sent.eventSent.input.value = json.dumps({ + "id": request_id, "op": "get", "parent": TEST_INSTANCE_ID, "input": token}) + reply = request.newEvents.add(eventId=2) + reply.eventRaised.name = request_id + reply.eventRaised.input.value = json.dumps({"result": token}) + unrelated = request.newEvents.add(eventId=3) + unrelated.eventRaised.name = "application-event" + unrelated.eventRaised.input.value = reply.eventRaised.input.value + worker = DurableFunctionsWorker() + execution = Mock(return_value=pb.OrchestratorResponse()) + monkeypatch.setattr(worker, "_run_orchestration", execution) + encoded = base64.b64encode(request.SerializeToString()).decode() + if storage_failure: + with pytest.raises(OSError) as raised: + await worker.execute_orchestration_request_async(Mock(), encoded) + assert raised.value is error + execution.assert_not_called() + store.download_async.assert_awaited_once() + else: + await worker.execute_orchestration_request_async(Mock(), encoded) + hydrated = execution.call_args.args[1] + assert json.loads(hydrated.newEvents[0].eventRaised.input.value)["result"] == value + assert hydrated.newEvents[1] == unrelated + if not modern_request: + assert hydrated.pastEvents[0] == sent + store.download_async.assert_awaited_once() + + +@pytest.mark.parametrize("use_async", [False, True]) +async def test_replay_skips_unavailable_historical_entity_input(monkeypatch, payload_store_factory, use_async): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + result_token = store.upload(b'"ok"') + missing_input = json.dumps("blob:v1:test-container:missing") + request_id = "63c281d7-02d7-412c-9f66-1d6d26a83948" + sent = pb.HistoryEvent(eventId=1) + sent.eventSent.instanceId = "@counter@one" + sent.eventSent.name = "op" + sent.eventSent.input.value = json.dumps({ + "id": request_id, "op": "set", "parent": TEST_INSTANCE_ID, "input": missing_input}) + reply = pb.HistoryEvent(eventId=2) + reply.eventRaised.name = request_id + reply.eventRaised.input.value = json.dumps({"result": json.dumps(result_token)}) + request = pb.OrchestratorRequest(instanceId=TEST_INSTANCE_ID) + request.pastEvents.extend([ + helpers.new_orchestrator_started_event(), + helpers.new_execution_started_event("entity-replay", TEST_INSTANCE_ID), + sent, + ]) + request.newEvents.extend([helpers.new_orchestrator_started_event(), reply]) + download = store.download + monkeypatch.setattr(store, "download", Mock(side_effect=download)) + monkeypatch.setattr(store, "download_async", AsyncMock(side_effect=download)) + + def orchestrator(context, value): + return (yield context.call_entity(EntityInstanceId("counter", "one"), "set", input="x" * 200)) + + worker = DurableFunctionsWorker() + encoded = base64.b64encode(request.SerializeToString()).decode() + result = (await worker.execute_orchestration_request_async(orchestrator, encoded) if use_async + else worker.execute_orchestration_request(orchestrator, encoded)) + completion = _get_completion_action(_decode_orchestrator_response(result)) + assert completion.orchestrationStatus == pb.ORCHESTRATION_STATUS_COMPLETED + assert json.loads(completion.result.value) == "ok" + if use_async: + store.download_async.assert_awaited_once_with(result_token) + store.download.assert_not_called() + else: + store.download.assert_called_once_with(result_token) + store.download_async.assert_not_called() + + +@pytest.mark.parametrize("use_async", [False, True]) +@pytest.mark.parametrize("missing_blobs", [False, True]) +async def test_replay_skips_scheduled_inputs(monkeypatch, payload_store_factory, use_async, missing_blobs): + store = payload_store_factory() + monkeypatch.setattr(payloads, "_payload_store", store) + request = pb.OrchestratorRequest(instanceId=TEST_INSTANCE_ID) + request.pastEvents.extend([ + helpers.new_orchestrator_started_event(), + helpers.new_execution_started_event("activity-replay", TEST_INSTANCE_ID), + ]) + for task_id in range(1, 21): + token = store.upload(json.dumps("x" * 300_000).encode()) + events = request.pastEvents if task_id < 20 else request.newEvents + events.extend([ + helpers.new_orchestrator_started_event(), + helpers.new_task_scheduled_event(task_id, "echo", encoded_input=json.dumps(token)), + helpers.new_task_completed_event(task_id, json.dumps("ok")), + ]) + if missing_blobs: + store._blobs.clear() + download = store.download + monkeypatch.setattr(store, "download", Mock(side_effect=download)) + monkeypatch.setattr(store, "download_async", AsyncMock(side_effect=download)) + + def orchestrator(context, value): + for index in range(20): + result = yield context.call_activity("echo", input="x" * 300_000) + assert result == "ok" + return "complete" + + worker = DurableFunctionsWorker() + encoded = base64.b64encode(request.SerializeToString()).decode() + result = (await worker.execute_orchestration_request_async(orchestrator, encoded) if use_async + else worker.execute_orchestration_request(orchestrator, encoded)) + completion = _get_completion_action(_decode_orchestrator_response(result)) + assert completion.orchestrationStatus == pb.ORCHESTRATION_STATUS_COMPLETED + assert json.loads(completion.result.value) == "complete" + store.download.assert_not_called() + store.download_async.assert_not_called() + assert request.pastEvents[3].taskScheduled.HasField("input") + + def _encode_orchestrator_request(name, encoded_input=None, instance_id=TEST_INSTANCE_ID): """Build a base64-encoded ``OrchestratorRequest`` for a single new dispatch.""" request = pb.OrchestratorRequest(instanceId=instance_id) diff --git a/tests/durabletask/test_large_payload.py b/tests/durabletask/test_large_payload.py index 8253151b..dc0f32d5 100644 --- a/tests/durabletask/test_large_payload.py +++ b/tests/durabletask/test_large_payload.py @@ -3,7 +3,10 @@ """Tests for large-payload externalization and de-externalization.""" -from unittest.mock import MagicMock +import asyncio +import gzip +import threading +from unittest.mock import AsyncMock, MagicMock import pytest from google.protobuf import wrappers_pb2 @@ -433,6 +436,72 @@ def test_is_known_token(self): # ------------------------------------------------------------------ +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["upload", "download"]) +@pytest.mark.parametrize("enable_compression", [False, True]) +async def test_async_blob_compression_yields(monkeypatch, operation, enable_compression): + pytest.importorskip("azure.storage.blob") + from durabletask.extensions.azure_blob_payloads import blob_payload_store as module + from durabletask.extensions.azure_blob_payloads import BlobPayloadStoreOptions + + data = b"payload" * 1000 + stored_data = gzip.compress(data) if enable_compression else data + loop = asyncio.get_running_loop() + loop_thread = threading.get_ident() + entered = asyncio.Event() + release = threading.Event() + gzip_method = "compress" if operation == "upload" else "decompress" + original = getattr(gzip, gzip_method) + + def gated_gzip(value): + assert enable_compression + assert threading.get_ident() != loop_thread + loop.call_soon_threadsafe(entered.set) + assert release.wait(5), "Event loop did not release compression" + return original(value) + + async def upload_blob(**kwargs): + assert threading.get_ident() == loop_thread + + async def readall(): + assert threading.get_ident() == loop_thread + return stored_data + + stream = MagicMock() + stream.readall = AsyncMock(side_effect=readall) + container = MagicMock() + container.create_container = AsyncMock() + container.upload_blob = AsyncMock(side_effect=upload_blob) + container.download_blob = AsyncMock(return_value=stream) + service = MagicMock() + service.get_container_client.return_value = container + service.close = AsyncMock() + monkeypatch.setattr(module, "BlobServiceClient", MagicMock()) + monkeypatch.setattr(module, "AsyncBlobServiceClient", MagicMock(return_value=service)) + monkeypatch.setattr(gzip, gzip_method, gated_gzip) + store = module.BlobPayloadStore(BlobPayloadStoreOptions( + account_url="https://example.blob.core.windows.net", enable_compression=enable_compression)) + transfer = asyncio.create_task(store.upload_async(data) if operation == "upload" + else store.download_async("blob:v1:durabletask-payloads:example")) + try: + if enable_compression: + await asyncio.wait_for(entered.wait(), timeout=2) + assert not transfer.done() + finally: + release.set() + try: + result = await transfer + finally: + store.close() + await store.close_async() + if operation == "upload": + uploaded = container.upload_blob.call_args.kwargs["data"] + assert (gzip.decompress(uploaded) if enable_compression else uploaded) == data + assert store.is_known_token(result) + else: + assert result == data + + class TestBlobPayloadStoreDefaults: def test_default_options(self): """Constructing with connection_string should use .NET SDK defaults."""