diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 4683602..134f726 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -6,6 +6,19 @@ on: pull_request: branches: [main] workflow_dispatch: + inputs: + server_commit: + description: Exact public Server candidate SHA, or empty for its default branch + type: string + default: "" + cooperative_qualification: + description: Qualify the candidate cooperative protocol 1.20 + type: boolean + default: false + native_commit: + description: Optional exact public Native source overlay SHA for candidate qualification + type: string + default: "" permissions: contents: read @@ -145,6 +158,11 @@ jobs: runs-on: ubuntu-latest timeout-minutes: 25 needs: [lint, test] + env: + DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION: ${{ inputs.cooperative_qualification && '1.20' || '1.19' }} + DURABLE_WORKFLOW_COOPERATIVE_QUALIFICATION: ${{ inputs.cooperative_qualification && '1' || '0' }} + DURABLE_WORKFLOW_NATIVE_SOURCE: ${{ github.workspace }}/native + COMPOSE_FILE: docker-compose.test.yml steps: - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: @@ -153,7 +171,32 @@ jobs: with: python-version: "3.12" - name: Check out public Server integration source - run: python sdk-python/scripts/ci/checkout-public-repository.py server server + env: + SERVER_COMMIT: ${{ inputs.server_commit }} + run: python sdk-python/scripts/ci/checkout-public-repository.py server server --commit "$SERVER_COMMIT" + - name: Require exact cooperative candidate identities + env: + COOPERATIVE_REQUESTED: ${{ inputs.cooperative_qualification }} + SERVER_COMMIT: ${{ inputs.server_commit }} + NATIVE_COMMIT: ${{ inputs.native_commit }} + run: | + if [ "$COOPERATIVE_REQUESTED" = true ]; then + [[ "$SERVER_COMMIT" =~ ^[0-9a-f]{40}$ ]] + fi + if [ -n "$NATIVE_COMMIT" ]; then + test "$COOPERATIVE_REQUESTED" = true + [[ "$NATIVE_COMMIT" =~ ^[0-9a-f]{40}$ ]] + fi + - name: Check out optional exact public Native source + if: ${{ inputs.native_commit != '' }} + env: + NATIVE_COMMIT: ${{ inputs.native_commit }} + run: python sdk-python/scripts/ci/checkout-public-repository.py workflow native --commit "$NATIVE_COMMIT" + - name: Select readonly Native source qualification + if: ${{ inputs.native_commit != '' }} + run: | + printf 'COMPOSE_FILE=docker-compose.test.yml:docker-compose.native-cancellation.yml\n' >> "$GITHUB_ENV" + printf 'DURABLE_WORKFLOW_NATIVE_SOURCE_QUALIFICATION=1\n' >> "$GITHUB_ENV" - run: pip install -e '.[dev]' working-directory: sdk-python - name: Configure isolated Docker project @@ -162,23 +205,47 @@ jobs: - name: Start Server stack working-directory: sdk-python run: | - docker compose --project-name "$COMPOSE_PROJECT_NAME" -f docker-compose.test.yml \ + docker compose --project-name "$COMPOSE_PROJECT_NAME" \ up -d --build --wait --timeout 300 + - name: Retain exact source and published image authority + working-directory: sdk-python + env: + NATIVE_COMMIT: ${{ inputs.native_commit }} + run: | + docker compose exec -T server cat /app/.package-provenance > integration-package-provenance.txt + docker compose images --format json > integration-images.json + jq --null-input --arg sdk "$GITHUB_SHA" --arg server "$(git -C ../server rev-parse HEAD)" --arg native "$NATIVE_COMMIT" \ + '{qualification:"source",sdk_commit:$sdk,server_commit:$server,native_source_overlay:($native | if length > 0 then . else null end)}' \ + > integration-source-provenance.json - name: Select and probe integration endpoint working-directory: sdk-python run: python scripts/ci/configure-integration-endpoint.py - name: Run integration tests + shell: bash working-directory: sdk-python env: DURABLE_WORKFLOW_AUTH_TOKEN: test-token - run: pytest tests/integration/ -v + run: pytest tests/integration/ -v --capture=tee-sys --junitxml=integration-results.xml 2>&1 | tee integration-scenarios.log + - name: Retain connected scenario results + if: ${{ always() && github.server_url == 'https://github.com' }} + uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7 + with: + name: connected-integration-${{ github.sha }} + path: | + sdk-python/integration-results.xml + sdk-python/integration-scenarios.log + sdk-python/integration-source-provenance.json + sdk-python/integration-package-provenance.txt + sdk-python/integration-images.json + if-no-files-found: warn + retention-days: ${{ inputs.cooperative_qualification && 90 || 7 }} - name: Emit integration diagnostics if: failure() working-directory: sdk-python run: | - for service in bootstrap server worker; do + for service in bootstrap server worker repair; do echo "::group::$service logs" - docker compose --project-name "$COMPOSE_PROJECT_NAME" -f docker-compose.test.yml \ + docker compose --project-name "$COMPOSE_PROJECT_NAME" \ logs --no-color --tail 200 "$service" || true echo "::endgroup::" done @@ -186,7 +253,7 @@ jobs: if: always() working-directory: sdk-python run: | - docker compose --project-name "$COMPOSE_PROJECT_NAME" -f docker-compose.test.yml \ + docker compose --project-name "$COMPOSE_PROJECT_NAME" \ down -v --rmi local target-branch-qualification: diff --git a/README.md b/README.md index 2ce8503..76a4b3f 100644 --- a/README.md +++ b/README.md @@ -127,11 +127,26 @@ pytest tests/ -m "not integration" Integration tests use Docker: ```bash +export COMPOSE_PROJECT_NAME=sdk-python-local docker compose -f docker-compose.test.yml up -d --build --wait -pytest tests/integration/ -v +SERVER_PORT=$(docker compose -f docker-compose.test.yml port server 8080 | sed 's/.*://') +DURABLE_WORKFLOW_SERVER_URL="http://127.0.0.1:$SERVER_PORT" DURABLE_WORKFLOW_AUTH_TOKEN=test-token pytest tests/integration/ -v docker compose -f docker-compose.test.yml down -v ``` +Candidate cooperative cancellation qualification is explicit. In a manual CI +run, supply an exact public `server_commit` and set `cooperative_qualification` +to true. CI verifies that checkout, builds the candidate Server, enables protocol +1.20, runs the connected cases and retains JUnit, raw observations, image +authority and exact source provenance. An optional exact `native_commit` mounts +that public Native checkout read-only into the test stack. The image's published +Composer authority stays intact and the evidence identifies the source overlay. +These are source qualification runs. For a local +candidate, set `DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION=1.20` before starting +Compose and `DURABLE_WORKFLOW_COOPERATIVE_QUALIFICATION=1` for pytest. These cases +fail if the runtime does not discover the required capability. Ordinary CI +keeps protocol 1.19 and skips this unpublished feature's connected cases. + ## License [MIT](LICENSE) diff --git a/docker-compose.native-cancellation.yml b/docker-compose.native-cancellation.yml new file mode 100644 index 0000000..6fffe46 --- /dev/null +++ b/docker-compose.native-cancellation.yml @@ -0,0 +1,14 @@ +# Source qualification only. Published Composer authority stays in the image. +services: + bootstrap: + volumes: + - ${DURABLE_WORKFLOW_NATIVE_SOURCE:?exact Native checkout}:/app/vendor/durable-workflow/workflow:ro + server: + volumes: + - ${DURABLE_WORKFLOW_NATIVE_SOURCE:?exact Native checkout}:/app/vendor/durable-workflow/workflow:ro + worker: + volumes: + - ${DURABLE_WORKFLOW_NATIVE_SOURCE:?exact Native checkout}:/app/vendor/durable-workflow/workflow:ro + repair: + volumes: + - ${DURABLE_WORKFLOW_NATIVE_SOURCE:?exact Native checkout}:/app/vendor/durable-workflow/workflow:ro diff --git a/docker-compose.test.yml b/docker-compose.test.yml index 9a1cc12..c82ed76 100644 --- a/docker-compose.test.yml +++ b/docker-compose.test.yml @@ -30,6 +30,8 @@ services: WORKFLOW_SERVER_AUTH_DRIVER: token WORKFLOW_SERVER_AUTH_TOKEN: "test-token" DW_WORKFLOW_TASK_TIMEOUT: 10 + # Candidate cooperative qualification is explicit. Published defaults remain 1.19. + DW_WORKER_PROTOCOL_VERSION: "${DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION:-1.19}" depends_on: mysql: condition: service_healthy @@ -70,6 +72,16 @@ services: redis: condition: service_healthy + repair: + build: + context: ../server + # Match the production scheduler's unscoped, unthrottled repair cadence. + command: sh -c 'while true; do php artisan workflow:v2:repair-pass --json; sleep 10; done' + environment: *server-env + depends_on: + server: + condition: service_healthy + mysql: image: mysql:8.0 environment: diff --git a/docs/cooperative-cancellation-design.md b/docs/cooperative-cancellation-design.md new file mode 100644 index 0000000..fe013bb --- /dev/null +++ b/docs/cooperative-cancellation-design.md @@ -0,0 +1,286 @@ +# Cooperative cancellation source design + +This describes the unfinished source candidate for shared cancellation issue 136. +The default worker protocol remains 1.19. Cooperation requires explicit protocol +1.20 opt-in and a compatible Server and Native backend. The shared specification +and published mixed-language qualification remain pending. + +## Recorded context + +After committed cancellation delivery, `WorkflowContext.cancellation_context` +provides the original immutable request object. The delivered `WorkflowCancelled` +exception carries the same object in `context`. Earlier workflow code sees +`None`. A poll or heartbeat observation does not expose future cancellation +metadata to earlier workflow execution. + +```python +try: + yield ctx.start_timer(60) +except WorkflowCancelled as cancelled: + request = cancelled.context + with ctx.cancellation_shield(): + yield ctx.schedule_activity("release-reservation", [ + request.root_request_id if request else None, + request.reason if request else None, + ]) + raise +``` + +`CancellationContext` includes local and root request IDs, root workflow instance +and run IDs, the immediate parent request ID, original reason, requester, source, +root request time, original cleanup deadline and ordered lineage. Requester +metadata is limited to caller type, ID and label. It is a read-only mapping and +lineage is a tuple of immutable `CancellationLineage` entries. Timestamp values +are immutable UTC dates. `to_dict()` returns a detached portable snapshot. + +A child preserves the root request time and budget even when its local request +arrives later. Canonical request history supplies the context on cold replay, +while transport observation retains the opaque refresh route. The parser rejects +mismatched local request/run identities, cycles, invalid budgets and a delivery +that changes the accepted snapshot. Older cancellation histories without rich +context continue delivering cancellation with `context is None`. + +Candidate context v2 exposes immutable `scope_origin`, a +`ScopedCancellationContext` with the original root context and every scope +address in order. Its `deadline` is the originating scope's budget, while +`root_deadline` retains the original global deadline. The child context's own +`deadline` may be earlier when parent authority is narrower. The immediate +parent request names the last scope hop, including multiple scopes in one run. +The parser verifies that run lineage derives from the complete tree and refuses +changed metadata, repeated addresses, reentry into an earlier run or a larger +budget. Cold replay and `remaining()` retain the original narrowed authority. +Reading this metadata does not enable scope execution. Both v1 and v2 contexts +remain readable. + +## Child policies + +`CancellationPolicy` and `ParentClosePolicy` are available from the package +root. Child commands accept these enums or their portable string values. + +```python +yield ctx.start_child_workflow( + "python.child", + [], + cancellation_policy=CancellationPolicy.WAIT_CANCELLATION_COMPLETED, + parent_close_policy=ParentClosePolicy.REQUEST_CANCELLATION, +) +``` + +`TRY_CANCEL` requests child cleanup and delivers parent cancellation without +waiting. `WAIT_CANCELLATION_COMPLETED` parks the parent until the child has a +recorded terminal outcome, releasing its task claim for other work. `ABANDON` +leaves the child independent and preserves the historical default. Parent +closure is a separate choice. `REQUEST_CANCELLATION` uses genuine cooperative +cleanup with the original lineage and budget. `REQUEST_CANCEL` retains legacy +terminal behavior. `TERMINATE` and `ABANDON` retain their existing meanings. + +Both command encoders preserve the policies. Cold replay compares them for +ordinary calls, parallel groups, selections and cancellation delivery. Omitted +historical fields mean the original `ABANDON` defaults. Later events with +missing fields keep the scheduled snapshot. Changed options and invalid or +conflicting history fail replay. A worker without the negotiated cooperation +capability refuses these choices before completion, reporting its identity and +the required protocol. Server also checks the immutable task claim and backend. + +## Remote Activity policies + +`ctx.schedule_activity()` accepts `cancellation_policy` as a `CancellationPolicy` +enum or its portable string value. Both command encoders preserve it. Omission +keeps the historical `TRY_CANCEL` behavior and wire shape. + +```python +yield ctx.schedule_activity( + "rust.remote-work", [], + cancellation_policy=CancellationPolicy.WAIT_CANCELLATION_COMPLETED, +) +``` + +`TRY_CANCEL` requests cancellation and continues without waiting for the stop +receipt. `WAIT_CANCELLATION_COMPLETED` delays workflow delivery until the +original remote attempt's physical stop acknowledgment is recorded. `ABANDON` +leaves the Activity independent after parent cancellation. It requires a finite +positive integer `schedule_to_close_timeout`, whose original deadline bounds +the independent work. It does not extend the parent's cleanup budget. + +Every explicit remote policy requires negotiated cooperation, protocol 1.20 +and a compatible installed backend. Missing worker capability is diagnosed +before submission. Server also checks the original immutable claim. Prepared +local Activity policies use their separate admission contract below. + +Replay compares authored policy with original canonical history for ordinary, +parallel and selection calls and cancellation delivery. Later events that omit +the policy retain the original value. Unknown, conflicting and changed policies +fail replay explicitly. Historical histories without a policy retain Try. + +The connected Source tests include explicit Try and Wait with both async and +sync callbacks and no application heartbeats. Wait checks physical stop receipt +ordering before workflow delivery. Bounded Abandon checks that a callback +survives parent closure, completes under its original total lifetime and cannot +publish a second outcome or reopen the cancelled parent. The exact published +mixed-language gate remains required. + +## Remote callback-stop transport + +`Client.acknowledge_activity_cancellation()` reports the original task, activity +attempt, lease owner and cancellation request. It requires explicit worker +protocol 1.20 and validates the Server's original receipt identity. Duplicate +retries retain those identities and share a five-second transport budget. +Refusals preserve the Server diagnostic. This receipt does not renew authority, +record an application heartbeat or extend the cleanup deadline. + +The cooperative remote worker stops and joins its callback and supervisor +before sending this receipt. Each attempt gets an explicit +spawn context, an independent supervisor, a callback process and private payload +files. Single-byte control channels keep owner disconnect independent of a +payload transfer or callback progress. The supervisor joins the callback before +reporting stop. The owner then joins the supervisor. A failed supervisor alone +does not prove a live callback stopped. + +The process tests cover a C call holding the callback interpreter's +GIL, ignored TERM followed by forced stop, actual owner SIGKILL, typed results, +application failure metadata, interceptors and authored heartbeats. The worker +observes ownership independently of callback progress, checks before result +encoding or failure reporting, and reports only the original canonical request +after confirmed stop. Unconfirmed stop retains activity capacity and refuses +successful worker shutdown. A receipt refusal cannot become result publication +or a new cleanup budget. Connected exact-source qualification remains required. + +After application drain expires, cooperative remote shutdown permits up to +15 seconds for process reaping and bounded receipt transport. Application work +is stopped during this phase. The run's original cancellation deadline stays +unchanged. Failure to confirm stop leaves the worker registration active and +raises an explicit shutdown error. + +Handlers, arguments and interceptors must be compatible +with Python's spawn serialization. Define importable handlers and protect the +application entry point with `if __name__ == "__main__"`. Captured memory changes +are local to the callback process. Open process-local connections in the +callback. Registration and remote polling refuse incompatible handler or +interceptor definitions before claiming work, naming the worker and activity. +Prepared local callbacks use the same physical process ownership with separate +durable admission and receipts. Legacy worker protocol 1.19 continues using its existing execution. +Cooperating downstream systems still need idempotency or reconciliation for +effects already performed. + +## Prepared sequential local callbacks + +`ctx.local_activity()` accepts `cancellation_policy=CancellationPolicy.TRY_CANCEL` +or `CancellationPolicy.WAIT_CANCELLATION_COMPLETED`. Omission, including Python's +optional `None`, preserves historical Try behavior and omits the wire field. +Explicit policies require prepared execution and Server discovery of +`prepared_local_activity_cancellation_policies`. The Worker advertises its policy +consumer only for the discovered installed policies. Local `ABANDON` is refused +because a callback owned by this workflow worker cannot outlive that ownership +under the prepared contract. + +```python +yield ctx.local_activity( + "release-reservation", [], + cancellation_policy=CancellationPolicy.WAIT_CANCELLATION_COMPLETED, +) +``` + +Wait parks cancellation delivery until the original local attempt's stop receipt +is recorded. Both supported policies physically stop and join the owned callback +before acknowledgment, without requiring application heartbeats. Neither grants +a new cleanup budget. Replay compares the original policy for completed and +unresolved calls, groups and the committed delivery boundary. Unsupported +policies refuse the entire authored group before callbacks or checkpoint +submission. Queries, updates and validators use the same negotiated replay +consumer. Explicit policies are unavailable through the legacy inline path. + +The source candidate can explicitly request both `cooperative_cancellation` and +`prepared_local_activities` in Worker capabilities. Registration requires source +protocol 1.20 and actual Server discovery of its installed admission bridge. +The manifest advertises `durable_sequential_admission`. Ordinary parallel groups +add explicit `prepared_local_activity_groups` capability and require the Server's +installed atomic admission bridge. Their manifest advertises +`durable_atomic_all_admission`. Selection and turn-closing waits remain refused. + +Replay captures the authored local call and sequence before application code +runs. Earlier side effects, version markers and metadata commands obtain a +retained-claim checkpoint, followed by canonical history refresh. The Server +then creates the local execution and original attempt. The worker validates its +workflow claim, epoch, owner, backend IDs, nonce, fixed deadlines and cleanup +authority before spawning. Native owns retries, backoff and execution timeouts. + +Independent control renews ownership without creating application heartbeat +history or advancing its timeout. Only a real callback heartbeat may do those +things. Requests, retries, payload transfer and result encoding share the +original conservative authority budget. A cancellation fence physically stops +and joins callback and supervisor before the stop receipt. An unconfirmed stop +retains workflow capacity and prevents successful worker deregistration. + +Canonical outcome history supplies the next replay value. A lost or malformed +receipt abandons the claim. Cold replay skips completed callbacks and requests +Native recovery of unfinished Started attempts. Recovery records unknown stop +and may release the claim for a durable retry, without claiming that the +replacement observed the original callback stop. + +Cleanup local calls require a shield after canonical delivery. Admission and +control preserve its original local request ID, root ID, delivery history event +ID and deadline. No replacement or duplicate request grants a fresh cleanup +budget. Connected exact-source qualification remains separate from publication. + +## Prepared local parallel groups + +An ordinary list may contain local activities, remote activities, children, +timers and nested lists, with at most 100 total leaves. The worker checkpoints +the complete authored batch atomically, including every nested position. It +validates all opening history and each canonical local execution identity before +preparing callbacks. Every local admission must validate before any callback +spawns. Earlier metadata and side effects use a separate retained checkpoint. + +Callbacks execute concurrently with independent authority observation. Native +owns outcome history, deadlines, retries and unknown-stop recovery. A retry, +receipt loss or authority loss stops and joins siblings before returning the +claim. Cancellation joins the entire group before workflow delivery or cleanup +replay. A stop acknowledgment proves its own callback physically joined, or +that no callback spawned. An unconfirmed join retains workflow capacity and +worker registration. + +Cold replay preserves completed siblings and recovers only unfinished Started +attempts. A durable retry must release the original claim before new admission. +Results preserve authored nested positions despite settlement order. A cleanup +group requires a shield after canonical delivery and every local member retains +the same original root, delivery event and immutable deadline. + +## Remaining qualification + +### Python deadline and remaining time + +The Source `CancellationContext.deadline` is the original immutable cleanup +deadline. `remaining()` returns fractional seconds left at the replay boundary +consumed by workflow code, clamped to zero. Committed cancellation delivery sets +the initial clock. Blocking activity, prepared local activity, child, timer, +condition, selection and awaited-handle outcomes advance it. A selection uses +its committed winner marker. A group failure excludes later sibling outcomes. +Recorded clock skew cannot increase an already consumed budget. + +Synchronous side effects, version markers, memo updates and inline local +callbacks preserve that clock because they return before their results are +persisted on first execution. Their later history timestamps cannot change +the same authored decision during cold replay. `WorkflowContext.now()` retains +its existing start-time contract. + +Only the active replay that delivered the context can use its remaining-time +clock. Detached metadata and calls after replay ends fail explicitly. Missing +or invalid boundary timestamps also fail, without a host-time fallback. +The runtime supervisor independently enforces the actual deadline and task +ownership even when authoring code cannot run. + +Connected process-loss scenarios record remaining time before SIGKILL and +require the same value and metadata in the replacement worker. Legacy inline +cleanup preserves that value through result persistence. Sequential and atomic +prepared cleanup, each under an original 30-second deadline, check the final +value against the committed completion timestamp and original deadline. + +Rust helpers, explicit local operation policies, nested scopes and competitive +qualification still need completion. + +Connected qualification must cover the PHP parent, Python child, Rust remote +activity and PHP local activity together. Callbacks must stop without application +heartbeats. A replacement after SIGKILL during cleanup must replay the same +boundary and finish before the original 30-second deadline. Record supported +workflow lease, heartbeat and repair settings with that scenario. Exact published +artifacts and one cascade inspection view remain required for release claims. diff --git a/scripts/ci/checkout-public-repository.py b/scripts/ci/checkout-public-repository.py index fe3c681..3b15ff8 100644 --- a/scripts/ci/checkout-public-repository.py +++ b/scripts/ci/checkout-public-repository.py @@ -5,6 +5,7 @@ import argparse import os +import re import subprocess from collections.abc import Sequence from pathlib import Path @@ -12,11 +13,14 @@ PUBLIC_REPOSITORIES = { "cli": "https://github.com/durable-workflow/cli.git", "server": "https://github.com/durable-workflow/server.git", + "workflow": "https://github.com/durable-workflow/workflow.git", } -def checkout(repository: str, destination: Path) -> None: +def checkout(repository: str, destination: Path, commit: str = "") -> None: """Clone a supported public repository without runner-host credentials.""" + if commit and re.fullmatch(r"[0-9a-f]{40}", commit) is None: + raise ValueError("integration source commit must be a full lowercase Git SHA") environment = os.environ.copy() environment["GIT_TERMINAL_PROMPT"] = "0" @@ -34,18 +38,31 @@ def checkout(repository: str, destination: Path) -> None: check=True, env=environment, ) + if commit: + git = ["git", "-c", "credential.helper=", "-C", str(destination)] + subprocess.run([*git, "fetch", "--depth=1", "origin", commit], check=True, env=environment) + subprocess.run([*git, "checkout", "--detach", commit], check=True, env=environment) + resolved = subprocess.run( + [*git, "rev-parse", "HEAD"], check=True, env=environment, capture_output=True, text=True, + ).stdout.strip() + if resolved != commit: + raise RuntimeError("integration checkout did not resolve the requested commit") + print(f"Integration {repository} source commit: {resolved}") def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("repository", choices=sorted(PUBLIC_REPOSITORIES)) parser.add_argument("destination", type=Path) + parser.add_argument( + "--commit", default="", help="Exact public candidate SHA. Defaults to the repository default branch.", + ) return parser.parse_args(argv) def main(argv: Sequence[str] | None = None) -> int: args = parse_args(argv) - checkout(args.repository, args.destination) + checkout(args.repository, args.destination, args.commit) return 0 diff --git a/src/durable_workflow/__init__.py b/src/durable_workflow/__init__.py index adaed75..5dcb0ff 100644 --- a/src/durable_workflow/__init__.py +++ b/src/durable_workflow/__init__.py @@ -16,6 +16,14 @@ AuthCompositionContractError, parse_auth_composition_contract, ) +from .cancellation import ( + CancellationContext, + CancellationLineage, + CancellationPolicy, + ParentClosePolicy, + ScopedCancellationContext, + ScopedCancellationLineage, +) from .client import ( BridgeAdapterOutcome, Client, @@ -237,6 +245,12 @@ "ActivityInterceptorContext", "ActivityRetryPolicy", "BridgeAdapterOutcome", + "CancellationContext", + "CancellationLineage", + "ScopedCancellationContext", + "ScopedCancellationLineage", + "CancellationPolicy", + "ParentClosePolicy", "ChildWorkflowRetryPolicy", "ChildWorkflowCancelled", "ChildWorkflowFailed", diff --git a/src/durable_workflow/_activity_process.py b/src/durable_workflow/_activity_process.py new file mode 100644 index 0000000..63e8d70 --- /dev/null +++ b/src/durable_workflow/_activity_process.py @@ -0,0 +1,322 @@ +"""Owned callback processes for the unfinished cooperative worker. + +The supervisor never unpickles or invokes the serialized application callback. +Single-byte control messages keep owner disconnect observable even if an +application payload is large or its writer is killed. Attempt-local private +files carry payloads. This module does not grant or report Server authority. +""" + +from __future__ import annotations + +import asyncio +import inspect +import multiprocessing +import pickle +import selectors +import shutil +import socket +import tempfile +import threading +import traceback +from collections.abc import Awaitable, Callable +from dataclasses import dataclass +from multiprocessing.process import BaseProcess +from pathlib import Path +from typing import Any, cast + +from .activity import ActivityContext, ActivityInfo, _set_context +from .errors import ActivityCancelled, NonRetryableError +from .interceptors import ActivityInterceptorContext, WorkerInterceptor + + +class CallbackProcessLost(RuntimeError): + """The owner could not prove the callback outcome or physical stop.""" + + +@dataclass(frozen=True) +class CallbackInvocation: + handler: Callable[..., Any] + args: tuple[Any, ...] + info: ActivityInfo + task: dict[str, Any] + interceptors: tuple[WorkerInterceptor, ...] = () + + def serialize(self) -> bytes: + try: + return pickle.dumps(self, protocol=pickle.HIGHEST_PROTOCOL) + except Exception as error: + raise ValueError( + f"cooperative activity {self.info.activity_type!r} requires a spawn-compatible " + "importable handler, arguments and interceptors" + ) from error + + +@dataclass(frozen=True) +class CallbackFailure: + message: str + failure_type: str + failure_class: str + failure_code: int | None + stack_trace: str + non_retryable: bool + cancelled: bool = False + + +@dataclass(frozen=True) +class CallbackOutcome: + value: Any = None + failure: CallbackFailure | None = None + + +def _write_payload(directory: str, name: str, payload: Any) -> None: + destination = Path(directory, name) + temporary = destination.with_suffix(".writing") + with temporary.open("wb") as stream: + pickle.dump(payload, stream, protocol=pickle.HIGHEST_PROTOCOL) + temporary.replace(destination) + + +def _read_payload(directory: str, name: str) -> Any: + with Path(directory, name).open("rb") as stream: + return pickle.load(stream) + + +def _callback(directory: str, channel: socket.socket) -> None: + # Deserialize only in the process that is permitted to run application code. + heartbeat_lock = threading.Lock() + + def heartbeat_request(details: dict[str, Any] | None) -> None: + with heartbeat_lock: + _write_payload(directory, "heartbeat", details) + channel.sendall(b"H") + if channel.recv(1) != b"Y": + raise ActivityCancelled() + + async def heartbeat(details: dict[str, Any] | None) -> None: + await asyncio.to_thread(heartbeat_request, details) + + async def invoke(invocation: CallbackInvocation) -> Any: + _set_context(ActivityContext(info=invocation.info, client=cast(Any, None), heartbeat_callback=heartbeat)) + context = ActivityInterceptorContext( + worker_id=invocation.info.worker_id, task_queue=invocation.info.task_queue, + task=invocation.task, activity_type=invocation.info.activity_type, args=invocation.args, + ) + + async def call(ctx: ActivityInterceptorContext) -> Any: + value = (invocation.handler(*ctx.args) if inspect.iscoroutinefunction(invocation.handler) + else await asyncio.to_thread(invocation.handler, *ctx.args)) + return await value if inspect.isawaitable(value) else value + + handler = call + for interceptor in reversed(invocation.interceptors): + next_handler = handler + + async def intercepted( + ctx: ActivityInterceptorContext, *, interceptor: WorkerInterceptor = interceptor, + next_handler: Callable[[ActivityInterceptorContext], Awaitable[Any]] = next_handler, + ) -> Any: + return await interceptor.execute_activity(ctx, next_handler) + + handler = intercepted + try: + return await handler(context) + finally: + _set_context(None) + + try: + invocation = pickle.loads(Path(directory, "invocation").read_bytes()) + if not isinstance(invocation, CallbackInvocation): + raise TypeError("invalid callback invocation") + channel.sendall(b"B") + try: + _write_payload(directory, "outcome", CallbackOutcome(value=asyncio.run(invoke(invocation)))) + except BaseException as error: + code = getattr(error, "code", None) + _write_payload(directory, "outcome", CallbackOutcome(failure=CallbackFailure( + message=str(error), failure_type=type(error).__name__, + failure_class=f"{type(error).__module__}.{type(error).__qualname__}", + failure_code=code if isinstance(code, int) and not isinstance(code, bool) else None, + stack_trace=traceback.format_exc(), non_retryable=isinstance(error, NonRetryableError), + cancelled=isinstance(error, ActivityCancelled), + ))) + finally: + channel.close() + + +def _stop_owned_callback(callback: BaseProcess) -> None: + if callback.is_alive(): + callback.terminate() + callback.join(timeout=0.5) + if callback.is_alive(): + callback.kill() + # If the OS cannot reap it yet, retain the supervisor and ownership. Never + # send stop evidence before an actual join, even after the owner disappears. + callback.join() + + +def _supervise(directory: str, owner: socket.socket) -> None: + context = multiprocessing.get_context("spawn") + supervisor_channel, callback_channel = socket.socketpair() + callback = context.Process(target=_callback, args=(directory, callback_channel)) + callback_started = False + stopped = False + try: + callback.start() + callback_started = True + callback_channel.close() + with selectors.DefaultSelector() as selector: + selector.register(owner, selectors.EVENT_READ, "owner") + selector.register(supervisor_channel, selectors.EVENT_READ, "callback") + while True: + for key, _ in selector.select(timeout=0.05): + message = cast(socket.socket, key.fileobj).recv(1) + if key.data == "owner": + if message in (b"", b"F"): + return + if message == b"S": + _stop_owned_callback(callback) + stopped = True + owner.sendall(b"X") + elif message in (b"Y", b"N") and not stopped: + supervisor_channel.sendall(message) + elif message in (b"B", b"H"): + owner.sendall(message) + elif message == b"": + selector.unregister(supervisor_channel) + if not stopped and callback.exitcode is not None: + callback.join() + stopped = True + owner.sendall(b"D") + except (BrokenPipeError, ConnectionResetError): + pass # Owner disappeared. The finally block still owns and reaps its callback. + finally: + if callback_started: + _stop_owned_callback(callback) + callback.close() + callback_channel.close() + supervisor_channel.close() + owner.close() + shutil.rmtree(directory, ignore_errors=True) + + +class SupervisedCallback: + """One owner, one supervisor and one isolated application callback. + + The owner must serialize access to result()/stop(). Cancellation of result() + leaves the process owned until stop() proves shutdown. A killed supervisor + never supplies stop evidence. The worker integration remains a separate gate. + """ + + def __init__(self, invocation: CallbackInvocation) -> None: + serialized = invocation.serialize() + self.directory = tempfile.mkdtemp(prefix="dw-activity-") + try: + Path(self.directory, "invocation").write_bytes(serialized) + except BaseException: + shutil.rmtree(self.directory, ignore_errors=True) + raise + self._owner, self._supervisor_channel = socket.socketpair() + self._owner.setblocking(False) + self._supervisor = multiprocessing.get_context("spawn").Process( + target=_supervise, args=(self.directory, self._supervisor_channel), + ) + self._started = False + self._closed = False + self._joined = False + self._exitcode: int | None = None + self.stopped = False + + async def start(self) -> None: + try: + # Pass only the private directory and control socket to spawn. The + # payload is already serialized, and this short start cannot leave + # a background spawn thread untracked if the owner is cancelled. + self._supervisor.start() + self._started = True + self._supervisor_channel.close() + message = await asyncio.wait_for(self._receive(), timeout=5.0) + if message != b"B": + raise CallbackProcessLost("callback did not prove spawn readiness") + except BaseException: + await self.close() + raise + + async def _receive(self) -> bytes: + try: + message = await asyncio.get_running_loop().sock_recv(self._owner, 1) + except (ConnectionResetError, OSError) as error: + raise CallbackProcessLost("callback supervisor disconnected") from error + if not message: + raise CallbackProcessLost("callback supervisor disconnected without stop evidence") + return message + + async def _send(self, message: bytes) -> None: + await asyncio.get_running_loop().sock_sendall(self._owner, message) + + async def result(self, heartbeat: Callable[[dict[str, Any] | None], Awaitable[None]]) -> CallbackOutcome: + while True: + message = await self._receive() + if message == b"H": + await heartbeat(_read_payload(self.directory, "heartbeat")) + await self._send(b"Y") + elif message == b"D": + outcome = _read_payload(self.directory, "outcome") + if not isinstance(outcome, CallbackOutcome): + raise CallbackProcessLost("callback did not produce a typed outcome") + await self._finish() + return outcome + else: + raise CallbackProcessLost("callback supervisor returned an unexpected result boundary") + + async def stop(self) -> None: + if self.stopped: + return + async def wait_for_stop() -> None: + while await self._receive() not in (b"D", b"X"): + pass + await self._finish() + + try: + await self._send(b"S") + await asyncio.wait_for(wait_for_stop(), timeout=5.0) + except BaseException: + await self.close() + raise + + async def _finish(self) -> None: + # Once stop evidence arrives, cancellation cannot interrupt the OS + # join and leave a physically exited callback falsely unconfirmed. + join = asyncio.create_task(self._confirm_join()) + cancelled = False + while not join.done(): + try: + await asyncio.shield(join) + except asyncio.CancelledError: + cancelled = True + join.result() + if cancelled: + raise asyncio.CancelledError + + async def _confirm_join(self) -> None: + await self._send(b"F") + await self.close() + if self._exitcode != 0: + raise CallbackProcessLost("callback supervisor did not confirm a clean join") + self.stopped = True + + async def close(self) -> None: + if self._joined: + return + if not self._closed: + self._closed = True + self._owner.close() + self._supervisor_channel.close() + if self._started: + await asyncio.to_thread(self._supervisor.join, 5.0) + if self._supervisor.is_alive(): + raise CallbackProcessLost("callback supervisor has not confirmed shutdown") + self._exitcode = self._supervisor.exitcode + else: + shutil.rmtree(self.directory, ignore_errors=True) + self._supervisor.close() + self._joined = True diff --git a/src/durable_workflow/_cooperative_cancellation.py b/src/durable_workflow/_cooperative_cancellation.py new file mode 100644 index 0000000..332f5cc --- /dev/null +++ b/src/durable_workflow/_cooperative_cancellation.py @@ -0,0 +1,279 @@ +"""Canonical cancellation state used by service workflow replay. + +An observation supplies request identity. Only the durable delivery event +authorizes an exception at an authored workflow call. +""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from datetime import datetime +from typing import Any + +from .cancellation import CancellationContext +from .errors import NonDeterministicReplayError + +REQUEST_EVENT = "CooperativeCancellationRequested" +DELIVERY_EVENT = "CooperativeCancellationDelivered" +_CALL_KINDS = {"activity", "local_activity", "timer", "condition", "signal", "child", "parallel", "selection_handle"} +_RESOLUTION_EVENTS = { + "ActivityCompleted", + "ActivityFailed", + "ActivityCancelled", + "ActivityTimedOut", + "TimerFired", + "TimerCancelled", + "ConditionWaitSatisfied", + "ConditionWaitTimedOut", + "SignalApplied", + "ChildRunCompleted", + "ChildRunFailed", + "ChildRunCancelled", + "ChildRunTerminated", +} + + +def _invalid(detail: str, sequence: int = 0) -> NonDeterministicReplayError: + return NonDeterministicReplayError( + sequence, + "canonical cooperative cancellation", + [REQUEST_EVENT, DELIVERY_EVENT], + detail=detail, + ) + + +def _text(value: Any, field: str) -> str: + if not isinstance(value, str) or not value.strip(): + raise _invalid(f"{field} must be a non-empty string") + return value + + +def _positive(value: Any, field: str, *, maximum: int = 2**63 - 1) -> int: + if type(value) is not int or not 1 <= value <= maximum: + raise _invalid(f"{field} must be a positive integer within {maximum}") + return value + + +def _timestamp(value: Any, field: str) -> str: + text = _text(value, field) + try: + parsed = datetime.fromisoformat(text.replace("Z", "+00:00")) + except ValueError as exc: + raise _invalid(f"{field} must be an ISO timestamp") from exc + if parsed.tzinfo is None: + raise _invalid(f"{field} must include its timezone") + return text + + +@dataclass(frozen=True) +class CancellationRequest: + request_id: str + requested_at: str + cleanup_deadline_at: str + history_refresh_page_token: str | None = None + context: CancellationContext | None = None + + @classmethod + def from_observation(cls, value: Mapping[str, Any]) -> CancellationRequest: + request = cls( + _text(value.get("request_id"), "request_id"), + _timestamp(value.get("requested_at"), "requested_at"), + _timestamp(value.get("cleanup_deadline_at"), "cleanup_deadline_at"), + _text(value.get("history_refresh_page_token"), "history_refresh_page_token"), + ) + request._validate_deadline() + return request + + def _validate_deadline(self) -> None: + requested = datetime.fromisoformat(self.requested_at.replace("Z", "+00:00")) + deadline = datetime.fromisoformat(self.cleanup_deadline_at.replace("Z", "+00:00")) + if deadline <= requested: + raise _invalid("cleanup deadline must follow the original request") + + +@dataclass(frozen=True) +class CancellationDelivery: + request_id: str + sequence: int + call_kind: str + sequence_span: int = 1 + operation_sequence: int | None = None + operation_sequence_span: int = 1 + + @classmethod + def from_payload(cls, payload: Mapping[str, Any]) -> CancellationDelivery: + request_id = _text(payload.get("workflow_command_id"), "workflow_command_id") + sequence = _positive(payload.get("sequence"), "sequence") + kind = _text(payload.get("call_kind"), "call_kind") + span = _positive(payload.get("sequence_span", 1), "sequence_span", maximum=1000) + operation = payload.get("operation_sequence") + operation_span = _positive(payload.get("operation_sequence_span", 1), "operation_sequence_span", maximum=1000) + if kind not in _CALL_KINDS or (kind != "parallel" and span != 1): + raise _invalid("delivery kind and span do not describe a durable call", sequence) + if sequence > 2**63 - 1 - span: + raise _invalid("delivery sequence range overflows", sequence) + if kind == "selection_handle": + operation = _positive(operation, "operation_sequence") + if operation >= sequence or operation_span > sequence - operation: + raise _invalid("selection handle must name an earlier authored operation", sequence) + elif operation is not None or operation_span != 1: + raise _invalid("only selection handles may carry an operation range", sequence) + return cls(request_id, sequence, kind, span, operation, operation_span) + + def interrupts(self, sequence: int) -> bool: + base = self.operation_sequence if self.operation_sequence is not None else self.sequence + span = self.operation_sequence_span if self.operation_sequence is not None else self.sequence_span + return base <= sequence < base + span + + +@dataclass(frozen=True) +class CancellationHistory: + request: CancellationRequest | None + delivery: CancellationDelivery | None + request_index: int + delivery_index: int | None + resolved_before_request: frozenset[int] + failed_before_request: frozenset[int] + selected_before_request: frozenset[tuple[int, int]] + + def eligible(self, sequence: int, span: int = 1) -> bool: + if self.request is None or self.delivery is not None: + return False + sequences = set(range(sequence, sequence + span)) + return ( + not sequences <= self.resolved_before_request + and not sequences & self.failed_before_request + and (sequence, span) not in self.selected_before_request + ) + + +def read_cancellation_history( + events: Sequence[dict[str, Any]], + *, + run_id: str = "", + observation: Mapping[str, Any] | None = None, +) -> CancellationHistory: + request = CancellationRequest.from_observation(observation) if observation is not None else None + request_index = len(events) + delivery: CancellationDelivery | None = None + delivery_index: int | None = None + saw_request = False + for index, event in enumerate(events): + kind = event.get("event_type") or event.get("type") + if kind not in {REQUEST_EVENT, DELIVERY_EVENT}: + continue + payload = event.get("payload") + if not isinstance(payload, Mapping): + raise _invalid("canonical event is missing its payload") + request_id = _text(payload.get("workflow_command_id"), "workflow_command_id") + if event.get("workflow_command_id", request_id) != request_id: + raise _invalid("event and payload disagree on the original request identity") + event_run = _text(payload.get("workflow_run_id"), "workflow_run_id") + if run_id and event_run != run_id: + raise _invalid("cancellation belongs to a different workflow run") + if kind == REQUEST_EVENT: + if saw_request or delivery is not None: + raise _invalid("history must contain one request before delivery") + context = None + if "cancellation" in payload: + snapshot = payload["cancellation"] + if not isinstance(snapshot, Mapping): + raise _invalid("canonical cancellation context must be an object") + try: + context = CancellationContext.from_dict(snapshot) + except ValueError as error: + raise _invalid(str(error)) from error + local = context.lineage[-1] + recorded_at = _timestamp(event.get("recorded_at", event.get("timestamp")), "recorded_at") + if ( + context.request_id != request_id or local.workflow_run_id != event_run + or ("workflow_instance_id" in payload + and local.workflow_instance_id != payload["workflow_instance_id"]) + or context.deadline != datetime.fromisoformat( + _timestamp(payload.get("cleanup_deadline_at"), "cleanup_deadline_at").replace("Z", "+00:00"), + ) + or context.requested_at > datetime.fromisoformat(recorded_at.replace("Z", "+00:00")) + or "reason" in payload and context.reason != payload["reason"] + ): + raise _invalid("canonical cancellation context does not match its request event") + canonical_request = CancellationRequest( + request_id, + context.requested_at.isoformat(timespec="microseconds").replace("+00:00", "Z") + if context is not None + else request.requested_at + if request is not None + else _timestamp( + event.get("recorded_at", event.get("timestamp")), + "requested_at", + ), + _timestamp(payload.get("cleanup_deadline_at"), "cleanup_deadline_at"), + request.history_refresh_page_token if request is not None else None, + context, + ) + if request is not None and ( + request.request_id != canonical_request.request_id + or datetime.fromisoformat(request.cleanup_deadline_at.replace("Z", "+00:00")) + != datetime.fromisoformat(canonical_request.cleanup_deadline_at.replace("Z", "+00:00")) + ): + raise _invalid("observation changes the original request or cleanup deadline") + request = canonical_request + request_index = index + saw_request = True + else: + if not saw_request or request is None or delivery is not None: + raise _invalid("delivery requires one earlier canonical request and one marker") + delivery = CancellationDelivery.from_payload(payload) + if delivery.request_id != request.request_id: + raise _invalid("delivery names a different request", delivery.sequence) + if "cancellation" in payload: + snapshot = payload["cancellation"] + if not isinstance(snapshot, Mapping) or request.context is None: + raise _invalid("delivery changes the canonical cancellation context", delivery.sequence) + try: + delivered_context = CancellationContext.from_dict(snapshot) + except ValueError as error: + raise _invalid(str(error), delivery.sequence) from error + if delivered_context != request.context: + raise _invalid("delivery changes the canonical cancellation context", delivery.sequence) + delivery_index = index + resolved: set[int] = set() + failed: set[int] = set() + selected: set[tuple[int, int]] = set() + for event in events[:request_index]: + kind = event.get("event_type") or event.get("type") + payload = event.get("payload") or {} + if not isinstance(payload, Mapping): + continue + sequence = payload.get("sequence") + if kind == "SelectionResolved": + base = payload.get("selection_group_base_sequence") + span = payload.get("selection_group_size") + if type(base) is int and base > 0 and type(span) is int and 1 <= span <= 1000: + selected.add((base, span)) + elif kind == "SelectionOperationCancelled": + base = payload.get("member_base_sequence") + span = payload.get("member_size") + if type(base) is int and base > 0 and type(span) is int and 1 <= span <= 1000: + resolved.update(range(base, base + span)) + if type(sequence) is int and sequence > 0 and kind in _RESOLUTION_EVENTS: + resolved.add(sequence) + if kind in { + "ActivityFailed", + "ActivityTimedOut", + "ActivityCancelled", + "ChildRunFailed", + "ChildRunCancelled", + "ChildRunTerminated", + }: + failed.add(sequence) + state = CancellationHistory( + request, delivery, request_index, delivery_index, frozenset(resolved), frozenset(failed), frozenset(selected), + ) + if delivery is not None: + base = delivery.operation_sequence or delivery.sequence + span = delivery.operation_sequence_span if delivery.operation_sequence is not None else delivery.sequence_span + interrupted = set(range(base, base + span)) + if interrupted <= resolved or interrupted & failed or (base, span) in selected: + raise _invalid("delivery cannot replace an earlier committed result", delivery.sequence) + return state diff --git a/src/durable_workflow/_prepared_local_activity.py b/src/durable_workflow/_prepared_local_activity.py new file mode 100644 index 0000000..647f051 --- /dev/null +++ b/src/durable_workflow/_prepared_local_activity.py @@ -0,0 +1,352 @@ +"""Source-only prepared local admission and supervised callback execution. + +Native owns retry policy, execution deadlines and terminal history. Control +observes authority independently of application heartbeats. A lost receipt +abandons the claim and never authorizes another callback or publication. +""" + +from __future__ import annotations + +import asyncio +import contextlib +import re +import time +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from datetime import datetime, timedelta +from typing import Any + +from . import serializer +from ._activity_process import CallbackInvocation, SupervisedCallback +from .cancellation import CancellationContext +from .client import Client +from .workflow import LocalActivityExecutionAborted + +_FIXED_DEADLINES = ("start_to_close_deadline_at", "schedule_to_close_deadline_at") + + +def _require(condition: bool, message: str) -> None: + if not condition: + raise LocalActivityExecutionAborted(message) + + +def _text(value: Any) -> str: + _require(isinstance(value, str) and bool(value.strip()), "prepared receipt lacks a nonempty identity") + return str(value) + + +def _timestamp(value: Any) -> datetime: + _require(isinstance(value, str) and re.fullmatch( + r"\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}(?:\.\d+)?(?:Z|[+-]\d{2}:\d{2})", value, + ) is not None, "prepared receipt requires an ISO timestamp with timezone") + try: + return datetime.fromisoformat(str(value).replace("Z", "+00:00")) + except ValueError as error: + raise LocalActivityExecutionAborted("prepared receipt timestamp is invalid") from error + + +def _same_deadline(actual: Any, expected: Any) -> bool: + return actual is None if expected is None else _timestamp(actual) == _timestamp(expected) + + +class PreparedCancellationObserved(LocalActivityExecutionAborted): + """Cancellation was fenced and its original callback physically joined.""" + + +@dataclass +class PreparedAttempt: + task_id: str + run_id: str + owner: str + epoch: int + nonce: str + execution_id: str + attempt_id: str + attempt_number: int + deadlines: dict[str, Any] + heartbeat_timeout: int | None + cleanup: Mapping[str, str] | None + clock_origin: float + authority_deadline: float + + @classmethod + def admitted( + cls, receipt: Mapping[str, Any], *, task_id: str, run_id: str, owner: str, epoch: int, nonce: str, + heartbeat_timeout: int | None, cleanup: Mapping[str, str] | None, request_started: float, + ) -> PreparedAttempt: + _require( + receipt.get("prepared") is True and type(receipt.get("duplicate")) is bool + and "reason" in receipt and receipt["reason"] is None + and receipt.get("workflow_task_id") == task_id and receipt.get("lease_owner") == owner + and type(receipt.get("workflow_task_attempt")) is int and receipt["workflow_task_attempt"] == epoch + and receipt.get("worker_attempt_id") == nonce + and type(receipt.get("attempt_number")) is int and receipt["attempt_number"] > 0, + "local admission changed its original workflow claim or nonce", + ) + deadlines = {field: receipt.get(field) for field in (*_FIXED_DEADLINES, "heartbeat_deadline_at")} + attempt = cls(task_id, run_id, owner, epoch, nonce, _text(receipt.get("activity_execution_id")), + _text(receipt.get("activity_attempt_id")), receipt["attempt_number"], deadlines, + heartbeat_timeout, cleanup, request_started - _timestamp(receipt.get("server_time")).timestamp(), + request_started) + attempt.validate_cleanup(receipt) + server_time = _timestamp(receipt.get("server_time")) + _require(_timestamp(receipt.get("lease_expires_at")) > server_time, "local admission has an expired lease") + cleanup_deadline = _timestamp(cleanup["cleanup_deadline_at"]) if cleanup is not None else None + for field, deadline in deadlines.items(): + _require(field in receipt and (deadline is None or _timestamp(deadline) > server_time), + "local admission omitted or exhausted an execution deadline") + if deadline is not None and cleanup_deadline is not None: + _require(_timestamp(deadline) <= cleanup_deadline, + "local admission extended its original cleanup budget") + heartbeat = deadlines["heartbeat_deadline_at"] + if heartbeat_timeout is not None: + _require(heartbeat is not None + and _timestamp(heartbeat) <= server_time + timedelta(seconds=heartbeat_timeout), + "local admission changed its authored application heartbeat timeout") + else: + _require(_same_deadline(heartbeat, cleanup_deadline.isoformat() if cleanup_deadline is not None else None), + "local admission invented an application heartbeat deadline") + if cleanup_deadline is not None: + _require(cleanup_deadline > server_time, "local admission exhausted its original cleanup budget") + attempt.accept_budget(receipt, request_started) + attempt.remaining() + return attempt + + def validate_cleanup(self, receipt: Mapping[str, Any]) -> None: + actual = receipt.get("cancellation_cleanup") + if self.cleanup is None: + _require(actual is None, "local receipt invented cleanup authority") + return + fields = {"request_id", "root_request_id", "delivery_history_event_id", "cleanup_deadline_at"} + if not isinstance(actual, Mapping) or set(actual) != fields or set(self.cleanup) != fields: + raise LocalActivityExecutionAborted("local receipt changed canonical cleanup authority") + for field in fields: + _require(_same_deadline(actual[field], self.cleanup[field]) if field == "cleanup_deadline_at" + else _text(actual[field]) == _text(self.cleanup[field]), + "local receipt changed canonical cleanup authority") + + def validate_identity(self, receipt: Mapping[str, Any]) -> None: + _require( + receipt.get("workflow_task_id") == self.task_id + and type(receipt.get("workflow_task_attempt")) is int and receipt["workflow_task_attempt"] == self.epoch + and receipt.get("activity_execution_id") == self.execution_id + and receipt.get("activity_attempt_id") == self.attempt_id + and ("lease_owner" not in receipt or receipt["lease_owner"] == self.owner), + "prepared receipt belongs to a different attempt or workflow claim", + ) + + def validate_control(self, receipt: Mapping[str, Any], *, heartbeat: bool = False) -> None: + self.validate_identity(receipt) + self.validate_cleanup(receipt) + active = receipt.get("active") + _require( + type(active) is bool and type(receipt.get("renewed")) is bool + and receipt.get("lease_owner") == self.owner + and type(receipt.get("stop_required")) is bool and receipt["stop_required"] is not active + and "reason" in receipt + and (receipt["reason"] is None and receipt["renewed"] is (not heartbeat) if active + else isinstance(receipt["reason"], str) and bool(receipt["reason"]) and receipt["renewed"] is False), + "local control lacks live authority or an explicit stop", + ) + recorded = receipt.get("heartbeat_recorded") + _require( + recorded is active and (isinstance(receipt.get("heartbeat_history_event_id"), str) + and bool(receipt["heartbeat_history_event_id"].strip()) if active + else receipt.get("heartbeat_history_event_id") is None) if heartbeat + else recorded is False and receipt.get("heartbeat_history_event_id") is None, + "supervisor control and application heartbeat receipts must remain distinct", + ) + server_time = _timestamp(receipt.get("server_time")) + for field in _FIXED_DEADLINES: + _require(field in receipt and _same_deadline(receipt[field], self.deadlines[field]), + "local control changed an original execution deadline") + _require("heartbeat_deadline_at" in receipt, "local control omitted the application heartbeat deadline") + updated = receipt["heartbeat_deadline_at"] + if heartbeat and active and self.heartbeat_timeout is not None: + _require(_timestamp(self.deadlines["heartbeat_deadline_at"]) > server_time + and _timestamp(updated) >= _timestamp(self.deadlines["heartbeat_deadline_at"]) + and _timestamp(updated) <= server_time + timedelta(seconds=self.heartbeat_timeout) + and (self.cleanup is None + or _timestamp(updated) <= _timestamp(self.cleanup["cleanup_deadline_at"])), + "application heartbeat extended a fixed budget or revived an expired deadline") + else: + _require(_same_deadline(updated, self.deadlines["heartbeat_deadline_at"]), + "local control changed an unacknowledged application heartbeat deadline") + if active: + for field in ("lease_expires_at", "workflow_lease_expires_at"): + _require(_timestamp(receipt.get(field)) > server_time, "local control returned an expired lease") + for deadline in (receipt[field] for field in (*_FIXED_DEADLINES, "heartbeat_deadline_at")): + _require(deadline is None or _timestamp(deadline) > server_time, + "local control returned active after an execution deadline") + if self.cleanup is not None: + _require(_timestamp(self.cleanup["cleanup_deadline_at"]) > server_time, + "local control returned active after the original cleanup deadline") + if heartbeat and active: + self.deadlines["heartbeat_deadline_at"] = updated + + def validate_outcome(self, receipt: Mapping[str, Any]) -> None: + self.validate_identity(receipt) + retry = bool(receipt.get("event_type") == "ActivityRetryScheduled") + created = receipt.get("created_task_ids") + _require( + receipt.get("recorded") is True and type(receipt.get("duplicate")) is bool + and "reason" in receipt and receipt["reason"] is None + and receipt.get("workflow_run_id") == self.run_id and receipt.get("worker_attempt_id") == self.nonce + and receipt.get("event_type") in {"ActivityCompleted", "ActivityFailed", "ActivityTimedOut", + "ActivityRetryScheduled"} + and receipt.get("claim_released") is retry and isinstance(created, list) and len(created) == int(retry), + "local outcome lacks a canonical receipt for its original attempt", + ) + _text(receipt.get("event_id")) + _timestamp(receipt.get("recorded_at")) + for identity in receipt["created_task_ids"]: + _text(identity) + + def accept_budget(self, receipt: Mapping[str, Any], started: float) -> None: + server_time = _timestamp(receipt.get("server_time")).timestamp() + values = [receipt.get(field) for field in ( + "lease_expires_at", "workflow_lease_expires_at", *_FIXED_DEADLINES, "heartbeat_deadline_at", + )] + if self.cleanup is not None: + values.append(self.cleanup["cleanup_deadline_at"]) + self.authority_deadline = min( + min(self.clock_origin + _timestamp(value).timestamp(), + started + _timestamp(value).timestamp() - server_time) + for value in values if value is not None + ) + + def remaining(self) -> float: + remaining = self.authority_deadline - time.monotonic() + _require(remaining > 0, "prepared local original authority budget expired") + return min(5.0, remaining) + + +class PreparedLocalRunner: + def __init__( + self, client: Client, attempt: PreparedAttempt, *, observe: Callable[[Mapping[str, Any]], Any], + shutdown: asyncio.Event, + ) -> None: + self.client = client + self.attempt = attempt + self.observe = observe + self.shutdown = shutdown + self.stop_request: CancellationContext | None = None + self.stop_acknowledged = False + self.lock = asyncio.Lock() + self.callback: SupervisedCallback | None = None + + async def operation(self, operation: str, body: Mapping[str, Any]) -> dict[str, Any]: + return await self.client.prepared_local_activity_operation( + task_id=self.attempt.task_id, lease_owner=self.attempt.owner, workflow_task_attempt=self.attempt.epoch, + activity_attempt_id=self.attempt.attempt_id, operation=operation, body=body, + timeout_seconds=self.attempt.remaining(), + ) + + async def control(self, details: dict[str, Any] | None = None, *, heartbeat: bool = False) -> None: + async with self.lock: + _require(not self.shutdown.is_set(), "worker shutdown abandoned its prepared local claim") + started = time.monotonic() + receipt = await self.operation("heartbeat" if heartbeat else "control", + {"progress": {"details": details} if details else {}} + if heartbeat else {"renew_lease": True}) + self.attempt.validate_control(receipt, heartbeat=heartbeat) + if receipt["active"]: + self.attempt.accept_budget(receipt, started) + self.attempt.remaining() + return + if receipt["reason"] not in {"cancellation_requested", "cancellation_deadline_expired"}: + raise LocalActivityExecutionAborted("prepared local callback lost its original authority") + try: + context = CancellationContext.from_dict(receipt.get("cancellation_request", {})) + except (ValueError, TypeError, AttributeError) as error: + raise LocalActivityExecutionAborted("prepared stop lacks canonical cancellation context") from error + _require(context.lineage[-1].workflow_run_id == self.attempt.run_id and receipt.get("fenced") is True, + "prepared cancellation did not fence this original callback") + _text(receipt.get("cancellation_history_event_id")) + self.observe({**context.to_dict(), "history_refresh_page_token": receipt.get("history_refresh_page_token")}) + self.stop_request = context + raise PreparedCancellationObserved("prepared callback observed cooperative cancellation") + + async def execute(self, invocation: CallbackInvocation, processes: set[SupervisedCallback]) -> dict[str, Any]: + async def observe() -> None: + while True: + await asyncio.sleep(min(1.0, self.attempt.remaining())) + await self.control() + + async def heartbeat(details: dict[str, Any] | None) -> None: + await self.control(details, heartbeat=True) + + try: + await self.control() + except PreparedCancellationObserved: + # Admission was fenced before any process existed. + await self.acknowledge_stop() + raise + callback = SupervisedCallback(invocation) + self.callback = callback + processes.add(callback) + background: list[asyncio.Task[Any]] = [] + try: + await asyncio.wait_for(callback.start(), timeout=self.attempt.remaining()) + result = asyncio.create_task(callback.result(heartbeat)) + observer = asyncio.create_task(observe()) + shutdown = asyncio.create_task(self.shutdown.wait()) + background = [result, observer, shutdown] + done, _ = await asyncio.wait(background, return_when=asyncio.FIRST_COMPLETED) + if shutdown in done: + raise LocalActivityExecutionAborted("worker shutdown abandoned its prepared local claim") + if observer in done: + await observer + raise LocalActivityExecutionAborted("prepared local authority observer stopped") + outcome = await result + await self.control() + report: dict[str, Any] + if outcome.failure is None: + report = {"outcome": "completed", "result": serializer.envelope(outcome.value), + "payload_codec": serializer.AVRO_CODEC} + else: + failure = outcome.failure + report = {"outcome": "failed", "message": failure.message, "exception_type": failure.failure_type, + "non_retryable": failure.non_retryable or failure.cancelled} + # Cancel the observer before publication. No control may race a + # committed result and mistake its closed attempt for lost authority. + observer.cancel() + await asyncio.gather(observer, return_exceptions=True) + async with self.lock: + receipt = await self.operation("outcome", {"report": report}) + self.attempt.validate_outcome(receipt) + return receipt + finally: + async def finish() -> None: + for pending in background: + if not pending.done(): + pending.cancel() + await asyncio.gather(*background, return_exceptions=True) + if not callback.stopped: + await callback.stop() + _require(callback.stopped, "prepared callback stop remains unconfirmed") + processes.discard(callback) + await self.acknowledge_stop() + + joined = asyncio.create_task(finish()) + while not joined.done(): + # The original cancellation still propagates after finally. + with contextlib.suppress(asyncio.CancelledError): + await asyncio.shield(joined) + joined.result() + + async def acknowledge_stop(self) -> None: + if self.stop_request is None or self.stop_acknowledged: + return + # A physical join, or the absence of any spawned callback, is proved + # before this diagnostic receipt. It grants no execution authority. + receipt = await self.client.prepared_local_activity_operation( + task_id=self.attempt.task_id, lease_owner=self.attempt.owner, + workflow_task_attempt=self.attempt.epoch, activity_attempt_id=self.attempt.attempt_id, + operation="acknowledge-cancellation", body={"request_id": self.stop_request.request_id}, + ) + _require(receipt.get("acknowledged") is True and type(receipt.get("duplicate")) is bool + and "reason" in receipt and receipt["reason"] is None, + "joined prepared callback stop was not acknowledged") + _text(receipt.get("history_event_id")) + self.stop_acknowledged = True diff --git a/src/durable_workflow/cancellation.py b/src/durable_workflow/cancellation.py new file mode 100644 index 0000000..edc2c94 --- /dev/null +++ b/src/durable_workflow/cancellation.py @@ -0,0 +1,386 @@ +"""Immutable cancellation metadata recorded by the workflow runtime.""" + +from __future__ import annotations + +import re +from collections.abc import Callable, Mapping +from dataclasses import dataclass, field, replace +from datetime import datetime, timezone +from enum import Enum +from types import MappingProxyType +from typing import Any + + +class CancellationPolicy(str, Enum): + """Cancellation at an awaiting operation. Activities default to Try, children to Abandon.""" + + TRY_CANCEL = "try_cancel" + WAIT_CANCELLATION_COMPLETED = "wait_cancellation_completed" + ABANDON = "abandon" + + +class ParentClosePolicy(str, Enum): + """Open-child behavior after parent closure. REQUEST_CANCEL is legacy terminal cancellation.""" + + ABANDON = "abandon" + REQUEST_CANCEL = "request_cancel" + REQUEST_CANCELLATION = "request_cancellation" + TERMINATE = "terminate" + + +def _canonical_child_policies(options: Mapping[str, Any]) -> dict[str, str]: + policies: dict[str, str] = {} + for name, enum in ( + ("parent_close_policy", ParentClosePolicy), + ("cancellation_policy", CancellationPolicy), + ): + value = options.get(name) + if value is None: + continue + if isinstance(value, Enum) and not isinstance(value, enum): + raise ValueError(f"child workflow {name} must be a supported policy") + if not isinstance(value, str): + raise ValueError(f"child workflow {name} must be a supported policy") + try: + policies[name] = enum(value).value + except ValueError as error: + raise ValueError(f"child workflow {name} must be a supported policy") from error + return policies + + +def _canonical_activity_policy(value: Any) -> str | None: + if value is None: + return None + if isinstance(value, Enum) and not isinstance(value, CancellationPolicy): + raise ValueError("remote activity cancellation_policy must be a supported policy") + if not isinstance(value, str): + raise ValueError("remote activity cancellation_policy must be a supported policy") + try: + return CancellationPolicy(value).value + except ValueError as error: + raise ValueError("remote activity cancellation_policy must be a supported policy") from error + + +def _text(snapshot: Mapping[str, Any], key: str) -> str: + value = snapshot.get(key) + if not isinstance(value, str) or not value.strip(): + raise ValueError(f"cancellation {key} must be a non-empty string") + return value + + +def _timestamp(value: str) -> datetime: + if not re.fullmatch(r"\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}(?:\.\d+)?(?:Z|[+-]\d{2}:\d{2})", value): + raise ValueError("cancellation timestamp requires an ISO date, time and timezone") + return datetime.fromisoformat(value.replace("Z", "+00:00")).astimezone(timezone.utc) + + +@dataclass(frozen=True) +class CancellationLineage: + """One local request in the ordered root-to-descendant lineage.""" + + request_id: str + workflow_instance_id: str + workflow_run_id: str + + def to_dict(self) -> dict[str, str]: + return { + "request_id": self.request_id, + "workflow_instance_id": self.workflow_instance_id, + "workflow_run_id": self.workflow_run_id, + } + + +@dataclass(frozen=True) +class CancellationContext: + """Original request metadata and budget, visible at committed delivery. + + Nested metadata is immutable. A child has a distinct local request ID + while retaining its root identity and request time. A scoped child can + narrow its cleanup deadline while preserving the original root budget. + """ + + request_id: str + root_request_id: str + root_workflow_instance_id: str + root_workflow_run_id: str + parent_request_id: str | None + reason: str | None + requester: Mapping[str, str] + source: str + requested_at: datetime + cleanup_deadline_at: datetime + lineage: tuple[CancellationLineage, ...] + _replay_clock: Callable[[], datetime] | None = field(default=None, repr=False, compare=False) + scope_origin: ScopedCancellationContext | None = None + + def __post_init__(self) -> None: + object.__setattr__(self, "requester", MappingProxyType(dict(self.requester))) + object.__setattr__(self, "lineage", tuple(self.lineage)) + + @property + def deadline(self) -> datetime: + """The original immutable cleanup deadline.""" + return self.cleanup_deadline_at + + def remaining(self) -> float: + """Seconds left at the consumed replay boundary, clamped to zero. + + Available only during the workflow replay that delivered this context. + Detached metadata has no clock. Host time never supplies this value. + """ + if self._replay_clock is None: + raise RuntimeError("cancellation remaining time requires active workflow replay") + return max(0.0, (self.cleanup_deadline_at - self._replay_clock()).total_seconds()) + + def _with_replay_clock(self, clock: Callable[[], datetime]) -> CancellationContext: + return replace(self, _replay_clock=clock) + + @classmethod + def from_dict(cls, snapshot: Mapping[str, Any]) -> CancellationContext: + schema = snapshot.get("schema") + if schema not in ("durable-workflow.cancellation-context/v1", "durable-workflow.cancellation-context/v2"): + raise ValueError("unsupported cancellation context schema") + scope_origin = None + if schema == "durable-workflow.cancellation-context/v2": + origin = snapshot.get("scope_origin") + if not isinstance(origin, Mapping): + raise ValueError("scoped run cancellation requires its original scope context") + scope_origin = ScopedCancellationContext.from_dict(origin) + elif "scope_origin" in snapshot or "scope_authority_deadline_at" in snapshot: + raise ValueError("legacy cancellation context cannot discard a scoped origin") + requester = snapshot.get("requester") + if not isinstance(requester, Mapping) or not requester: + raise ValueError("cancellation requester must identify its caller") + normalized_requester: dict[str, str] = {} + for key, value in requester.items(): + if key not in {"type", "id", "label"} or not isinstance(value, str) or value == "": + raise ValueError("cancellation requester contains unsupported metadata") + normalized_requester[key] = value + lineage = snapshot.get("lineage") + if not isinstance(lineage, list) or not lineage: + raise ValueError("cancellation lineage must contain the root request") + normalized = [] + for entry in lineage: + if not isinstance(entry, Mapping): + raise ValueError("cancellation lineage entry is invalid") + normalized.append(CancellationLineage( + _text(entry, "request_id"), _text(entry, "workflow_instance_id"), _text(entry, "workflow_run_id"), + )) + request_id = _text(snapshot, "request_id") + root_request_id = _text(snapshot, "root_request_id") + root_instance_id = _text(snapshot, "root_workflow_instance_id") + root_run_id = _text(snapshot, "root_workflow_run_id") + parent_request_id = snapshot.get("parent_request_id") + reason = snapshot.get("reason") + if ( + parent_request_id is not None and (not isinstance(parent_request_id, str) or not parent_request_id) + or reason is not None and not isinstance(reason, str) + ): + raise ValueError("cancellation parent identity or reason is invalid") + if len({entry.request_id for entry in normalized}) != len(normalized) or len({ + entry.workflow_run_id for entry in normalized + }) != len(normalized): + raise ValueError("cancellation lineage cannot contain a cycle") + expected_parent = normalized[-2].request_id if len(normalized) > 1 else None + if ( + normalized[0] != CancellationLineage(root_request_id, root_instance_id, root_run_id) + or normalized[-1].request_id != request_id + or scope_origin is None and parent_request_id != expected_parent + ): + raise ValueError("cancellation lineage does not match its request identities") + requested_at = _timestamp(_text(snapshot, "requested_at")) + deadline = _timestamp(_text(snapshot, "cleanup_deadline_at")) + if deadline <= requested_at: + raise ValueError("cancellation deadline must follow the original request") + if scope_origin is not None: + _assert_scope_origin(snapshot, scope_origin, normalized, normalized_requester, requested_at, deadline) + return cls( + request_id, root_request_id, root_instance_id, root_run_id, parent_request_id, + reason, normalized_requester, _text(snapshot, "source"), requested_at, deadline, tuple(normalized), + scope_origin=scope_origin, + ) + + def to_dict(self) -> dict[str, Any]: + """Return detached metadata in the portable context schema.""" + snapshot: dict[str, Any] = { + "schema": "durable-workflow.cancellation-context/v1" if self.scope_origin is None + else "durable-workflow.cancellation-context/v2", + "request_id": self.request_id, + "root_request_id": self.root_request_id, + "root_workflow_instance_id": self.root_workflow_instance_id, + "root_workflow_run_id": self.root_workflow_run_id, + "parent_request_id": self.parent_request_id, + "reason": self.reason, + "requester": dict(self.requester), + "source": self.source, + "requested_at": self.requested_at.isoformat(timespec="microseconds").replace("+00:00", "Z"), + "cleanup_deadline_at": self.cleanup_deadline_at.isoformat(timespec="microseconds").replace("+00:00", "Z"), + "lineage": [entry.to_dict() for entry in self.lineage], + } + if self.scope_origin is not None: + snapshot["scope_origin"] = self.scope_origin.to_dict() + snapshot["scope_authority_deadline_at"] = snapshot["cleanup_deadline_at"] + return snapshot + + +@dataclass(frozen=True) +class ScopedCancellationLineage: + """One immutable original scope address and its bounded cleanup budget.""" + + request_id: str + workflow_instance_id: str + workflow_run_id: str + scope_id: str + cleanup_deadline_at: datetime + + def to_dict(self) -> dict[str, str]: + return { + "request_id": self.request_id, + "workflow_instance_id": self.workflow_instance_id, + "workflow_run_id": self.workflow_run_id, + "scope_id": self.scope_id, + "cleanup_deadline_at": self.cleanup_deadline_at.isoformat(timespec="microseconds").replace("+00:00", "Z"), + } + + +@dataclass(frozen=True) +class ScopedCancellationContext: + """Original scope ancestry carried by a candidate cooperative child request. + + This is immutable metadata, not permission to enter a scope body. + """ + + root_context: CancellationContext + lineage: tuple[ScopedCancellationLineage, ...] + + def __post_init__(self) -> None: + object.__setattr__(self, "lineage", tuple(self.lineage)) + + @property + def request_id(self) -> str: + return self.lineage[-1].request_id + + @property + def parent_request_id(self) -> str | None: + return self.lineage[-2].request_id if len(self.lineage) > 1 else None + + @property + def workflow_instance_id(self) -> str: + return self.lineage[-1].workflow_instance_id + + @property + def workflow_run_id(self) -> str: + return self.lineage[-1].workflow_run_id + + @property + def scope_id(self) -> str: + return self.lineage[-1].scope_id + + @property + def root_scope_id(self) -> str: + return self.lineage[0].scope_id + + @property + def requested_at(self) -> datetime: + return self.root_context.requested_at + + @property + def root_deadline(self) -> datetime: + return self.root_context.deadline + + @property + def deadline(self) -> datetime: + return self.lineage[-1].cleanup_deadline_at + + @classmethod + def from_dict(cls, snapshot: Mapping[str, Any]) -> ScopedCancellationContext: + if set(snapshot) != {"schema", "root_context", "lineage"}: + raise ValueError("scoped cancellation context has missing or unsupported fields") + root_snapshot = snapshot["root_context"] + if (snapshot["schema"] != "durable-workflow.scoped-cancellation-context/v1" + or not isinstance(root_snapshot, Mapping) + or root_snapshot.get("schema") != "durable-workflow.cancellation-context/v1"): + raise ValueError("scoped cancellation requires its original root context") + root = CancellationContext.from_dict(root_snapshot) + if (set(root_snapshot) != set(root.to_dict()) or root.request_id != root.root_request_id + or root.parent_request_id is not None or len(root.lineage) != 1): + raise ValueError("scoped cancellation root must contain the original root request") + lineage = snapshot["lineage"] + if not isinstance(lineage, list) or not lineage: + raise ValueError("scoped cancellation lineage must contain its root address") + requests: set[str] = set() + addresses: set[tuple[str, str]] = set() + instances_by_run: dict[str, str] = {} + last_run: str | None = None + deadline = root.deadline + normalized = [] + for index, entry in enumerate(lineage): + if not isinstance(entry, Mapping) or set(entry) != { + "request_id", "workflow_instance_id", "workflow_run_id", "scope_id", "cleanup_deadline_at", + }: + raise ValueError("scoped cancellation address has missing or unsupported fields") + address = ScopedCancellationLineage( + _text(entry, "request_id"), _text(entry, "workflow_instance_id"), _text(entry, "workflow_run_id"), + _text(entry, "scope_id"), _timestamp(_text(entry, "cleanup_deadline_at")), + ) + if index == 0 and ( + address.request_id != root.request_id or address.workflow_instance_id != root.root_workflow_instance_id + or address.workflow_run_id != root.root_workflow_run_id or address.cleanup_deadline_at != root.deadline + ): + raise ValueError("scoped cancellation root address does not match its request") + key = (address.workflow_run_id, address.scope_id) + if address.request_id in requests or key in addresses: + raise ValueError("scoped cancellation cannot repeat a request or address") + if address.workflow_run_id in instances_by_run and ( + instances_by_run[address.workflow_run_id] != address.workflow_instance_id + or last_run != address.workflow_run_id + ): + raise ValueError("scoped cancellation cannot reenter or reassign an earlier run") + if address.cleanup_deadline_at <= root.requested_at or address.cleanup_deadline_at > deadline: + raise ValueError("scoped cancellation cannot extend a descendant budget") + normalized.append(address) + requests.add(address.request_id) + addresses.add(key) + instances_by_run[address.workflow_run_id] = address.workflow_instance_id + last_run = address.workflow_run_id + deadline = address.cleanup_deadline_at + return cls(root, tuple(normalized)) + + def to_dict(self) -> dict[str, Any]: + return { + "schema": "durable-workflow.scoped-cancellation-context/v1", + "root_context": self.root_context.to_dict(), + "lineage": [entry.to_dict() for entry in self.lineage], + } + + +def _assert_scope_origin( + snapshot: Mapping[str, Any], origin: ScopedCancellationContext, lineage: list[CancellationLineage], + requester: dict[str, str], requested_at: datetime, deadline: datetime, +) -> None: + root = origin.root_context + last = lineage[-1] + if ( + snapshot.get("parent_request_id") != origin.request_id + or snapshot["root_request_id"] != root.root_request_id + or snapshot["root_workflow_instance_id"] != root.root_workflow_instance_id + or snapshot["root_workflow_run_id"] != root.root_workflow_run_id + or snapshot.get("reason") != root.reason or requester != root.requester + or _text(snapshot, "source") != root.source or requested_at != root.requested_at + or deadline != _timestamp(_text(snapshot, "scope_authority_deadline_at")) or deadline > origin.deadline + or last.request_id in {entry.request_id for entry in origin.lineage} + or last.workflow_run_id in {entry.workflow_run_id for entry in origin.lineage} + ): + raise ValueError("run cancellation does not preserve its original scope context") + expected = [root.lineage[0]] + for entry in origin.lineage: + if entry.workflow_run_id == root.root_workflow_run_id: + continue + address = CancellationLineage(entry.request_id, entry.workflow_instance_id, entry.workflow_run_id) + if expected[-1].workflow_run_id == entry.workflow_run_id: + expected[-1] = address + else: + expected.append(address) + expected.append(last) + if lineage != expected: + raise ValueError("run cancellation lineage discards or replaces its original scope ancestry") diff --git a/src/durable_workflow/client.py b/src/durable_workflow/client.py index 7dd684f..05dfe31 100644 --- a/src/durable_workflow/client.py +++ b/src/durable_workflow/client.py @@ -37,11 +37,13 @@ import httpx from . import serializer +from ._cooperative_cancellation import CancellationDelivery, CancellationRequest from .errors import ( ExternalPayloadIntegrityMismatch, ExternalPayloadOversized, ExternalPayloadUnavailable, ExternalPayloadUnsupported, + NonDeterministicReplayError, RuntimeCapabilityUnsupported, RuntimeDiscoveryUnavailable, ServerError, @@ -87,6 +89,7 @@ CONTROL_PLANE_REQUEST_CONTRACT_SCHEMA = "durable-workflow.v2.control-plane-request.contract" CONTROL_PLANE_REQUEST_CONTRACT_VERSION = 1 _QUERY_TASKS_DISCOVERY_PATH = "worker_protocol.server_capabilities.query_tasks" +_COOPERATIVE_CANCELLATION_DISCOVERY_PATH = "worker_protocol.server_capabilities.cooperative_cancellation" _UPDATE_WAIT_STAGES_DISCOVERY_PATH = ( "control_plane.request_contract.operations.update.fields.wait_for.canonical_values" ) @@ -105,10 +108,32 @@ ) _RUNTIME_EXTERNAL_PAYLOAD_ERROR_BODY_LIMIT = 64 * 1024 _PAYLOAD_COMPLETION_SCHEMA = "durable-workflow.v2.payload-completion-context.v1" +_PREPARED_PAYLOAD_COMPLETION_SCHEMA = "durable-workflow.v2.payload-completion-context.v2" _PAYLOAD_COMPLETION_HEADER = "X-Durable-Workflow-Payload-Completion" -def _payload_completion_context(path: str, body: Any) -> dict[str, Any] | None: +def _payload_completion_context(path: str, body: Any, *, allow_prepared: bool = False) -> dict[str, Any] | None: + prepared = re.fullmatch( + r"/worker/workflow-tasks/([^/]+)/local-activities/(?:(checkpoint|checkpoint-group|prepare|recover)|([^/]+)/outcome)", + path.split("?")[0], + ) + if allow_prepared and prepared is not None and isinstance(body, dict): + task_id, operation, activity_attempt_id = prepared.groups() + operation = operation or "outcome" + identity = ({"checkpoint_id": body.get("checkpoint_id")} if operation in {"checkpoint", "checkpoint-group"} + else {"activity_attempt_id": unquote(activity_attempt_id)} if activity_attempt_id is not None + else {"sequence": body.get("sequence")}) + value = next(iter(identity.values())) + owner, attempt = body.get("lease_owner"), body.get("workflow_task_attempt") + identity_valid = (type(value) is int and value > 0) if operation in {"prepare", "recover"} \ + else isinstance(value, str) and bool(value.strip()) + if (not isinstance(owner, str) or not owner.strip() or type(attempt) is not int or attempt < 1 + or not identity_valid): + return None + return {"schema": _PREPARED_PAYLOAD_COMPLETION_SCHEMA, "kind": "workflow", "task_id": unquote(task_id), + "attempt": attempt, "lease_owner": owner, + "operation": "local_activity_group_checkpoint" if operation == "checkpoint-group" + else "local_activity_" + operation, **identity} match = re.fullmatch(r"/worker/(activity|workflow|query)-tasks/([^/]+)/(complete|fail)", path.split("?")[0]) if match is None or not isinstance(body, dict): return None @@ -183,6 +208,13 @@ def _worker_protocol_supports_message_streams() -> bool: return (int(parts[0]), int(parts[1])) >= _MESSAGE_STREAMS_MINIMUM_WORKER_PROTOCOL +def _supports_cooperative_cancellation_protocol(version: Any) -> bool: + if not isinstance(version, str): + return False + match = re.fullmatch(r"1\.([0-9]{1,4})", version) + return match is not None and int(match.group(1)) >= 20 + + def _normalize_base_url(base_url: str) -> str: parsed = urlsplit(base_url) if parsed.query or parsed.fragment: @@ -301,6 +333,7 @@ class _RuntimeExternalPayloadTransport: request_timeout_seconds: float status: str completion_context: bool = False + prepared_completion_context: bool = False @dataclass @@ -1330,6 +1363,15 @@ async def cancel(self, *, reason: str | None = None) -> None: """Close this workflow's current run as cancelled. See :meth:`Client.cancel_workflow`.""" await self._client.cancel_workflow(self.workflow_id, reason=reason) + async def request_cancellation( + self, *, reason: str | None = None, cleanup_timeout_seconds: int | None = None, + ) -> dict[str, Any]: + """Request bounded cleanup for this run. See :meth:`Client.request_workflow_cancellation`.""" + return await self._client.request_workflow_cancellation( + self.workflow_id, run_id=self.run_id, reason=reason, + cleanup_timeout_seconds=cleanup_timeout_seconds, + ) + async def terminate(self, *, reason: str | None = None) -> None: """Forcefully stop this workflow. See :meth:`Client.terminate_workflow`.""" await self._client.terminate_workflow(self.workflow_id, reason=reason) @@ -1777,7 +1819,8 @@ async def _request( transport=transport, uploaded={}, completion=( - _payload_completion_context(path, json) if worker and transport.completion_context else None + _payload_completion_context(path, json, allow_prepared=transport.prepared_completion_context) + if worker and transport.completion_context else None ), ) @@ -1958,6 +2001,11 @@ async def _runtime_external_payload_transport( completion_context=isinstance(completion, dict) and completion.get("schema") == _PAYLOAD_COMPLETION_SCHEMA and completion.get("header") == _PAYLOAD_COMPLETION_HEADER, + prepared_completion_context=isinstance(completion, dict) + and completion.get("prepared_schema") == _PREPARED_PAYLOAD_COMPLETION_SCHEMA + and _supports_cooperative_cancellation_protocol(_protocol_version_from_env( + "DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", PROTOCOL_VERSION, + )), ) self._runtime_external_payload_transport_cache = transport self._runtime_external_payload_transport_resolved = True @@ -2522,6 +2570,32 @@ async def _require_query_support(self) -> None: ), ) + async def _require_cooperative_cancellation_support(self) -> None: + operation = "Client.request_workflow_cancellation" + info = await self._runtime_discovery( + operation=operation, + required_path=_COOPERATIVE_CANCELLATION_DISCOVERY_PATH, + ) + protocol = info.get("worker_protocol") + capabilities = protocol.get("server_capabilities") if isinstance(protocol, dict) else None + supported = capabilities.get("cooperative_cancellation") if isinstance(capabilities, dict) else None + if supported is False: + raise RuntimeCapabilityUnsupported( + operation, _COOPERATIVE_CANCELLATION_DISCOVERY_PATH, + "This runtime does not support cooperative cancellation. " + "Existing cancel and terminate close immediately.", + ) + if supported is not True: + raise RuntimeDiscoveryUnavailable( + operation, _COOPERATIVE_CANCELLATION_DISCOVERY_PATH, + "Runtime discovery did not advertise cooperative cancellation support.", + ) + if not isinstance(protocol, dict) or not _supports_cooperative_cancellation_protocol(protocol.get("version")): + raise RuntimeDiscoveryUnavailable( + operation, "worker_protocol.version", + "Cooperative cancellation requires an advertised compatible worker protocol of at least 1.20.", + ) + async def _require_update_wait_stage(self, wait_for: str) -> None: info = await self._runtime_discovery( operation="Client.update_workflow", @@ -4140,14 +4214,61 @@ async def query_workflow( "POST", f"/workflows/{workflow_id}/query/{query_name}", json=body, context=workflow_id ) + async def request_workflow_cancellation( + self, + workflow_id: str, + *, + run_id: str | None = None, + reason: str | None = None, + cleanup_timeout_seconds: int | None = None, + ) -> dict[str, Any]: + """Request cancellation with bounded workflow-authored cleanup. + + Requires explicit runtime capability discovery and protocol 1.20. + Repeated requests return Server's original request ID and cleanup + deadline. No client-generated identity or deadline replaces them. + ``run_id`` fences the request to that current run when supplied. + """ + if not isinstance(workflow_id, str) or not workflow_id.strip(): + raise ValueError("workflow_id must be a non-empty string") + if run_id is not None and (not isinstance(run_id, str) or not run_id.strip()): + raise ValueError("run_id must be a non-empty string") + if cleanup_timeout_seconds is not None and ( + type(cleanup_timeout_seconds) is not int or not 1 <= cleanup_timeout_seconds <= 3600 + ): + raise ValueError("cleanup_timeout_seconds must be an integer from 1 to 3600") + await self._require_cooperative_cancellation_support() + path = f"/workflows/{quote(workflow_id, safe='._:-')}" + if run_id is not None: + path += f"/runs/{quote(run_id, safe='._:-')}" + body: dict[str, Any] = {} + if reason is not None: + body["reason"] = reason + if cleanup_timeout_seconds is not None: + body["cleanup_timeout_seconds"] = cleanup_timeout_seconds + result = await self._request("POST", f"{path}/request-cancellation", json=body, context=workflow_id) + if ( + not isinstance(result, dict) or result.get("accepted") is not True + or type(result.get("duplicate")) is not bool or result.get("workflow_id") != workflow_id + or not isinstance(result.get("run_id"), str) or not result["run_id"].strip() + or (run_id is not None and result["run_id"] != run_id) + or not isinstance(result.get("cancellation_request"), dict) + ): + raise ServerError(200, {"reason": "invalid_cooperative_cancellation_response"}) + try: + CancellationRequest.from_observation(result["cancellation_request"]) + except NonDeterministicReplayError as error: + raise ServerError(200, {"reason": "invalid_cooperative_cancellation_response"}) from error + return result + async def cancel_workflow(self, workflow_id: str, *, reason: str | None = None) -> None: """Close the current run as cancelled immediately. Server cancels open tasks and timers; it does not resume workflow code to run saga or ``finally`` cleanup. :meth:`terminate_workflow` also - closes immediately, with a distinct terminal outcome. Embedded - Laravel's cooperative ``requestCancellation()`` is not yet available - through this service-mode API. + closes immediately, with a distinct terminal outcome. + For capable runtimes, :meth:`request_workflow_cancellation` separately + requests bounded workflow-authored cleanup. """ body: dict[str, Any] = {} if reason is not None: @@ -5053,6 +5174,90 @@ async def heartbeat_workflow_task( }, ) + async def deliver_workflow_cancellation( + self, + *, + task_id: str, + lease_owner: str, + workflow_task_attempt: int, + request_id: str, + sequence: int, + call_kind: str, + sequence_span: int = 1, + operation_sequence: int | None = None, + operation_sequence_span: int = 1, + ) -> dict[str, Any]: + """Commit delivery at one authored call using the current task lease. + + Requires worker protocol 1.20 and a capable recorded claim. Retry + with the same request, attempt and operation range after a lost + acknowledgment, then reload canonical history before cleanup replay. + A child wait may explicitly release the claim. Return to polling and + replay the boundary on the successor claim when the child finishes. + Successful delivery does not release or renew the lease. + """ + version = _protocol_version_from_env("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", PROTOCOL_VERSION) + if not _supports_cooperative_cancellation_protocol(version): + raise ValueError("cooperative cancellation delivery requires worker protocol 1.20 or newer") + if not isinstance(task_id, str) or not task_id.strip(): + raise ValueError("task_id must be a non-empty string") + if not isinstance(lease_owner, str) or not lease_owner.strip(): + raise ValueError("lease_owner must be a non-empty string") + if type(workflow_task_attempt) is not int or workflow_task_attempt < 1: + raise ValueError("workflow_task_attempt must be a positive integer") + try: + delivery = CancellationDelivery.from_payload({ + "workflow_command_id": request_id, "sequence": sequence, "call_kind": call_kind, + "sequence_span": sequence_span, "operation_sequence": operation_sequence, + "operation_sequence_span": operation_sequence_span, + }) + except NonDeterministicReplayError as error: + raise ValueError("cancellation delivery must name a valid authored call and operation range") from error + body: dict[str, Any] = { + "lease_owner": lease_owner, "workflow_task_attempt": workflow_task_attempt, + "request_id": delivery.request_id, "sequence": delivery.sequence, + "call_kind": delivery.call_kind, "sequence_span": delivery.sequence_span, + } + if delivery.operation_sequence is not None: + body["operation_sequence"] = delivery.operation_sequence + body["operation_sequence_span"] = delivery.operation_sequence_span + result = await self._request( + "POST", f"/worker/workflow-tasks/{quote(task_id, safe='._:-')}/deliver-cancellation", + worker=True, json=body, + ) + if ( + isinstance(result, dict) + and result.get("delivered") is False + and ( + (result.get("reason") == "cancellation_waiting_for_child" + and delivery.call_kind in {"child", "parallel", "selection_handle"}) + or (result.get("reason") == "cancellation_waiting_for_activity" + and delivery.call_kind in {"activity", "local_activity", "parallel", "selection_handle"}) + ) + and result.get("claim_released") is True + and result.get("task_id") == task_id + and all(result.get(field) is None for field in ( + "request_id", "sequence", "call_kind", "sequence_span", + "operation_sequence", "operation_sequence_span", + )) + ): + return result + if not isinstance(result, dict) or result.get("delivered") is not True or result.get("task_id") != task_id: + raise ServerError(200, {"reason": "invalid_cooperative_cancellation_delivery"}) + try: + recorded = CancellationDelivery.from_payload({ + "workflow_command_id": result.get("request_id"), + "sequence": result.get("sequence"), "call_kind": result.get("call_kind"), + "sequence_span": result.get("sequence_span"), + "operation_sequence": result.get("operation_sequence"), + "operation_sequence_span": result.get("operation_sequence_span"), + }) + except NonDeterministicReplayError as error: + raise ServerError(200, {"reason": "invalid_cooperative_cancellation_delivery"}) from error + if recorded != delivery: + raise ServerError(200, {"reason": "invalid_cooperative_cancellation_delivery"}) + return result + async def workflow_task_history( self, *, @@ -5396,6 +5601,138 @@ async def fail_activity_task( "POST", f"/worker/activity-tasks/{task_id}/fail", worker=True, json=body ) + async def activity_task_status( + self, + *, + task_id: str, + activity_attempt_id: str, + lease_owner: str, + ) -> Any: + """Observe one cooperative activity claim without recording progress. + + Requires explicit worker protocol 1.20. This readonly observation does + not renew the activity lease or its heartbeat deadline. A positive reply + is not a reservation of ownership for a subsequent completion. + """ + if not _supports_cooperative_cancellation_protocol(_protocol_version_from_env( + "DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", PROTOCOL_VERSION, + )): + raise ValueError("activity ownership observation requires explicit worker protocol 1.20") + return await asyncio.wait_for(self._request( + "POST", f"/worker/activity-tasks/{task_id}/status", worker=True, + json={"activity_attempt_id": activity_attempt_id, "lease_owner": lease_owner}, timeout=5.0, + ), timeout=5.0) + + async def acknowledge_activity_cancellation( + self, + *, + task_id: str, + activity_attempt_id: str, + lease_owner: str, + request_id: str, + ) -> dict[str, Any]: + """Report that the original owner's remote callback has stopped. + + The caller must stop and join the callback before calling this method. + The receipt records diagnostic evidence and does not renew a lease, + heartbeat, publication authority or cleanup budget. Requires explicit + worker protocol 1.20. Retries share one five-second transport budget. + """ + if not _supports_cooperative_cancellation_protocol(_protocol_version_from_env( + "DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", PROTOCOL_VERSION, + )): + raise ValueError("activity cancellation acknowledgment requires explicit worker protocol 1.20") + identities = { + "task_id": task_id, + "activity_attempt_id": activity_attempt_id, + "lease_owner": lease_owner, + "request_id": request_id, + } + if any( + not isinstance(value, str) or not value.strip() or len(value.encode("utf-8")) > 255 + for value in identities.values() + ): + raise ValueError( + "activity cancellation acknowledgment requires bounded, nonempty claim and request identities" + ) + result = await asyncio.wait_for(self._request( + "POST", f"/worker/activity-tasks/{quote(task_id, safe='._:-')}/acknowledge-cancellation", + worker=True, + json={"activity_attempt_id": activity_attempt_id, "lease_owner": lease_owner, "request_id": request_id}, + timeout=5.0, + ), timeout=5.0) + if ( + not isinstance(result, dict) + or any(result.get(key) != value for key, value in identities.items()) + or result.get("acknowledged") is not True + or not isinstance(result.get("duplicate"), bool) + or result.get("reason") is not None + or result.get("heartbeat_recorded") is not False + or not isinstance(result.get("history_event_id"), str) + or not result["history_event_id"].strip() + ): + raise ServerError(200, {"reason": "invalid_activity_cancellation_acknowledgement"}) + return result + + async def prepared_local_activity_operation( + self, + *, + task_id: str, + lease_owner: str, + workflow_task_attempt: int, + operation: str, + body: Mapping[str, Any] | None = None, + activity_attempt_id: str | None = None, + timeout_seconds: float = 5.0, + ) -> dict[str, Any]: + """Perform one prepared-local operation on its original workflow claim. + + Source protocol 1.20 only. This transport does not grant permission to + invoke a callback. The worker must discover the installed bridge and + validate its original admission receipt, deadlines and canonical history. + Every retry and payload transfer shares this total authority budget. + """ + if not _supports_cooperative_cancellation_protocol(_protocol_version_from_env( + "DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", PROTOCOL_VERSION, + )): + raise ValueError("prepared local operations require explicit worker protocol 1.20") + admission = {"checkpoint", "checkpoint-group", "prepare", "recover"} + attempts = {"control", "heartbeat", "outcome", "acknowledge-cancellation"} + if operation not in admission | attempts: + raise ValueError("unsupported prepared local operation") + identifiers = [task_id, lease_owner] + if operation in attempts: + if activity_attempt_id is None: + raise ValueError("prepared local attempt operation requires its original backend attempt identity") + identifiers.append(activity_attempt_id) + elif activity_attempt_id is not None: + raise ValueError("prepared local admission cannot carry an activity attempt identity") + if any(not isinstance(value, str) or not value.strip() or len(value.encode("utf-8")) > 255 + for value in identifiers): + raise ValueError("prepared local operations require bounded nonempty original claim identities") + if type(workflow_task_attempt) is not int or workflow_task_attempt < 1: + raise ValueError("prepared local operations require a positive original workflow task epoch") + if (isinstance(timeout_seconds, bool) or not isinstance(timeout_seconds, int | float) + or not math.isfinite(timeout_seconds) or not 0 < timeout_seconds <= 5): + raise ValueError("prepared local authority budget must be positive, finite and at most five seconds") + if body is not None and not isinstance(body, Mapping): + raise TypeError("prepared local operation body must be an object") + payload = dict(body or {}) + if "lease_owner" in payload or "workflow_task_attempt" in payload: + raise ValueError("prepared local operation body cannot replace original workflow claim authority") + payload = {"lease_owner": lease_owner, "workflow_task_attempt": workflow_task_attempt, **payload} + json_module.dumps(payload, allow_nan=False) + path = f"/worker/workflow-tasks/{quote(task_id, safe='._:-')}/local-activities/" + if activity_attempt_id is not None: + path += quote(activity_attempt_id, safe="._:-") + "/" + result = await asyncio.wait_for( + self._request("POST", path + operation, worker=True, json=payload, timeout=timeout_seconds), + timeout=timeout_seconds, + ) + if not isinstance(result, dict): + raise ServerError(200, {"reason": "invalid_prepared_local_activity_receipt"}) + return result + async def heartbeat_activity_task( self, *, diff --git a/src/durable_workflow/errors.py b/src/durable_workflow/errors.py index 0f8ac20..bb1f0cd 100644 --- a/src/durable_workflow/errors.py +++ b/src/durable_workflow/errors.py @@ -20,6 +20,8 @@ class explicitly (``except (ActivityCancelled, ...):``). This mirrors the way from typing import Any +from .cancellation import CancellationContext + class DurableWorkflowError(Exception): """Base class for every exception raised by the SDK.""" @@ -655,8 +657,13 @@ class WorkflowCancelled(BaseException): class by name. """ - def __init__(self, message: str = "workflow was cancelled") -> None: + def __init__( + self, message: str = "workflow was cancelled", *, request_id: str | None = None, + context: CancellationContext | None = None, + ) -> None: super().__init__(message) + self.request_id = request_id + self.context = context class ActivityCancelled(BaseException): diff --git a/src/durable_workflow/worker.py b/src/durable_workflow/worker.py index 8046e82..37a1d2b 100644 --- a/src/durable_workflow/worker.py +++ b/src/durable_workflow/worker.py @@ -17,10 +17,12 @@ import asyncio import contextlib +import contextvars import hashlib import inspect import json import logging +import pickle import sys import threading import time @@ -28,12 +30,16 @@ import types import uuid from collections.abc import Awaitable, Callable, Coroutine, Iterable, Mapping +from concurrent.futures import Future, ThreadPoolExecutor from datetime import datetime, timezone from functools import wraps from types import FunctionType from typing import Annotated, Any, Concatenate, Literal, ParamSpec, TypeVar, Union, get_args, get_origin, get_type_hints from . import serializer +from ._activity_process import CallbackFailure, CallbackInvocation, CallbackProcessLost, SupervisedCallback +from ._cooperative_cancellation import CancellationRequest, read_cancellation_history +from ._prepared_local_activity import PreparedAttempt, PreparedCancellationObserved, PreparedLocalRunner from .activity import ActivityContext, ActivityInfo, _set_context from .auth_composition import ( AUTH_COMPOSITION_CONTRACT_SCHEMA, @@ -41,6 +47,7 @@ AuthCompositionContractError, parse_auth_composition_contract, ) +from .cancellation import CancellationContext from .client import ( CONTROL_PLANE_REQUEST_CONTRACT_SCHEMA, CONTROL_PLANE_REQUEST_CONTRACT_VERSION, @@ -49,6 +56,8 @@ PROTOCOL_VERSION, Client, WorkflowExecution, + _protocol_version_from_env, + _supports_cooperative_cancellation_protocol, ) from .errors import ( ActivityCancelled, @@ -85,8 +94,12 @@ FailWorkflow, LocalActivityExecutionAborted, NexusServiceCall, + PreparedLocalActivityCall, RecordLocalActivity, RecordSideEffect, + ReplayOutcome, + ScheduleActivity, + StartChildWorkflow, UpsertMemo, apply_update, commands_to_server_commands, @@ -153,12 +166,30 @@ _R = TypeVar("_R") +class _RemoteActivityExecutionAborted(Exception): + """Ownership cannot authorize another remote callback boundary.""" + + +class _RemoteActivityApplicationFailure(Exception): + def __init__(self, failure: CallbackFailure) -> None: + super().__init__(failure.message) + self.failure = failure + + class _LocalActivityTimedOut(Exception): def __init__(self, kind: str) -> None: super().__init__(f"local activity {kind} timeout elapsed") self.kind = kind +class _CooperativeCancellationObserved(LocalActivityExecutionAborted): + """Return transport observation to the worker, never to authored cleanup.""" + + +class _WorkflowClaimDeferred(LocalActivityExecutionAborted): + """Server released this claim until cancellation acknowledgments resolve.""" + + class _InvalidLocalActivityReport(NonRetryableError): pass @@ -997,6 +1028,20 @@ def __init__( } self.activities = {_activity_name(a): a for a in activities} self.capabilities = tuple(dict.fromkeys(capability.strip() for capability in capabilities)) + self._cooperative_cancellation_supported = False + self._prepared_local_activities_supported = False + self._prepared_local_activity_groups_supported = False + self._local_activity_cancellation_policies: tuple[str, ...] = () + if "prepared_local_activity_cancellation_policies" in self.capabilities and ( + "prepared_local_activities" not in self.capabilities + ): + raise ValueError("prepared local cancellation policies require prepared_local_activities capability") + if "prepared_local_activity_groups" in self.capabilities and ( + "prepared_local_activities" not in self.capabilities + ): + raise ValueError("prepared local groups require prepared_local_activities capability") + if "prepared_local_activities" in self.capabilities and "cooperative_cancellation" not in self.capabilities: + raise ValueError("prepared local activities require cooperative_cancellation capability") if any(not capability for capability in self.capabilities): raise ValueError("worker capabilities must be non-empty strings") self.worker_id = worker_id or f"py-worker-{uuid.uuid4().hex[:8]}" @@ -1016,14 +1061,21 @@ def __init__( if heartbeat_interval <= 0: raise ValueError("heartbeat_interval must be positive") - self._poll_timeout = poll_timeout # Client supplies HTTP grace separately from this requested poll window. - self._poll_http_timeout = poll_timeout + self._poll_timeout = poll_timeout self.max_concurrent_workflow_tasks = max_concurrent_workflow_tasks self.max_concurrent_activity_tasks = max_concurrent_activity_tasks self.max_concurrent_worker_sessions = max_concurrent_worker_sessions self._worker_sessions: dict[str, WorkerSession] = {} self._stop = asyncio.Event() + self._local_activity_shutdown = asyncio.Event() + self._local_activity_executor: ThreadPoolExecutor | None = None + self._remote_activity_executor: ThreadPoolExecutor | None = None + self._remote_activity_threads: set[Future[Any]] = set() + self._remote_activity_processes: set[SupervisedCallback] = set() + self._prepared_local_activity_processes: set[SupervisedCallback] = set() + self._abandoned_prepared_local_claims: set[str] = set() + self._remote_activity_thread_slots = asyncio.Semaphore(max_concurrent_activity_tasks) self._wf_semaphore = asyncio.Semaphore(max_concurrent_workflow_tasks) self._act_semaphore = asyncio.Semaphore(max_concurrent_activity_tasks) self._shutdown_timeout = shutdown_timeout @@ -1154,6 +1206,52 @@ async def _register(self) -> None: raise RuntimeError(f"Server compatibility error: unable to read /api/cluster/info: {e}") from e _validate_server_compatibility(info) + protocol = info.get("worker_protocol") + server_capabilities = protocol.get("server_capabilities") if isinstance(protocol, Mapping) else None + self._cooperative_cancellation_supported = ( + "cooperative_cancellation" in self.capabilities + and isinstance(protocol, Mapping) + and isinstance(server_capabilities, Mapping) + and server_capabilities.get("cooperative_cancellation") is True + and _supports_cooperative_cancellation_protocol(protocol.get("version")) + and _supports_cooperative_cancellation_protocol(_protocol_version_from_env( + "DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", PROTOCOL_VERSION, + )) + ) + if "cooperative_cancellation" in self.capabilities and not self._cooperative_cancellation_supported: + raise RuntimeError("cooperative cancellation requires explicit compatible runtime and worker protocol 1.20") + self._prepared_local_activities_supported = ( + "prepared_local_activities" in self.capabilities and self._cooperative_cancellation_supported + and isinstance(server_capabilities, Mapping) + and server_capabilities.get("prepared_local_activities") is True + ) + if "prepared_local_activities" in self.capabilities and not self._prepared_local_activities_supported: + raise RuntimeError( + "prepared_local_activity_not_supported: Server must advertise its installed admission bridge", + ) + self._prepared_local_activity_groups_supported = ( + "prepared_local_activity_groups" in self.capabilities and self._prepared_local_activities_supported + and isinstance(server_capabilities, Mapping) + and server_capabilities.get("prepared_local_activity_groups") is True + ) + if "prepared_local_activity_groups" in self.capabilities and not self._prepared_local_activity_groups_supported: + raise RuntimeError("prepared_local_group_not_supported: Server must advertise its installed atomic bridge") + policies = ( + server_capabilities.get("prepared_local_activity_cancellation_policies") + if isinstance(server_capabilities, Mapping) else None + ) + self._local_activity_cancellation_policies = tuple( + policy for policy in ("try_cancel", "wait_cancellation_completed") + if isinstance(policies, list) and policy in policies + ) if self._prepared_local_activities_supported else () + if "prepared_local_activity_cancellation_policies" in self.capabilities and ( + not self._local_activity_cancellation_policies + ): + raise RuntimeError( + "prepared_local_activity_cancellation_policy_not_supported: " + "Server must advertise installed prepared policies", + ) + self._validate_cooperative_activity_handlers() self._query_tasks_supported = _server_supports_query_tasks(info) self._workflow_memo_updates_supported = _server_supports_workflow_memo_updates(info) has_update_validators = any( @@ -1184,6 +1282,11 @@ async def _register(self) -> None: capabilities.append(WORKFLOW_UPDATES_CAPABILITY) capabilities.append(MESSAGE_STREAMS_CAPABILITY) capabilities.extend(self.capabilities) + if ( + self._local_activity_cancellation_policies + and "prepared_local_activity_cancellation_policies" not in capabilities + ): + capabilities.append("prepared_local_activity_cancellation_policies") if PORTABLE_WORKER_AFFINITY_CAPABILITY_MANIFEST["worker_sessions"]["supported"]: capabilities.append("worker_sessions") @@ -1199,7 +1302,21 @@ async def _register(self) -> None: max_concurrent_worker_sessions=self.max_concurrent_worker_sessions, build_id=self.build_id, capabilities=capabilities, - capability_manifest=PORTABLE_WORKER_AFFINITY_CAPABILITY_MANIFEST, + capability_manifest={ + **PORTABLE_WORKER_AFFINITY_CAPABILITY_MANIFEST, + **({"prepared_local_activities": { + "supported": True, "minimum_protocol_version": "1.20", + "implementation": "durable_sequential_admission", + }} if self._prepared_local_activities_supported else {}), + **({"prepared_local_activity_groups": { + "supported": True, "minimum_protocol_version": "1.20", + "implementation": "durable_atomic_all_admission", + }} if self._prepared_local_activity_groups_supported else {}), + **({"prepared_local_activity_cancellation_policies": { + "supported": True, "minimum_protocol_version": "1.20", + "implementation": "prepared_local_policy_admission_and_replay", + }} if self._local_activity_cancellation_policies else {}), + }, task_slots=self._current_task_slots(), process_metrics=self._current_process_metrics(), ) @@ -1346,6 +1463,8 @@ async def _resolve_workflow_nexus_commands( raise RuntimeError("workflow yielded too many consecutive Nexus service calls") async def _renew_local_workflow_lease(self, task: dict[str, Any]) -> None: + if self._cooperative_cancellation_supported and self._local_activity_shutdown.is_set(): + raise LocalActivityExecutionAborted("worker shutdown abandoned its local workflow claim") task_id = str(task["task_id"]) attempt = int(task.get("workflow_task_attempt", 1)) try: @@ -1356,6 +1475,8 @@ async def _renew_local_workflow_lease(self, task: dict[str, Any]) -> None: ) except Exception as exc: raise LocalActivityExecutionAborted("workflow task lease renewal failed") from exc + if self._cooperative_cancellation_supported and self._local_activity_shutdown.is_set(): + raise LocalActivityExecutionAborted("worker shutdown abandoned its local workflow claim") if not isinstance(response, Mapping) or any(( response.get("task_id") != task_id, response.get("lease_owner") != self.worker_id, @@ -1363,6 +1484,472 @@ async def _renew_local_workflow_lease(self, task: dict[str, Any]) -> None: response.get("renewed") is not True, )): raise LocalActivityExecutionAborted("workflow task lease renewal was not fenced and acknowledged") + observation = response.get("cancellation_request") + if observation is not None: + self._observe_workflow_cancellation(task, observation) + if task.get("_delivered_cancellation_request_id") != task["cancellation_request"]["request_id"]: + raise _CooperativeCancellationObserved("cooperative request observed on the actual task heartbeat") + + def _observe_workflow_cancellation(self, task: dict[str, Any], observation: Any) -> CancellationRequest: + if not self._cooperative_cancellation_supported or not isinstance(observation, Mapping): + raise LocalActivityExecutionAborted("workflow cancellation observation is not negotiated") + try: + current = CancellationRequest.from_observation(observation) + previous = task.get("cancellation_request") + if previous is not None and not isinstance(previous, Mapping): + raise ValueError("previous observation is not an object") + original = CancellationRequest.from_observation(previous) if previous is not None else None + except Exception as error: + raise LocalActivityExecutionAborted("workflow cancellation observation is malformed") from error + if original is not None and (original.request_id, original.requested_at, original.cleanup_deadline_at) != ( + current.request_id, current.requested_at, current.cleanup_deadline_at, + ): + raise LocalActivityExecutionAborted("workflow cancellation observation changed its original identity") + task["cancellation_request"] = dict(observation) + return current + + async def _load_workflow_claim_history( + self, task: dict[str, Any], *, first_page_token: str | None = None, + ) -> list[dict[str, Any]]: + history = [] if first_page_token is not None else list(task.get("history_events", [])) + token = first_page_token if first_page_token is not None else task.get("next_history_page_token") + seen: set[str] = set() + while token is not None: + if not isinstance(token, str) or not token or token in seen: + raise LocalActivityExecutionAborted("workflow history paging did not advance its opaque token") + seen.add(token) + page = await self.client.workflow_task_history( + task_id=task["task_id"], next_history_page_token=token, + lease_owner=self.worker_id, workflow_task_attempt=task.get("workflow_task_attempt", 1), + ) + if not isinstance(page, Mapping) or not isinstance(page.get("history_events"), list): + raise LocalActivityExecutionAborted("workflow history page was not acknowledged") + if any(not isinstance(event, dict) for event in page["history_events"]): + raise LocalActivityExecutionAborted("workflow history page contains a malformed event") + history.extend(page["history_events"]) + token = page.get("next_history_page_token") + return history + + async def _refresh_cancellation_history( + self, task: dict[str, Any], observed: CancellationRequest, + ) -> list[dict[str, Any]]: + try: + return await self._load_workflow_claim_history(task, first_page_token=observed.history_refresh_page_token) + except Exception as error: + raise LocalActivityExecutionAborted( + "canonical cancellation history could not be loaded on this claim", + ) from error + + async def _replay_workflow_claim( + self, cls: type, task: dict[str, Any], history: list[dict[str, Any]], start_input: list[Any], + *, payload_codec: str | None, execute_local: Callable[[RecordLocalActivity], Any], + ) -> tuple[ReplayOutcome, list[dict[str, Any]]]: + observation = task.get("cancellation_request") + if observation is not None: + observed = self._observe_workflow_cancellation(task, observation) + history = await self._refresh_cancellation_history(task, observed) + delivery_attempts = 0 + for _ in range(1000): + state = read_cancellation_history( + history, run_id=task.get("run_id", ""), observation=task.get("cancellation_request"), + ) + if state.request is not None and not self._cooperative_cancellation_supported: + raise LocalActivityExecutionAborted("canonical cancellation requires a negotiated capable worker") + if state.delivery is not None: + task["_delivered_cancellation_request_id"] = state.delivery.request_id + try: + outcome = await asyncio.to_thread( + replay, cls, history, start_input, workflow_id=task.get("workflow_id"), + run_id=task.get("run_id", ""), + workflow_command_id=( + _string_or_none(task.get("workflow_command_id")) or _string_or_none(task.get("task_id")) + ), + payload_codec=payload_codec, external_storage=self.external_storage, + external_storage_cache=self.external_storage_cache, + cancel_requested=bool(task.get("cancel_requested", False)) and state.request is None, + cancellation_request=task.get("cancellation_request"), local_activity_executor=execute_local, + prepare_local_activities=self._prepared_local_activities_supported, + prepare_local_activity_groups=self._prepared_local_activity_groups_supported, + local_activity_cancellation_policies=self._local_activity_cancellation_policies, + ) + if outcome.prepared_local_activity_group is not None: + history = await self._execute_prepared_local_activity_group(task, history, outcome) + continue + if outcome.prepared_local_activity is not None: + history = await self._execute_prepared_local_activity(task, history, outcome) + continue + except PreparedCancellationObserved as error: + # Control carries the root request time. Workflow observations + # carry this run's local admission time, which can be later in a + # cascade. Read the actual task observation instead of replacing + # either timestamp with the other. + expected = CancellationContext.from_dict(task["_prepared_cancellation_context"]) + with contextlib.suppress(_CooperativeCancellationObserved): + await self._renew_local_workflow_lease(task) + observed = self._observe_workflow_cancellation(task, task.get("cancellation_request")) + if observed.request_id != expected.request_id or datetime.fromisoformat( + observed.cleanup_deadline_at.replace("Z", "+00:00"), + ) != expected.deadline: + raise LocalActivityExecutionAborted( + "prepared stop changed the original request or deadline", + ) from error + history = await self._refresh_cancellation_history(task, observed) + committed = read_cancellation_history(history, run_id=task.get("run_id", ""), + observation=task.get("cancellation_request")) + if committed.request is None or committed.request.context != expected: + raise LocalActivityExecutionAborted( + "prepared stop context differs from canonical request history", + ) from error + continue + except _CooperativeCancellationObserved: + observed = self._observe_workflow_cancellation(task, task.get("cancellation_request")) + history = await self._refresh_cancellation_history(task, observed) + continue + intent = outcome.cancellation_delivery + if intent is None or outcome.commands: + # Earlier authored commands must commit on this claim first. + # Completion releases the claim; a successor replays their durable results. + return outcome, history + delivery_attempts += 1 + if delivery_attempts > 3: + raise LocalActivityExecutionAborted("cancellation replay did not converge on its canonical delivery") + observed = self._observe_workflow_cancellation(task, task.get("cancellation_request")) + delivery_error: Exception | None = None + try: + delivery_reply = await self.client.deliver_workflow_cancellation( + task_id=task["task_id"], lease_owner=self.worker_id, + workflow_task_attempt=task.get("workflow_task_attempt", 1), + request_id=intent.request_id, sequence=intent.sequence, call_kind=intent.call_kind, + sequence_span=intent.sequence_span, operation_sequence=intent.operation_sequence, + operation_sequence_span=intent.operation_sequence_span, + ) + if delivery_reply.get("delivered") is False: + raise _WorkflowClaimDeferred("parent awaits canonical child cleanup on a new claim") + except _WorkflowClaimDeferred: + raise + except Exception as error: + delivery_error = error + history = await self._refresh_cancellation_history(task, observed) + committed = read_cancellation_history( + history, run_id=task.get("run_id", ""), observation=task.get("cancellation_request"), + ) + if committed.delivery != intent: + raise LocalActivityExecutionAborted( + "delivery was not proved by matching canonical history", + ) from delivery_error + raise LocalActivityExecutionAborted("workflow exceeded the prepared local replay admission limit") + + async def _refresh_prepared_local_history( + self, task: dict[str, Any], receipt: Mapping[str, Any], + ) -> list[dict[str, Any]]: + token = receipt.get("history_refresh_page_token") + if not isinstance(token, str) or not token.strip(): + raise LocalActivityExecutionAborted("prepared operation lacks a Server-issued canonical history cursor") + return await self._load_workflow_claim_history(task, first_page_token=token) + + async def _execute_prepared_local_activity( + self, task: dict[str, Any], history: list[dict[str, Any]], outcome: ReplayOutcome, + ) -> list[dict[str, Any]]: + call = outcome.prepared_local_activity + if call is None: + raise LocalActivityExecutionAborted("prepared replay did not identify its authored local call") + task_id = task["task_id"] + epoch = task.get("workflow_task_attempt", 1) + codec = _validate_payload_codec(task.get("payload_codec")) or serializer.AVRO_CODEC + try: + if outcome.commands: + return await self._checkpoint_prepared_prefix(task, outcome.commands, call.sequence, codec) + descriptor = call.descriptor(codec) + if call.recover: + original = next((event.get("payload", {}) for event in reversed(history) + if event.get("event_type", event.get("type")) == "ActivityStarted" + and event.get("payload", {}).get("sequence", event.get("payload", {}).get( + "workflow_sequence")) == call.sequence), None) + receipt = await self.client.prepared_local_activity_operation( + task_id=task_id, lease_owner=self.worker_id, workflow_task_attempt=epoch, + operation="recover", body={"sequence": call.sequence, "descriptor": descriptor}, + ) + kind = receipt.get("event_type") + retry = bool(kind == "ActivityRetryScheduled") + created = receipt.get("created_task_ids") + if (receipt.get("recovered") is not True or type(receipt.get("duplicate")) is not bool + or "reason" not in receipt or receipt["reason"] is not None + or receipt.get("workflow_task_id") != task_id or original is None + or not isinstance(receipt.get("activity_execution_id"), str) + or not receipt["activity_execution_id"].strip() + or not isinstance(receipt.get("activity_attempt_id"), str) + or not receipt["activity_attempt_id"].strip() + or receipt["activity_execution_id"] != original.get("activity_execution_id") + or receipt["activity_attempt_id"] != original.get("activity_attempt_id") + or receipt.get("callback_stop_state") != "unknown" + or kind not in {"ActivityRetryScheduled", "ActivityFailed", "ActivityTimedOut"} + or receipt.get("claim_released") is not retry + or not isinstance(created, list) or len(created) != int(retry) + or any(not isinstance(value, str) or not value.strip() for value in created) + or not isinstance(receipt.get("event_id"), str) or not receipt["event_id"].strip()): + raise LocalActivityExecutionAborted("prepared recovery lacks its canonical unknown-stop receipt") + if retry: + raise _WorkflowClaimDeferred("durable prepared local retry released the original workflow claim") + refreshed = await self._refresh_prepared_local_history(task, receipt) + event = next((event for event in refreshed if event.get("id") == receipt["event_id"]), None) + payload = event.get("payload", {}) if event is not None else {} + recovery = payload.get("local_recovery", {}) + if (event is None or event.get("event_type", event.get("type")) != kind + or payload.get("sequence", payload.get("workflow_sequence")) != call.sequence + or payload.get("activity_execution_id") != receipt["activity_execution_id"] + or payload.get("activity_attempt_id") != receipt["activity_attempt_id"] + or recovery.get("workflow_task_id") != task_id or recovery.get("workflow_task_attempt") != epoch + or recovery.get("lease_owner") != self.worker_id + or recovery.get("callback_stop_state") != "unknown"): + raise LocalActivityExecutionAborted( + "prepared recovery is absent from this claim's canonical history", + ) + return refreshed + runner, invocation = await self._prepare_local_callback(task, call, codec) + attempt = runner.attempt + try: + receipt = await runner.execute(invocation, self._prepared_local_activity_processes) + finally: + if runner.callback is not None and runner.callback in self._prepared_local_activity_processes: + self._abandoned_prepared_local_claims.add(task_id) + if receipt["claim_released"]: + raise _WorkflowClaimDeferred("durable prepared local retry released the original workflow claim") + refreshed = await self._refresh_prepared_local_history(task, receipt) + event = next((event for event in refreshed if event.get("id") == receipt["event_id"]), None) + payload = event.get("payload", {}) if event is not None else {} + if (event is None or event.get("event_type", event.get("type")) != receipt["event_type"] + or payload.get("sequence", payload.get("workflow_sequence")) != call.sequence + or payload.get("activity_execution_id") != attempt.execution_id + or payload.get("activity_attempt_id") != attempt.attempt_id): + raise LocalActivityExecutionAborted("prepared outcome is absent from this claim's canonical history") + return refreshed + except (_WorkflowClaimDeferred, LocalActivityExecutionAborted): + raise + except ServerError as error: + if error.reason() == "cancellation_requested": + await self._renew_local_workflow_lease(task) + raise LocalActivityExecutionAborted("prepared operation has an unknown or refused outcome") from error + except Exception as error: + raise LocalActivityExecutionAborted( + "prepared operation could not validate its original authority", + ) from error + + async def _checkpoint_prepared_prefix( + self, task: dict[str, Any], commands: list[Command], next_sequence: int, codec: str, + ) -> list[dict[str, Any]]: + wire = commands_to_server_commands(commands, self.task_queue, payload_codec=codec) + if any(command["type"] not in { + "record_side_effect", "record_version_marker", "upsert_memo", "upsert_search_attributes", + } for command in wire): + raise LocalActivityExecutionAborted("prepared prefix has no supported authored sequence range") + if any(command["type"] == "upsert_memo" for command in wire) and not self._workflow_memo_updates_supported: + raise LocalActivityExecutionAborted("prepared memo prefix requires negotiated workflow memo updates") + receipt = await self._checkpoint_prepared_commands(task, wire, next_sequence) + return await self._refresh_prepared_local_history(task, receipt) + + async def _checkpoint_prepared_commands( + self, task: dict[str, Any], commands: list[dict[str, Any]], next_sequence: int, *, group: bool = False, + ) -> dict[str, Any]: + task_id = task["task_id"] + epoch = task.get("workflow_task_attempt", 1) + start = next_sequence - len(commands) + if start < 1: + raise LocalActivityExecutionAborted("prepared checkpoint has no authored sequence range") + checkpoint = hashlib.sha256(json.dumps( + [task_id, self.worker_id, epoch, start, commands], sort_keys=True, allow_nan=False, + ).encode()).hexdigest() + receipt = await self.client.prepared_local_activity_operation( + task_id=task_id, lease_owner=self.worker_id, workflow_task_attempt=epoch, + operation="checkpoint-group" if group else "checkpoint", + body={"checkpoint_id": checkpoint, "start_sequence": start, "commands": commands}, + ) + if (receipt.get("checkpointed") is not True or type(receipt.get("duplicate")) is not bool + or receipt.get("checkpoint_id") != checkpoint + or receipt.get("task_id") != task_id or receipt.get("workflow_run_id") != task["run_id"] + or type(receipt.get("workflow_task_attempt")) is not int + or receipt["workflow_task_attempt"] != epoch or receipt.get("lease_owner") != self.worker_id + or type(receipt.get("start_sequence")) is not int or receipt["start_sequence"] != start + or type(receipt.get("next_sequence")) is not int or receipt["next_sequence"] != next_sequence + or "reason" not in receipt or receipt["reason"] is not None): + raise LocalActivityExecutionAborted("prepared checkpoint lacks an original-claim receipt") + return receipt + + async def _prepare_local_callback( + self, task: dict[str, Any], call: PreparedLocalActivityCall, codec: str, + ) -> tuple[PreparedLocalRunner, CallbackInvocation]: + task_id = task["task_id"] + epoch = task.get("workflow_task_attempt", 1) + nonce = hashlib.sha256(json.dumps( + [task_id, task["run_id"], self.worker_id, epoch, call.sequence], + ).encode()).hexdigest() + started = time.monotonic() + receipt = await self.client.prepared_local_activity_operation( + task_id=task_id, lease_owner=self.worker_id, workflow_task_attempt=epoch, operation="prepare", + body={"sequence": call.sequence, "worker_attempt_id": nonce, "descriptor": call.descriptor(codec)}, + ) + attempt = PreparedAttempt.admitted( + receipt, task_id=task_id, run_id=task["run_id"], owner=self.worker_id, epoch=epoch, nonce=nonce, + heartbeat_timeout=call.command.heartbeat_timeout, cleanup=call.cleanup, request_started=started, + ) + handler = self.activities.get(call.command.activity_type) + if handler is None: + raise LocalActivityExecutionAborted("prepared local activity has no registered handler") + + def observe(value: Mapping[str, Any]) -> None: + previous = task.get("_prepared_cancellation_context") + if previous is not None and CancellationContext.from_dict(previous) != CancellationContext.from_dict(value): + raise LocalActivityExecutionAborted("prepared group members observed different cancellation authority") + task["_prepared_cancellation_context"] = dict(value) + + runner = PreparedLocalRunner(self.client, attempt, shutdown=self._local_activity_shutdown, observe=observe) + invocation = CallbackInvocation( + handler, tuple(call.command.arguments), ActivityInfo( + task_id, call.command.activity_type, attempt.attempt_id, attempt.attempt_number, + self.task_queue, self.worker_id, + ), {**task, "activity_attempt_id": attempt.attempt_id}, self.interceptors, + ) + return runner, invocation + + async def _execute_prepared_local_activity_group( + self, task: dict[str, Any], history: list[dict[str, Any]], outcome: ReplayOutcome, + ) -> list[dict[str, Any]]: + group = outcome.prepared_local_activity_group + if group is None or not self._prepared_local_activity_groups_supported: + raise LocalActivityExecutionAborted("prepared group lacks negotiated atomic admission") + codec = _validate_payload_codec(task.get("payload_codec")) or serializer.AVRO_CODEC + members: list[tuple[PreparedLocalRunner, CallbackInvocation]] = [] + executions: dict[asyncio.Task[dict[str, Any]], int] = {} + receipts: dict[int, dict[str, Any]] = {} + try: + if outcome.commands: + return await self._checkpoint_prepared_prefix(task, outcome.commands, group.base_sequence, codec) + if not group.committed: + calls = {call.sequence: call for call in group.calls} + commands: list[dict[str, Any]] = [] + for sequence, command in enumerate(group.commands, group.base_sequence): + if sequence in calls: + commands.append({**calls[sequence].descriptor(codec), "type": "prepare_local_activity"}) + else: + commands.extend(commands_to_server_commands([command], self.task_queue, payload_codec=codec)) + receipt = await self._checkpoint_prepared_commands( + task, commands, group.base_sequence + group.size, group=True, + ) + locals_ = receipt.get("local_activities") + identities: set[str] = set() + executions_by_sequence: dict[int, str] = {} + if not isinstance(locals_, list) or len(locals_) != len(group.calls): + raise LocalActivityExecutionAborted("atomic checkpoint omitted an authored local member") + for call, local in zip(group.calls, locals_, strict=True): + identity = local.get("activity_execution_id") if isinstance(local, Mapping) else None + if (not isinstance(local, Mapping) or type(local.get("sequence")) is not int + or local["sequence"] != call.sequence or not isinstance(identity, str) + or not identity.strip() or identity in identities): + raise LocalActivityExecutionAborted("atomic checkpoint changed its authored local members") + identities.add(identity) + executions_by_sequence[call.sequence] = identity + # Replay the entire canonical batch before preparing any callback. + refreshed = await self._refresh_prepared_local_history(task, receipt) + opening_kinds = { + "prepare_local_activity": {"ActivityScheduled"}, "schedule_activity": {"ActivityScheduled"}, + "start_timer": {"TimerScheduled"}, + "start_child_workflow": {"ChildWorkflowScheduled", "ChildRunStarted"}, + } + for sequence, wire_command in enumerate(commands, group.base_sequence): + opening = next((event for event in refreshed + if event.get("event_type", event.get("type")) in opening_kinds[wire_command["type"]] + and event.get("payload", {}).get("sequence", event.get("payload", {}).get( + "workflow_sequence")) == sequence), None) + payload = opening.get("payload", {}) if opening is not None else {} + if opening is None or payload.get("parallel_group_path") != wire_command["parallel_group_path"]: + raise LocalActivityExecutionAborted("atomic checkpoint is absent from canonical group history") + if sequence in calls and payload.get("activity_execution_id") != executions_by_sequence[sequence]: + raise LocalActivityExecutionAborted( + "atomic checkpoint changed canonical local execution identity", + ) + return refreshed + for call in group.calls: + if call.recover: + return await self._execute_prepared_local_activity( + task, history, ReplayOutcome([], prepared_local_activity=call), + ) + for call in group.calls: + members.append(await self._prepare_local_callback(task, call, codec)) + for index, (runner, invocation) in enumerate(members): + execution = asyncio.create_task(runner.execute(invocation, self._prepared_local_activity_processes)) + executions[execution] = index + pending = set(executions) + while pending: + done, pending = await asyncio.wait(pending, return_when=asyncio.FIRST_COMPLETED) + for execution in done: + receipt = await execution + if receipt["claim_released"]: + raise _WorkflowClaimDeferred("Native released the group claim for a durable local retry") + receipts[executions[execution]] = receipt + # A cursor from a later-settled receipt includes all prior commits. + last = next(reversed(receipts.values())) + refreshed = await self._refresh_prepared_local_history(task, last) + for index, receipt in receipts.items(): + call = group.calls[index] + attempt = members[index][0].attempt + event = next((event for event in refreshed if event.get("id") == receipt["event_id"]), None) + payload = event.get("payload", {}) if event is not None else {} + if (event is None or event.get("event_type", event.get("type")) != receipt["event_type"] + or payload.get("sequence", payload.get("workflow_sequence")) != call.sequence + or payload.get("activity_execution_id") != attempt.execution_id + or payload.get("activity_attempt_id") != attempt.attempt_id): + raise LocalActivityExecutionAborted("prepared group outcome is absent from canonical history") + return refreshed + except PreparedCancellationObserved: + # Stop every sibling first, including admitted members that did not + # spawn. Only then report their canonical cancellation fences. + await self._join_prepared_group(executions, members, task) + for execution, index in executions.items(): + if not execution.cancelled() and execution.exception() is None: + receipts[index] = execution.result() + for index, (runner, _) in enumerate(members): + if index in receipts or runner.stop_acknowledged: + continue + try: + await runner.control() + except PreparedCancellationObserved: + await runner.acknowledge_stop() + continue + raise LocalActivityExecutionAborted( + "joined group member lacks its canonical cancellation fence", + ) from None + raise + except ServerError as error: + await self._join_prepared_group(executions, members, task) + if error.reason() == "cancellation_requested": + # No group process can start until all admissions have validated. + for runner, _ in members: + try: + await runner.control() + except PreparedCancellationObserved: + await runner.acknowledge_stop() + await self._renew_local_workflow_lease(task) + raise LocalActivityExecutionAborted( + f"prepared group operation was refused: {error.reason() or 'unknown'}", + ) from error + except (_WorkflowClaimDeferred, LocalActivityExecutionAborted): + raise + except Exception as error: + raise LocalActivityExecutionAborted("prepared group could not validate original authority") from error + finally: + await self._join_prepared_group(executions, members, task) + + async def _join_prepared_group( + self, executions: Mapping[asyncio.Task[dict[str, Any]], int], + members: list[tuple[PreparedLocalRunner, CallbackInvocation]], task: dict[str, Any], + ) -> None: + for execution in executions: + if not execution.done(): + execution.cancel() + await asyncio.gather(*executions, return_exceptions=True) + if any(runner.callback is not None and runner.callback in self._prepared_local_activity_processes + for runner, _ in members): + self._abandoned_prepared_local_claims.add(task["task_id"]) + raise LocalActivityExecutionAborted("prepared group callback stop remains unconfirmed") def _maybe_externalize_local_payload(self, envelope: dict[str, str]) -> dict[str, Any]: storage = self.external_storage @@ -1375,6 +1962,48 @@ def _maybe_externalize_local_payload(self, envelope: dict[str, str]) -> dict[str reference = store_external_payload(storage, data, codec=envelope["codec"]) return {"codec": envelope["codec"], "external_storage": reference.to_dict()} + async def _execute_cooperative_local_callable( + self, task: dict[str, Any], command: RecordLocalActivity, handler: Callable[..., Any], + attempt_state: dict[str, Any], + ) -> Any: + async def observe_lease() -> None: + while True: + await asyncio.sleep(min(5.0, self._heartbeat_interval)) + await self._renew_local_workflow_lease(task) + + invocation = asyncio.create_task(self._execute_activity_callable( + task, command.activity_type, tuple(command.arguments), handler, run_sync_in_thread=True, + )) + observation = asyncio.create_task(observe_lease()) + shutdown = asyncio.create_task(self._local_activity_shutdown.wait()) + try: + done, _ = await asyncio.wait([invocation, observation, shutdown], return_when=asyncio.FIRST_COMPLETED) + if shutdown in done: + raise LocalActivityExecutionAborted("worker shutdown abandoned its local workflow claim") + if observation in done: + await observation # Propagate transport observation or lost lease, never a workflow cancellation. + raise LocalActivityExecutionAborted("local lease observer stopped without an acknowledgment") + return await invocation + finally: + shutdown.cancel() + with contextlib.suppress(asyncio.CancelledError): + await shutdown + observation.cancel() + with contextlib.suppress(asyncio.CancelledError, Exception): + await observation + if not invocation.done(): + attempt_state["lease_aborted"] = True + invocation.cancel() + + def discard_late_result(future: asyncio.Task[Any]) -> None: + if not future.cancelled(): + future.exception() + + # Python cannot forcibly stop a synchronous thread or a callable + # that suppresses cancellation. Its attempt is fenced and its + # eventual result cannot become a durable command. + invocation.add_done_callback(discard_late_result) + async def _execute_local_activity( self, task: dict[str, Any], @@ -1452,7 +2081,13 @@ def check_boundary( now = time.monotonic() if state["lease_aborted"]: raise LocalActivityExecutionAborted("local activity lost its workflow task lease") - if self._stop.is_set() or task.get("cancel_requested") is True: + if self._cooperative_cancellation_supported and self._local_activity_shutdown.is_set(): + state["lease_aborted"] = True + raise LocalActivityExecutionAborted("worker shutdown abandoned its local workflow claim") + if (self._stop.is_set() and not self._cooperative_cancellation_supported) or ( + task.get("cancel_requested") is True and task.get("cancellation_request") is None + and task.get("_delivered_cancellation_request_id") is None + ): raise ActivityCancelled("local activity cancelled") if ( command.heartbeat_timeout is not None @@ -1515,11 +2150,16 @@ async def heartbeat( ) _set_context(ActivityContext(info=info, client=self.client, heartbeat_callback=heartbeat)) try: - result = await self._execute_activity_callable( - task, command.activity_type, tuple(command.arguments), handler, - ) + if self._cooperative_cancellation_supported: + result = await self._execute_cooperative_local_callable(task, command, handler, attempt_state) + else: + result = await self._execute_activity_callable( + task, command.activity_type, tuple(command.arguments), handler, + ) finally: _set_context(None) + if self._cooperative_cancellation_supported: + await self._renew_local_workflow_lease(task) check_boundary() attempts.append({ "attempt_id": attempt_id, @@ -1610,21 +2250,7 @@ async def _run_workflow_task_core(self, task: dict[str, Any]) -> list[dict[str, task_id: str = task["task_id"] attempt: int = task.get("workflow_task_attempt", 1) wf_type: str = task.get("workflow_type", "") - history = task.get("history_events", []) - - # The worker requests bounded history pages when polling. Do not replay - # an incomplete history if fetching a later page fails. - next_page_token = task.get("next_history_page_token") - while next_page_token: - page_data = await self.client.workflow_task_history( - task_id=task_id, - next_history_page_token=next_page_token, - lease_owner=self.worker_id, - workflow_task_attempt=attempt, - ) - if page_data and page_data.get("history_events"): - history.extend(page_data["history_events"]) - next_page_token = page_data.get("next_history_page_token") if page_data else None + history = await self._load_workflow_claim_history(task) start_input: list[Any] = [] codec = task.get("payload_codec") @@ -1703,6 +2329,9 @@ async def _run_workflow_task_core(self, task: dict[str, Any]) -> list[dict[str, payload_codec=codec, external_storage=self.external_storage, external_storage_cache=self.external_storage_cache, + prepare_local_activities=self._prepared_local_activities_supported, + prepare_local_activity_groups=self._prepared_local_activity_groups_supported, + local_activity_cancellation_policies=self._local_activity_cancellation_policies, ) command = update_command.to_server_command( self.task_queue, @@ -1747,23 +2376,12 @@ def execute_local(command: RecordLocalActivity) -> Any: return future.result() try: - outcome = await asyncio.to_thread( - replay, - cls, - history, - start_input, - workflow_id=task.get("workflow_id"), - run_id=run_id, - workflow_command_id=( - _string_or_none(task.get("workflow_command_id")) - or _string_or_none(task.get("task_id")) - ), - payload_codec=codec, - external_storage=self.external_storage, - external_storage_cache=self.external_storage_cache, - cancel_requested=bool(task.get("cancel_requested", False)), - local_activity_executor=execute_local, + outcome, history = await self._replay_workflow_claim( + cls, task, history, start_input, payload_codec=codec, execute_local=execute_local, ) + except _WorkflowClaimDeferred: + log.info("workflow task %s released for durable successor work", task_id) + return None except LocalActivityExecutionAborted as e: log.warning("abandoning workflow task %s before local activity commit: %s", task_id, e) return None @@ -1859,6 +2477,43 @@ def execute_local(command: RecordLocalActivity) -> Any: log.warning("failed to report workflow memo capability failure: %s", failure_error) return None + if not self._cooperative_cancellation_supported and any( + isinstance(command, ScheduleActivity) and command.cancellation_policy is not None + for command in workflow_commands + ): + message = ( + f"activity_cancellation_policy_not_supported: Python worker {self.worker_id} requires " + "cooperative_cancellation capability, worker protocol 1.20 and a compatible Server/Native backend" + ) + try: + await self.client.fail_workflow_task( + task_id=task_id, lease_owner=self.worker_id, workflow_task_attempt=attempt, + message=message, failure_type="RuntimeCapabilityUnsupported", + ) + except Exception as failure_error: + log.warning("failed to report activity cancellation policy capability failure: %s", failure_error) + return None + + if not self._cooperative_cancellation_supported and any( + isinstance(command, StartChildWorkflow) and ( + command.parent_close_policy == "request_cancellation" + or command.cancellation_policy in ("try_cancel", "wait_cancellation_completed") + ) + for command in workflow_commands + ): + message = ( + f"child_cancellation_policy_not_supported: Python worker {self.worker_id} requires " + "cooperative_cancellation capability, worker protocol 1.20 and a compatible Server/Native backend" + ) + try: + await self.client.fail_workflow_task( + task_id=task_id, lease_owner=self.worker_id, workflow_task_attempt=attempt, + message=message, failure_type="RuntimeCapabilityUnsupported", + ) + except Exception as failure_error: + log.warning("failed to report child cancellation policy capability failure: %s", failure_error) + return None + def serialize_commands(source: list[Command]) -> list[dict[str, Any]]: return commands_to_server_commands( source, @@ -1979,6 +2634,167 @@ async def _report_workflow_task_after_completion_error( fail_error, ) + def _validate_cooperative_activity_handlers(self) -> None: + if not self._cooperative_cancellation_supported: + return + for activity_type, handler in self.activities.items(): + try: + pickle.dumps((handler, self.interceptors), protocol=pickle.HIGHEST_PROTOCOL) + except Exception as error: + raise RuntimeError( + f"worker {self.worker_id!r} cannot supervise cooperative activity {activity_type!r}: " + "handler and interceptors must be importable and spawn-compatible" + ) from error + + async def _acknowledge_stopped_remote_activity(self, task: dict[str, Any]) -> None: + # The caller has joined both the callback and its supervisor. Observation + # binds only this original claim to canonical cancellation metadata. + task_id = task["task_id"] + attempt_id = task.get("activity_attempt_id") or task.get("attempt_id", "") + try: + status = await asyncio.wait_for(self.client.activity_task_status( + task_id=task_id, activity_attempt_id=attempt_id, lease_owner=self.worker_id, + ), timeout=5.0) + receipt = status.get("cancellation_acknowledgement") if isinstance(status, Mapping) else None + if (not isinstance(status, Mapping) or status.get("task_id") != task_id + or status.get("activity_attempt_id") != attempt_id or status.get("lease_owner") != self.worker_id + or status.get("can_continue") is not False or status.get("cancel_requested") is not True + or status.get("heartbeat_recorded") is not False or not isinstance(receipt, Mapping) + or receipt.get("callback_state") not in ("unknown", "stopped")): + return + for field in ("request_id", "root_request_id", "cleanup_deadline_at", "cancellation_history_event_id"): + if not isinstance(receipt.get(field), str) or not receipt[field].strip(): + return + reply = await asyncio.wait_for(self.client.acknowledge_activity_cancellation( + task_id=task_id, activity_attempt_id=attempt_id, lease_owner=self.worker_id, + request_id=receipt["request_id"], + ), timeout=5.0) + if (not isinstance(reply, Mapping) or reply.get("task_id") != task_id + or reply.get("activity_attempt_id") != attempt_id or reply.get("lease_owner") != self.worker_id + or reply.get("request_id") != receipt["request_id"] or reply.get("acknowledged") is not True + or not isinstance(reply.get("duplicate"), bool) or reply.get("reason") is not None + or reply.get("heartbeat_recorded") is not False + or not isinstance(reply.get("history_event_id"), str) or not reply["history_event_id"].strip()): + raise _RemoteActivityExecutionAborted("Server did not prove the original callback-stop receipt") + log.info("remote activity %s callback stopped for request %s, root %s, deadline %s", + task_id, receipt["request_id"], receipt["root_request_id"], receipt["cleanup_deadline_at"]) + except Exception as error: + log.warning("remote activity %s stopped but its cancellation receipt failed: %s", task_id, error) + + async def _assert_remote_activity_claim(self, task: dict[str, Any]) -> None: + try: + if self._local_activity_shutdown.is_set(): + raise _RemoteActivityExecutionAborted("worker shutdown abandoned its remote activity claim") + reply = await self.client.activity_task_status( + task_id=task["task_id"], + activity_attempt_id=task.get("activity_attempt_id") or task.get("attempt_id", ""), + lease_owner=self.worker_id, + ) + if (not isinstance(reply, Mapping) + or reply.get("task_id") != task["task_id"] + or reply.get("activity_attempt_id") != ( + task.get("activity_attempt_id") or task.get("attempt_id", "")) + or reply.get("lease_owner") != self.worker_id + or reply.get("can_continue") is not True + or reply.get("cancel_requested") is not False + or reply.get("heartbeat_recorded") is not False + or reply.get("reason") is not None): + raise _RemoteActivityExecutionAborted("remote activity observation refused its ownership fence") + bounds = [reply.get("lease_expires_at")] + deadlines = reply.get("deadlines") + if deadlines is not None: + if not isinstance(deadlines, Mapping): + raise _RemoteActivityExecutionAborted("remote activity observation returned invalid deadlines") + bounds.extend(deadlines[kind] for kind in ("heartbeat", "start_to_close", "schedule_to_close") + if deadlines.get(kind) is not None) + session = reply.get("worker_session") + if session is not None: + if (not isinstance(session, Mapping) or session.get("status") != "active" + or session.get("lease_owner") != self.worker_id): + raise _RemoteActivityExecutionAborted("remote activity observation lost its worker session") + bounds.extend([session.get("lease_expires_at"), session.get("ttl_expires_at")]) + for bound in bounds: + if not isinstance(bound, str) or "T" not in bound: + raise _RemoteActivityExecutionAborted("remote activity observation returned an invalid deadline") + deadline = datetime.fromisoformat(bound.replace("Z", "+00:00")) + if deadline.tzinfo is None or deadline <= datetime.now(timezone.utc): + raise _RemoteActivityExecutionAborted("remote activity ownership or execution deadline elapsed") + if self._local_activity_shutdown.is_set(): + raise _RemoteActivityExecutionAborted("worker shutdown abandoned its remote activity claim") + except _RemoteActivityExecutionAborted: + raise + except Exception as error: + raise _RemoteActivityExecutionAborted("remote activity ownership observation failed") from error + + async def _execute_cooperative_remote_callable( + self, task: dict[str, Any], handler: Callable[..., Any], args: tuple[Any, ...], info: ActivityInfo, + ) -> Any: + async def send_heartbeat(details: dict[str, Any] | None) -> None: + await self._assert_remote_activity_claim(task) + try: + reply = await asyncio.wait_for(self.client.heartbeat_activity_task( + task_id=info.task_id, activity_attempt_id=info.activity_attempt_id, + lease_owner=self.worker_id, details=details, + ), timeout=5.0) + if (not isinstance(reply, Mapping) or reply.get("task_id") != info.task_id + or reply.get("activity_attempt_id") != info.activity_attempt_id + or reply.get("lease_owner") != self.worker_id or reply.get("can_continue") is not True + or reply.get("cancel_requested") is not False): + raise _RemoteActivityExecutionAborted("remote activity heartbeat lost its ownership fence") + except Exception as error: + raise _RemoteActivityExecutionAborted("remote activity user heartbeat failed") from error + await self._assert_remote_activity_claim(task) + + async def observe_ownership() -> None: + while True: + await asyncio.sleep(1.0) + await self._assert_remote_activity_claim(task) + + await self._assert_remote_activity_claim(task) + callback = SupervisedCallback(CallbackInvocation(handler, args, info, task, self.interceptors)) + self._remote_activity_processes.add(callback) + invocation: asyncio.Task[Any] | None = None + abandoned = False + observation = asyncio.create_task(observe_ownership()) + shutdown = asyncio.create_task(self._local_activity_shutdown.wait()) + try: + await callback.start() + invocation = asyncio.create_task(callback.result(send_heartbeat)) + done, _ = await asyncio.wait([invocation, observation, shutdown], return_when=asyncio.FIRST_COMPLETED) + if shutdown in done: + raise _RemoteActivityExecutionAborted("worker shutdown abandoned its remote activity claim") + if observation in done: + await observation + raise _RemoteActivityExecutionAborted("remote ownership observer stopped without a response") + try: + outcome = await invocation + except CallbackProcessLost as error: + raise _RemoteActivityExecutionAborted("remote callback supervisor lost stop authority") from error + await self._assert_remote_activity_claim(task) + if outcome.failure is not None: + raise _RemoteActivityApplicationFailure(outcome.failure) + return outcome.value + except (CallbackProcessLost, _RemoteActivityExecutionAborted, asyncio.CancelledError): + abandoned = True + raise + finally: + backgrounds = [background for background in (invocation, observation, shutdown) if background is not None] + for background in backgrounds: + if not background.done(): + background.cancel() + await asyncio.gather(*backgrounds, return_exceptions=True) + try: + if not callback.stopped: + await callback.stop() + except Exception as error: + log.warning("remote activity %s callback stop remains unconfirmed: %s", info.task_id, error) + if callback.stopped: + self._remote_activity_processes.discard(callback) + if abandoned: + await self._acknowledge_stopped_remote_activity(task) + # An unconfirmed callback retains its capacity and never emits a + # stop receipt, result or failure. Owner shutdown closes its channel. + async def _run_activity_task(self, task: dict[str, Any]) -> str: self._track_worker_session_from_task(task) task_id: str = task["task_id"] @@ -2065,7 +2881,30 @@ async def _run_activity_task(self, task: dict[str, Any]) -> str: ) _set_context(act_ctx) try: - result = await self._execute_activity_callable(task, activity_type, tuple(args), fn) + if self._cooperative_cancellation_supported: + result = await self._execute_cooperative_remote_callable(task, fn, tuple(args), act_ctx.info) + else: + result = await self._execute_activity_callable(task, activity_type, tuple(args), fn) + except _RemoteActivityExecutionAborted as error: + log.warning("remote activity %s claim abandoned: %s", task_id, error) + return "claim_aborted" + except CallbackProcessLost as error: + log.warning("remote activity %s supervision failed: %s", task_id, error) + return "claim_aborted" + except _RemoteActivityApplicationFailure as error: + failure = error.failure + try: + await self.client.fail_activity_task( + task_id=task_id, activity_attempt_id=attempt_id, lease_owner=self.worker_id, + message="activity cancelled" if failure.cancelled else failure.message, + failure_type=failure.failure_type, failure_class=failure.failure_class, + failure_code=failure.failure_code, stack_trace=failure.stack_trace, + non_retryable=failure.non_retryable or failure.cancelled, + ) + except Exception as report_error: + log.warning("failed to report remote activity failure: %s", report_error) + return ("cancelled" if failure.cancelled else + "failed_non_retryable" if failure.non_retryable else "failed") except ActivityCancelled: log.info("activity %s cancelled via heartbeat", task_id) try: @@ -2141,6 +2980,7 @@ async def _execute_activity_callable( activity_type: str, args: tuple[Any, ...], fn: Callable[..., Any], + *, run_sync_in_thread: bool = False, remote: bool = False, ) -> Any: context = ActivityInterceptorContext( worker_id=self.worker_id, @@ -2151,7 +2991,43 @@ async def _execute_activity_callable( ) async def call_activity(ctx: ActivityInterceptorContext) -> Any: - result = fn(*ctx.args) + if run_sync_in_thread and not inspect.iscoroutinefunction(fn): + if remote: + await self._remote_activity_thread_slots.acquire() + if self._remote_activity_executor is None: + self._remote_activity_executor = ThreadPoolExecutor( + max_workers=self.max_concurrent_activity_tasks, thread_name_prefix="dw-remote-activity", + ) + owner_loop = asyncio.get_running_loop() + try: + future = self._remote_activity_executor.submit(contextvars.copy_context().run, fn, *ctx.args) + except BaseException: + self._remote_activity_thread_slots.release() + raise + self._remote_activity_threads.add(future) + + def thread_finished(completed: Future[Any]) -> None: + def release_capacity() -> None: + self._remote_activity_threads.discard(completed) + self._remote_activity_thread_slots.release() + with contextlib.suppress(RuntimeError): + owner_loop.call_soon_threadsafe(release_capacity) + + # Cancellation of the await cannot release a running thread's + # slot. Only actual completion or cancellation before start can. + future.add_done_callback(thread_finished) + result = await asyncio.wrap_future(future) + else: + if self._local_activity_executor is None: + # Replay threads wait for local results, so sharing their pool can deadlock. + self._local_activity_executor = ThreadPoolExecutor( + max_workers=self.max_concurrent_workflow_tasks, thread_name_prefix="dw-local-activity", + ) + result = await asyncio.get_running_loop().run_in_executor( + self._local_activity_executor, contextvars.copy_context().run, fn, *ctx.args, + ) + else: + result = fn(*ctx.args) if asyncio.iscoroutine(result): return await result return result @@ -2276,6 +3152,9 @@ async def _run_query_task_core(self, task: dict[str, Any], *, client: Client | N payload_codec=codec, external_storage=self.external_storage, external_storage_cache=self.external_storage_cache, + prepare_local_activities=self._prepared_local_activities_supported, + prepare_local_activity_groups=self._prepared_local_activity_groups_supported, + local_activity_cancellation_policies=self._local_activity_cancellation_policies, ) if inspect.isawaitable(result): result = await result @@ -2490,7 +3369,10 @@ def _admit_workflow_work( # The callback owns the reservation so cancellation releases capacity # even when the dispatch coroutine is cancelled before it first runs. - dispatched.add_done_callback(lambda _: self._release_workflow_capacity()) + dispatched.add_done_callback(lambda _: ( + self._release_workflow_capacity() + if task_kind != "workflow" or task.get("task_id") not in self._abandoned_prepared_local_claims else None + )) return dispatched async def _poll_workflow_tasks(self) -> None: @@ -2504,7 +3386,7 @@ async def _poll_workflow_tasks(self) -> None: task = await self.client.poll_workflow_task( worker_id=self.worker_id, task_queue=self.task_queue, - timeout=self._poll_http_timeout, + timeout=self._poll_timeout, build_id=self.build_id, task_kinds=self._workflow_poll_task_kinds(), history_page_size=WORKFLOW_HISTORY_PAGE_SIZE, @@ -2601,6 +3483,11 @@ async def _report_unhandled_workflow_task_error( async def _poll_activity_tasks(self) -> None: while not self._stop.is_set(): + self._validate_cooperative_activity_handlers() + if (len(self._remote_activity_threads) + len(self._remote_activity_processes) + >= self.max_concurrent_activity_tasks): + await asyncio.sleep(0.1) + continue await self._act_semaphore.acquire() if self._stop.is_set(): self._act_semaphore.release() @@ -2610,7 +3497,7 @@ async def _poll_activity_tasks(self) -> None: task = await self.client.poll_activity_task( worker_id=self.worker_id, task_queue=self.task_queue, - timeout=self._poll_http_timeout, + timeout=self._poll_timeout, build_id=self.build_id, ) except asyncio.CancelledError: @@ -2669,7 +3556,7 @@ async def _poll_query_tasks(self, *, client: Client | None = None, track_tasks: task = await client.poll_query_task( worker_id=self.worker_id, task_queue=self.task_queue, - timeout=self._poll_http_timeout, + timeout=self._poll_timeout, build_id=self.build_id, ) except Exception as e: @@ -2759,6 +3646,9 @@ async def _run_update_validation_task(self, task: dict[str, Any]) -> str: payload_codec=codec, external_storage=self.external_storage, external_storage_cache=self.external_storage_cache, + prepare_local_activities=self._prepared_local_activities_supported, + prepare_local_activity_groups=self._prepared_local_activity_groups_supported, + local_activity_cancellation_policies=self._local_activity_cancellation_policies, ) if inspect.isawaitable(result): result = await result @@ -3035,7 +3925,9 @@ def _current_task_slots(self) -> dict[str, int]: 0, self.max_concurrent_workflow_tasks - self._workflow_reserved ), "activity_available": max( - 0, self.max_concurrent_activity_tasks - self._activity_inflight + 0, min(self.max_concurrent_activity_tasks - self._activity_inflight, + self.max_concurrent_activity_tasks - len(self._remote_activity_threads) + - len(self._remote_activity_processes)), ), "session_available": max( 0, self.max_concurrent_worker_sessions @@ -3194,7 +4086,7 @@ async def _run_until_loop( task = await self.client.poll_workflow_task( worker_id=self.worker_id, task_queue=self.task_queue, - timeout=self._poll_http_timeout, + timeout=self._poll_timeout, build_id=self.build_id, task_kinds=self._workflow_poll_task_kinds(), history_page_size=WORKFLOW_HISTORY_PAGE_SIZE, @@ -3239,11 +4131,16 @@ async def _run_until_loop( next_task_kind = "workflow" continue + self._validate_cooperative_activity_handlers() + if len(self._remote_activity_processes) >= self.max_concurrent_activity_tasks: + next_task_kind = "workflow" + await asyncio.sleep(poll_interval) + continue poll_start = time.perf_counter() task = await self.client.poll_activity_task( worker_id=self.worker_id, task_queue=self.task_queue, - timeout=self._poll_http_timeout, + timeout=self._poll_timeout, build_id=self.build_id, ) if self._stop.is_set(): @@ -3348,12 +4245,19 @@ async def _shutdown(self) -> None: in_flight, timeout=self._remaining_shutdown_time(deadline), ) + if pending: + self._local_activity_shutdown.set() for t in pending: t.cancel() if pending: log.warning("cancelled %d task(s) after shutdown timeout", len(pending)) await asyncio.sleep(0) still_running = {task for task in pending if not task.done()} + if still_running and (self._remote_activity_processes or self._prepared_local_activity_processes): + # Application drain has ended. Allow bounded process reaping and + # stop-report transport without granting more callback work or + # extending the run's original cancellation deadline. + _, still_running = await asyncio.wait(still_running, timeout=15.0) if still_running: raise RuntimeError( "worker shutdown timed out while cancelling " @@ -3362,6 +4266,18 @@ async def _shutdown(self) -> None: ) await asyncio.gather(*in_flight, return_exceptions=True) + if self._local_activity_executor is not None: + self._local_activity_executor.shutdown(wait=False, cancel_futures=True) + if self._remote_activity_executor is not None: + self._remote_activity_executor.shutdown(wait=False, cancel_futures=True) + + if self._remote_activity_processes or self._prepared_local_activity_processes: + kind = "remote" if self._remote_activity_processes else "prepared local" + raise RuntimeError( + f"worker shutdown has unconfirmed {kind} callback stop(s); " + "the worker registration remains active" + ) + for session in self._worker_sessions.values(): if not session.active: continue diff --git a/src/durable_workflow/workflow.py b/src/durable_workflow/workflow.py index 4810f43..986a04f 100644 --- a/src/durable_workflow/workflow.py +++ b/src/durable_workflow/workflow.py @@ -26,13 +26,24 @@ import random import re import uuid +import weakref from collections.abc import Callable, Generator, Iterable, Mapping, Sequence +from contextvars import ContextVar from copy import copy from dataclasses import dataclass, field from datetime import datetime, timezone from typing import Any, TypeVar, cast from . import serializer +from ._cooperative_cancellation import CancellationDelivery, read_cancellation_history +from .cancellation import ( + CancellationContext, + CancellationPolicy, + ParentClosePolicy, + _canonical_activity_policy, + _canonical_child_policies, + _timestamp, +) from .client import WorkflowStreamAppendItem from .errors import ( ActivityFailed, @@ -469,6 +480,7 @@ class ScheduleActivity: schedule_to_close_timeout: int | None = None heartbeat_timeout: int | None = None worker_session: WorkerSessionOptions | None = None + cancellation_policy: str | CancellationPolicy | None = None _parallel_group_path: list[dict[str, Any]] | None = field( default=None, init=False, @@ -476,6 +488,18 @@ class ScheduleActivity: compare=False, ) + def __post_init__(self) -> None: + self._validate_cancellation_policy() + + def _validate_cancellation_policy(self) -> None: + self.cancellation_policy = _canonical_activity_policy(self.cancellation_policy) + if self.cancellation_policy == CancellationPolicy.ABANDON.value and ( + not isinstance(self.schedule_to_close_timeout, int) + or isinstance(self.schedule_to_close_timeout, bool) + or self.schedule_to_close_timeout < 1 + ): + raise ValueError("remote activity Abandon requires a finite positive schedule_to_close_timeout") + def to_server_command( self, task_queue: str, @@ -486,6 +510,7 @@ def to_server_command( external_storage: ExternalStorageDriver | None = None, external_storage_threshold_bytes: int | None = None, ) -> dict[str, Any]: + self._validate_cancellation_policy() self._validate_timeouts() command: dict[str, Any] = { @@ -522,6 +547,8 @@ def to_server_command( command["heartbeat_timeout"] = self.heartbeat_timeout if self.worker_session is not None: command["worker_session"] = self.worker_session.to_wire() + if self.cancellation_policy is not None: + command["cancellation_policy"] = _canonical_activity_policy(self.cancellation_policy) return command def _validate_timeouts(self) -> None: @@ -567,11 +594,19 @@ class RecordLocalActivity: start_to_close_timeout: int | None = None schedule_to_close_timeout: int | None = None heartbeat_timeout: int | None = None + cancellation_policy: str | CancellationPolicy | None = None outcome: dict[str, Any] | None = field(default=None, init=False, repr=False) arguments_envelope: dict[str, Any] | None = field(default=None, init=False, repr=False) result_envelope: dict[str, Any] | None = field(default=None, init=False, repr=False) + _parallel_group_path: list[dict[str, Any]] | None = field(default=None, init=False, repr=False, compare=False) def __post_init__(self) -> None: + try: + self.cancellation_policy = _canonical_activity_policy(self.cancellation_policy) + except ValueError as error: + raise ValueError("local activity cancellation_policy must be a supported prepared policy") from error + if self.cancellation_policy == CancellationPolicy.ABANDON.value: + raise ValueError("local activity abandon is not supported by prepared callback ownership") if not isinstance(self.activity_type, str): raise TypeError("local activity type must be a string") self.activity_type = self.activity_type.strip() @@ -605,6 +640,11 @@ def to_server_command( size_warning: serializer.PayloadSizeWarningConfig | None = serializer.DEFAULT_PAYLOAD_SIZE_WARNING, warning_context: PayloadWarningContext = None, ) -> dict[str, Any]: + if self.cancellation_policy is not None: + raise LocalActivityExecutionAborted( + "prepared_local_activity_cancellation_policy_not_supported: " + "explicit policies require prepared admission", + ) if self.outcome is None or self.arguments_envelope is None: raise ValueError("local activity has no recorded terminal outcome") status = self.outcome.get("outcome") @@ -629,6 +669,61 @@ def to_server_command( return command +@dataclass(frozen=True) +class PreparedLocalActivityCall: + """An authored local call awaiting admission, with no application execution.""" + + command: RecordLocalActivity + sequence: int + recover: bool = False + cleanup: Mapping[str, str] | None = None + + def descriptor(self, payload_codec: str) -> dict[str, Any]: + command = self.command + descriptor: dict[str, Any] = { + "type": "record_local_activity", "activity_type": command.activity_type, + "execution_mode": "local", "arguments": serializer.envelope(command.arguments, codec=payload_codec), + "payload_codec": payload_codec, + } + if command.retry_policy is not None: + descriptor["retry_policy"] = dict(command.retry_policy) + if command.cancellation_policy is not None: + descriptor["cancellation_policy"] = command.cancellation_policy + for name in ("start_to_close_timeout", "schedule_to_close_timeout", "heartbeat_timeout"): + value = getattr(command, name) + if value is not None: + descriptor[name] = value + if self.cleanup is not None: + descriptor["cancellation_cleanup"] = { + "request_id": self.cleanup["request_id"], + "delivery_history_event_id": self.cleanup["delivery_history_event_id"], + } + _apply_parallel_group_metadata(command, descriptor) + return descriptor + + +@dataclass(frozen=True) +class PreparedLocalActivityGroup: + """A complete ordinary parallel group, with only unresolved local calls.""" + + base_sequence: int + size: int + commands: tuple[Command, ...] + calls: tuple[PreparedLocalActivityCall, ...] + committed: bool + + def __post_init__(self) -> None: + sequences = [call.sequence for call in self.calls] + if ( + self.base_sequence < 1 or not 1 <= self.size <= 100 or not self.calls + or len(self.commands) != (0 if self.committed else self.size) + or len(set(sequences)) != len(sequences) + or any(not self.base_sequence <= call.sequence < self.base_sequence + self.size + or not call.command._parallel_group_path for call in self.calls) + ): + raise LocalActivityExecutionAborted("prepared group lacks complete bounded authored membership") + + class LocalActivityExecutionAborted(Exception): """The workflow task lease could not be trusted after local execution began.""" @@ -916,10 +1011,11 @@ class StartChildWorkflow: workflow_type: str arguments: list[Any] = field(default_factory=list) task_queue: str | None = None - parent_close_policy: str | None = None + parent_close_policy: str | ParentClosePolicy | None = None retry_policy: ChildWorkflowRetryPolicyInput | None = None execution_timeout_seconds: int | None = None run_timeout_seconds: int | None = None + cancellation_policy: str | CancellationPolicy | None = None _parallel_group_path: list[dict[str, Any]] | None = field( default=None, init=False, @@ -927,6 +1023,12 @@ class StartChildWorkflow: compare=False, ) + def __post_init__(self) -> None: + _canonical_child_policies({ + "parent_close_policy": self.parent_close_policy, + "cancellation_policy": self.cancellation_policy, + }) + def to_server_command( self, task_queue: str, @@ -959,8 +1061,10 @@ def to_server_command( cmd["queue"] = self.task_queue else: cmd["queue"] = task_queue - if self.parent_close_policy is not None: - cmd["parent_close_policy"] = self.parent_close_policy + cmd.update(_canonical_child_policies({ + "parent_close_policy": self.parent_close_policy, + "cancellation_policy": self.cancellation_policy, + })) if self.retry_policy is not None: cmd["retry_policy"] = ( self.retry_policy.to_dict() @@ -1373,6 +1477,7 @@ def commands_to_server_commands( for command in commands: if isinstance(command, ScheduleActivity): + command._validate_cancellation_policy() queue = command.queue or task_queue server_command: dict[str, Any] = { "type": "schedule_activity", @@ -1406,6 +1511,8 @@ def commands_to_server_commands( server_command["heartbeat_timeout"] = command.heartbeat_timeout if command.worker_session is not None: server_command["worker_session"] = command.worker_session.to_wire() + if command.cancellation_policy is not None: + server_command["cancellation_policy"] = _canonical_activity_policy(command.cancellation_policy) server_commands.append(server_command) continue @@ -1489,8 +1596,10 @@ def commands_to_server_commands( task_queue=queue, ), )) - if command.parent_close_policy is not None: - server_command["parent_close_policy"] = command.parent_close_policy + server_command.update(_canonical_child_policies({ + "parent_close_policy": command.parent_close_policy, + "cancellation_policy": command.cancellation_policy, + })) if command.retry_policy is not None: server_command["retry_policy"] = ( command.retry_policy.to_dict() @@ -1680,16 +1789,17 @@ def compensate(self, initiating_failure: BaseException) -> Any: if self._closed: raise RuntimeError("a saga instance can only compensate once") self._closed = True - for compensation in reversed(self._compensations): - try: - yield compensation.command - except (Exception, WorkflowCancelled) as compensation_failure: - raise SagaCompensationFailed( - initiating_failure, - compensation_failure, - compensation_activity_type=compensation.command.activity_type, - compensation_registration_order=compensation.registration_order, - ) from compensation_failure + with self._context.cancellation_shield(): + for compensation in reversed(self._compensations): + try: + yield compensation.command + except (Exception, WorkflowCancelled) as compensation_failure: + raise SagaCompensationFailed( + initiating_failure, + compensation_failure, + compensation_activity_type=compensation.command.activity_type, + compensation_registration_order=compensation.registration_order, + ) from compensation_failure def run(self, forward: Callable[[Saga], Any]) -> Any: """Run ``forward`` and compensate if it fails or is cancelled.""" @@ -1708,6 +1818,9 @@ def run(self, forward: Callable[[Saga], Any]) -> Any: raise +_ACTIVE_WORKFLOW_REPLAY: ContextVar[WorkflowContext | None] = ContextVar("active_workflow_replay", default=None) + + class WorkflowContext: """Replay-safe helper surface passed to workflow ``run`` methods.""" @@ -1727,6 +1840,11 @@ def __init__( self._current_time = current_time or datetime.now(timezone.utc) self._workflow_command_id = workflow_command_id or run_id or workflow_id self._cancel_requested = bool(cancel_requested) + self._cancellation_request_id: str | None = None + self._cancellation_context: CancellationContext | None = None + self._cancellation_replay_time: datetime | None = None + self._cancellation_replay_time_available = False + self._cancellation_shield_depth = 0 seed = int(hashlib.sha256(run_id.encode()).hexdigest()[:16], 16) self._rng = random.Random(seed) self._uuid7_counter = 0 @@ -1810,17 +1928,66 @@ def _accept_message_stream(self, arguments: list[Any]) -> None: @property def is_cancellation_requested(self) -> bool: - """Whether this task carries a cooperative cancellation request. + """Whether cancellation has reached this replay's authored boundary. - Server's current ``/cancel`` route is terminal and does not set this - flag. Service-mode cooperative cancellation is not yet available. + Transport observation alone does not set this flag. Server's terminal + ``/cancel`` route does not deliver a cooperative workflow exception. """ return self._cancel_requested def throw_if_cancellation_requested(self) -> None: """Raise :class:`WorkflowCancelled` at an explicit safe point.""" - if self._cancel_requested: - raise WorkflowCancelled("workflow cancellation was requested") + if self._cancel_requested and self._cancellation_shield_depth == 0: + raise WorkflowCancelled( + "workflow cancellation was requested", request_id=self._cancellation_request_id, + context=self._cancellation_context, + ) + + @property + def cancellation_context(self) -> CancellationContext | None: + """Original metadata, visible only at committed cancellation delivery.""" + return self._cancellation_context + + def _observe_cancellation_replay_time(self, event: Mapping[str, Any] | None) -> None: + timestamp = event.get("timestamp") if event is not None else None + if timestamp is None and event is not None: + timestamp = event.get("recorded_at") + self._cancellation_replay_time_available = False + if not isinstance(timestamp, str): + return + try: + recorded_time = _timestamp(timestamp) + except ValueError: + return + if self._cancellation_replay_time is None or recorded_time > self._cancellation_replay_time: + self._cancellation_replay_time = recorded_time + self._cancellation_replay_time_available = True + + def _bind_cancellation_context(self, context: CancellationContext) -> CancellationContext: + reference = weakref.ref(self) + + def replay_time() -> datetime: + active = reference() + if active is None or _ACTIVE_WORKFLOW_REPLAY.get() is not active: + raise RuntimeError("cancellation remaining time requires active workflow replay") + if not active._cancellation_replay_time_available or active._cancellation_replay_time is None: + raise RuntimeError("cancellation replay boundary requires a recorded timestamp") + return active._cancellation_replay_time + + return context._with_replay_clock(replay_time) + + @contextlib.contextmanager + def cancellation_shield(self) -> Generator[None, None, None]: + """Permit cleanup calls without delivering the same request again. + + The Server still enforces the original cleanup deadline, task lease + and termination fence. This scope does not extend any of them. + """ + self._cancellation_shield_depth += 1 + try: + yield + finally: + self._cancellation_shield_depth -= 1 def saga(self) -> Saga: """Create a deterministic reverse-order compensation scope.""" @@ -1838,6 +2005,7 @@ def schedule_activity( schedule_to_close_timeout: int | None = None, heartbeat_timeout: int | None = None, worker_session: WorkerSessionOptions | None = None, + cancellation_policy: str | CancellationPolicy | None = None, ) -> ScheduleActivity: return ScheduleActivity( activity_type=activity_type, @@ -1849,6 +2017,7 @@ def schedule_activity( schedule_to_close_timeout=schedule_to_close_timeout, heartbeat_timeout=heartbeat_timeout, worker_session=worker_session, + cancellation_policy=cancellation_policy, ) def local_activity( @@ -1860,6 +2029,7 @@ def local_activity( start_to_close_timeout: int | None = None, schedule_to_close_timeout: int | None = None, heartbeat_timeout: int | None = None, + cancellation_policy: str | CancellationPolicy | None = None, ) -> RecordLocalActivity: """Yield an activity executed in this workflow worker process.""" return RecordLocalActivity( @@ -1869,6 +2039,7 @@ def local_activity( start_to_close_timeout=start_to_close_timeout, schedule_to_close_timeout=schedule_to_close_timeout, heartbeat_timeout=heartbeat_timeout, + cancellation_policy=cancellation_policy, ) def start_timer(self, seconds: int) -> StartTimer: @@ -2053,7 +2224,8 @@ def start_child_workflow( arguments: list[Any] | None = None, *, task_queue: str | None = None, - parent_close_policy: str | None = None, + parent_close_policy: str | ParentClosePolicy | None = None, + cancellation_policy: str | CancellationPolicy | None = None, retry_policy: ChildWorkflowRetryPolicyInput | None = None, execution_timeout_seconds: int | None = None, run_timeout_seconds: int | None = None, @@ -2063,6 +2235,7 @@ def start_child_workflow( arguments=list(arguments) if arguments is not None else [], task_queue=task_queue, parent_close_policy=parent_close_policy, + cancellation_policy=cancellation_policy, retry_policy=retry_policy, execution_timeout_seconds=execution_timeout_seconds, run_timeout_seconds=run_timeout_seconds, @@ -2241,6 +2414,9 @@ class ReplayOutcome: commands: list[Command] message_stream_cursors: list[dict[str, Any]] = field(default_factory=list) message_stream_waits: list[dict[str, Any]] = field(default_factory=list) + cancellation_delivery: CancellationDelivery | None = None + prepared_local_activity: PreparedLocalActivityCall | None = None + prepared_local_activity_group: PreparedLocalActivityGroup | None = None class Replayer: @@ -2328,6 +2504,7 @@ class _PendingReceiver: name: str args: list[Any] condition_wait_id: str | None = None + after_cancellation_delivery: bool = False @dataclass(frozen=True) @@ -2336,6 +2513,7 @@ class _RecordedStep: shape: str event_types: list[str] details: dict[str, Any] + history_index: int | None = None @dataclass(frozen=True) @@ -2611,7 +2789,11 @@ def replay( external_storage: ExternalStorageDriver | None = None, external_storage_cache: ExternalPayloadCache | None = None, cancel_requested: bool = False, + cancellation_request: Mapping[str, Any] | None = None, local_activity_executor: Callable[[RecordLocalActivity], Any] | None = None, + prepare_local_activities: bool = False, + prepare_local_activity_groups: bool = False, + local_activity_cancellation_policies: tuple[str, ...] = (), ) -> ReplayOutcome: return _replay_state( workflow_cls, @@ -2624,7 +2806,11 @@ def replay( external_storage=external_storage, external_storage_cache=external_storage_cache, cancel_requested=cancel_requested, + cancellation_request=cancellation_request, local_activity_executor=local_activity_executor, + prepare_local_activities=prepare_local_activities, + prepare_local_activity_groups=prepare_local_activity_groups, + local_activity_cancellation_policies=local_activity_cancellation_policies, ).outcome @@ -2640,6 +2826,9 @@ def query_state( payload_codec: str | None = None, external_storage: ExternalStorageDriver | None = None, external_storage_cache: ExternalPayloadCache | None = None, + prepare_local_activities: bool = False, + prepare_local_activity_groups: bool = False, + local_activity_cancellation_policies: tuple[str, ...] = (), ) -> Any: """Replay a workflow to current state and invoke a registered query. @@ -2662,6 +2851,9 @@ def query_state( external_storage=external_storage, external_storage_cache=external_storage_cache, stop_at_uncommitted_cancellation=True, + prepare_local_activities=prepare_local_activities, + prepare_local_activity_groups=prepare_local_activity_groups, + local_activity_cancellation_policies=local_activity_cancellation_policies, ) except Exception as exc: raise QueryFailed(f"workflow replay failed before query: {exc}") from exc @@ -2697,6 +2889,9 @@ def apply_update( payload_codec: str | None = None, external_storage: ExternalStorageDriver | None = None, external_storage_cache: ExternalPayloadCache | None = None, + prepare_local_activities: bool = False, + prepare_local_activity_groups: bool = False, + local_activity_cancellation_policies: tuple[str, ...] = (), ) -> CompleteUpdate | FailUpdate: """Replay current workflow state and run one accepted update handler. @@ -2717,6 +2912,9 @@ def apply_update( payload_codec=payload_codec, external_storage=external_storage, external_storage_cache=external_storage_cache, + prepare_local_activities=prepare_local_activities, + prepare_local_activity_groups=prepare_local_activity_groups, + local_activity_cancellation_policies=local_activity_cancellation_policies, ) except Exception as exc: return _fail_update_from_exception( @@ -2798,6 +2996,9 @@ def validate_update( payload_codec: str | None = None, external_storage: ExternalStorageDriver | None = None, external_storage_cache: ExternalPayloadCache | None = None, + prepare_local_activities: bool = False, + prepare_local_activity_groups: bool = False, + local_activity_cancellation_policies: tuple[str, ...] = (), ) -> Any: """Replay state and invoke only the declared pre-accept update validator. @@ -2818,6 +3019,9 @@ def validate_update( payload_codec=payload_codec, external_storage=external_storage, external_storage_cache=external_storage_cache, + prepare_local_activities=prepare_local_activities, + prepare_local_activity_groups=prepare_local_activity_groups, + local_activity_cancellation_policies=local_activity_cancellation_policies, ) except Exception as exc: raise UpdateValidationFailed( @@ -3306,6 +3510,8 @@ def _recorded_step_details(payload: Mapping[str, Any]) -> dict[str, Any]: for key in ( "workflow_type", "child_workflow_type", + "parent_close_policy", + "cancellation_policy", "timer_kind", "change_id", "condition_key", @@ -3439,6 +3645,13 @@ def _recorded_detail_mismatch(command: Any, step: _RecordedStep) -> str | None: f"Recorded activity_type {recorded!r}, but current workflow " f"scheduled {command.activity_type!r}." ) + actual_policy = _canonical_activity_policy(command.cancellation_policy) or "try_cancel" + recorded_policy = step.details.get("cancellation_policy", "try_cancel") + if recorded_policy != actual_policy: + return ( + f"activity_cancellation_policy_changed: recorded {recorded_policy!r}, " + f"but current workflow requested {actual_policy!r}." + ) elif isinstance(command, RecordLocalActivity): if step.details.get("execution_mode") != "local": return "Recorded remote activity cannot replay as a local activity." @@ -3448,6 +3661,13 @@ def _recorded_detail_mismatch(command: Any, step: _RecordedStep) -> str | None: f"Recorded local activity_type {recorded!r}, but current workflow " f"requested {command.activity_type!r}." ) + actual_policy = command.cancellation_policy or "try_cancel" + recorded_policy = step.details.get("cancellation_policy", "try_cancel") + if recorded_policy != actual_policy: + return ( + f"local_activity_cancellation_policy_changed: recorded {recorded_policy!r}, " + f"but current workflow requested {actual_policy!r}." + ) elif isinstance(command, StartChildWorkflow): recorded = step.details.get("workflow_type") or step.details.get("child_workflow_type") if isinstance(recorded, str) and recorded != command.workflow_type: @@ -3455,6 +3675,20 @@ def _recorded_detail_mismatch(command: Any, step: _RecordedStep) -> str | None: f"Recorded child workflow_type {recorded!r}, but current workflow " f"started {command.workflow_type!r}." ) + actual_policies = { + "parent_close_policy": "abandon", "cancellation_policy": "abandon", + **_canonical_child_policies({ + "parent_close_policy": command.parent_close_policy, + "cancellation_policy": command.cancellation_policy, + }), + } + for field, actual_policy in actual_policies.items(): + recorded_policy = step.details.get(field, "abandon") + if recorded_policy != actual_policy: + return ( + f"child_workflow_policy_changed: recorded {field} {recorded_policy!r}, " + f"but current workflow requested {actual_policy!r}." + ) elif isinstance(command, RecordVersionMarker): recorded = step.details.get("change_id") if isinstance(recorded, str) and recorded != command.change_id: @@ -3594,6 +3828,26 @@ def _first_yield_failure(values: Iterable[Any]) -> ActivityFailed | ChildWorkflo return None +def _assert_cancellation_scope_replay_supported(events: list[dict[str, Any]]) -> None: + """Refuse unqualified scope execution before constructing application code.""" + for event in events: + unsupported = _history_event_type(event) in { + "CancellationScopeOpened", + "CancellationScopeRequested", + "CancellationScopeDelivered", + "CancellationScopeRequestConflicted", + } + payload = event.get("payload") + if isinstance(payload, Mapping): + for container in (payload, *(payload.get(name) for name in ("activity", "timer", "child_workflow"))): + if isinstance(container, Mapping) and "cancellation_scope_id" in container: + unsupported = unsupported or container["cancellation_scope_id"] != "root" + if unsupported: + raise LocalActivityExecutionAborted( + "cancellation_scope_execution_not_supported: Python worker cannot replay scoped cancellation history", + ) + + def _replay_state( workflow_cls: type, history_events: Iterable[dict[str, Any]], @@ -3606,9 +3860,15 @@ def _replay_state( external_storage: ExternalStorageDriver | None = None, external_storage_cache: ExternalPayloadCache | None = None, cancel_requested: bool = False, + cancellation_request: Mapping[str, Any] | None = None, local_activity_executor: Callable[[RecordLocalActivity], Any] | None = None, + prepare_local_activities: bool = False, + prepare_local_activity_groups: bool = False, + local_activity_cancellation_policies: tuple[str, ...] = (), stop_at_uncommitted_cancellation: bool = False, ) -> _ReplayState: + if prepare_local_activity_groups and not prepare_local_activities: + raise LocalActivityExecutionAborted("prepared groups require prepared local activity admission") if payload_codec is not None and payload_codec != serializer.AVRO_CODEC: try: serializer.decode("", codec=payload_codec) @@ -3622,9 +3882,17 @@ def _replay_state( ) from exception events = list(history_events) + _assert_cancellation_scope_replay_supported(events) + cancellation = read_cancellation_history(events, run_id=run_id, observation=cancellation_request) + cancellation_consumed = False + cancellation_intent: CancellationDelivery | None = None + authored_sequence = 1 + current_call_sequence = 1 workflow_id = workflow_id or _workflow_id_from_history(events) event_types_by_sequence: dict[int, list[str]] = {} details_by_sequence: dict[int, dict[str, Any]] = {} + child_policies_by_sequence: dict[int, dict[str, str]] = {} + activity_policies_by_sequence: dict[int, str] = {} resolved_sequences: set[int] = set() condition_wait_ids_by_sequence: dict[int, str] = {} selected_condition_wait_ids_by_sequence: dict[int, str] = {} @@ -3646,6 +3914,52 @@ def _replay_state( # envelopes into the public non-determinism diagnostic. recorded_details = {} details_by_sequence.setdefault(sequence, {}).update(recorded_details) + if event_type in ( + "ActivityScheduled", "ActivityStarted", "ActivityCompleted", + "ActivityFailed", "ActivityTimedOut", "ActivityCancelled", + ): + policy = activity_policies_by_sequence.get(sequence) + snapshot = payload.get("activity") + for source in (payload, snapshot if isinstance(snapshot, Mapping) else {}): + if "cancellation_policy" not in source: + continue + incoming = source["cancellation_policy"] + if not isinstance(incoming, str) or incoming not in tuple(item.value for item in CancellationPolicy): + raise NonDeterministicReplayError( + sequence, "activity", [event_type], + detail="invalid_activity_cancellation_policy_history: expected a supported policy.", + ) + if policy is not None and policy != incoming: + raise NonDeterministicReplayError( + sequence, "activity", [event_type], + detail="activity_cancellation_policy_history_conflict: policy changed between events.", + ) + policy = incoming + activity_policies_by_sequence[sequence] = policy or "try_cancel" + details_by_sequence[sequence]["cancellation_policy"] = activity_policies_by_sequence[sequence] + if event_type in ( + "ChildWorkflowScheduled", "ChildRunStarted", "ChildRunCompleted", + "ChildRunFailed", "ChildRunCancelled", "ChildRunTerminated", + ): + try: + incoming_policies = _canonical_child_policies(payload) + except ValueError as error: + raise NonDeterministicReplayError( + sequence, "child workflow", [event_type], + detail=f"invalid_child_workflow_policy_history: {error}", + ) from error + previous_policies = child_policies_by_sequence.get(sequence) + for field, value in incoming_policies.items(): + if previous_policies is not None and previous_policies[field] != value: + raise NonDeterministicReplayError( + sequence, "child workflow", [event_type], + detail=f"child_workflow_policy_history_conflict: {field} changed between history events.", + ) + child_policies_by_sequence[sequence] = { + "parent_close_policy": "abandon", "cancellation_policy": "abandon", + **(previous_policies or {}), **incoming_policies, + } + details_by_sequence[sequence].update(child_policies_by_sequence[sequence]) if event_type == "ConditionWaitOpened": wait_id = payload.get("condition_wait_id") if isinstance(wait_id, str) and wait_id: @@ -3721,10 +4035,16 @@ def _replay_state( ): ctx._accept_message_stream(arguments) - def _state(commands: list[Command]) -> _ReplayState: + def _state( + commands: list[Command], prepared_local_activity: PreparedLocalActivityCall | None = None, + prepared_local_activity_group: PreparedLocalActivityGroup | None = None, + ) -> _ReplayState: return _ReplayState( outcome=ReplayOutcome( commands=commands, + cancellation_delivery=cancellation_intent, + prepared_local_activity=prepared_local_activity, + prepared_local_activity_group=prepared_local_activity_group, message_stream_cursors=[ {"stream_name": name, "through_position": position} for name, position in sorted(ctx._message_stream_cursors.items()) @@ -3769,15 +4089,18 @@ def _state(commands: list[Command]) -> _ReplayState: # loser complete before or after unrelated successor work without changing # replay binding. selection_resolutions: dict[int, Any] = {} + selection_resolution_indexes: dict[int, int] = {} selection_steps: dict[int, list[_RecordedStep]] = {} selection_payloads: dict[int, list[dict[str, Any]]] = {} selection_opening_payloads: dict[int, dict[str, Any]] = {} selection_condition_terminal_payloads: dict[int, list[dict[str, Any]]] = {} selection_condition_sequences: dict[str, int] = {} selection_markers: list[dict[str, Any]] = [] + selection_marker_indexes: list[int] = [] selection_marker_cursor = 0 consumed_selection_sequences: set[int] = set() cancelled_selection_members: dict[tuple[str, int], dict[str, Any]] = {} + cancelled_selection_indexes: dict[tuple[str, int], int] = {} validated_selection_cancellations: set[tuple[str, int]] = set() authored_selection_handles: dict[int, DurableOperationHandle] = {} @@ -3800,9 +4123,11 @@ def _append_resolved_result(value: Any, shape: str, event: Mapping[str, Any]) -> shape=shape, event_types=event_types, details=details, + history_index=event_index, ) if _selection_group_metadata(payload) is not None: selection_resolutions[workflow_sequence] = value + selection_resolution_indexes[workflow_sequence] = event_index _append_selection_step(shape, event, fallback_sequence=workflow_sequence) return resolved_results.append(value) @@ -3832,11 +4157,16 @@ def _recorded_step( else {} ) details.update(_recorded_step_details(payload)) + if shape == "child workflow": + details.update(child_policies_by_sequence.get(workflow_sequence, {})) + if shape == "activity": + details["cancellation_policy"] = activity_policies_by_sequence.get(workflow_sequence, "try_cancel") return _RecordedStep( workflow_sequence=workflow_sequence, shape=shape, event_types=event_types, details=details, + history_index=event_index, ) def _append_selection_step( @@ -3898,6 +4228,7 @@ def _append_selection_condition_resolution( fallback_sequence=selection_sequence, ) selection_resolutions[selection_sequence] = value + selection_resolution_indexes[selection_sequence] = event_index return True def _assert_step_matches(command: Any, step: _RecordedStep) -> None: @@ -3924,6 +4255,15 @@ def _assert_next_step_matches(command: Any, offset: int = 0) -> None: step_index = result_cursor + offset if step_index >= len(recorded_steps): return + if ( + (cancellation.request is not None or prepare_local_activities) + and recorded_steps[step_index].workflow_sequence != current_call_sequence + offset + ): + raise NonDeterministicReplayError( + current_call_sequence + offset, _command_diagnostic_shape(command), + recorded_steps[step_index].event_types, + detail="recorded result belongs to a different authored call", + ) _assert_step_matches(command, recorded_steps[step_index]) def _unconsumed_recorded_steps() -> list[_RecordedStep]: @@ -3949,6 +4289,11 @@ def _assert_pending_step_matches(command: Any, offset: int = 0) -> None: _assert_step_matches(command, candidates[offset]) def _assert_no_unconsumed_history(terminal_shape: str) -> None: + if cancellation.delivery is not None and not cancellation_consumed: + raise NonDeterministicReplayError( + cancellation.delivery.sequence, terminal_shape, ["CooperativeCancellationDelivered"], + detail="workflow did not reach the committed cancellation boundary", + ) step = _next_unconsumed_recorded_step() if step is not None: raise NonDeterministicReplayError( @@ -4007,6 +4352,8 @@ def _is_external_receiver_event(event_type: str | None) -> bool: return event_type in ("SignalReceived", "UpdateApplied") def _receiver_binding_boundary_kind(event_type: str | None, payload: Mapping[str, Any]) -> str | None: + if event_type == "CooperativeCancellationDelivered": + return "condition" if payload.get("call_kind") == "condition" else "step" if event_type in ( "ConditionWaitSatisfied", "ConditionWaitTimedOut", @@ -4154,6 +4501,25 @@ def _receiver_condition_wait_bindings() -> dict[int, str | None]: raw_payload = ev.get("payload") or {} payload = dict(raw_payload) if isinstance(raw_payload, Mapping) else {} sequence = _workflow_sequence(payload) + if ( + cancellation.delivery is not None and sequence is not None + and cancellation.delivery.interrupts(sequence) + and etype in { + "ActivityScheduled", "ActivityStarted", "ActivityCompleted", "ActivityFailed", "ActivityTimedOut", + "ActivityCancelled", "TimerScheduled", "TimerFired", "TimerCancelled", "ChildWorkflowScheduled", + "ChildRunStarted", "ChildRunCompleted", "ChildRunFailed", "ChildRunCancelled", "ChildRunTerminated", + "ConditionWaitOpened", "ConditionWaitSatisfied", "ConditionWaitTimedOut", "SignalWaitOpened", + "SignalApplied", + } + and not (cancellation.delivery.call_kind == "selection_handle" and etype in { + "ActivityScheduled", "ActivityStarted", "TimerScheduled", "ChildWorkflowScheduled", + "ChildRunStarted", "ConditionWaitOpened", "SignalWaitOpened", + }) + and not (cancellation.delivery.call_kind == "condition" and etype == "ConditionWaitOpened") + ): + # The delivery marker owns the interruption. Its activity/timer + # terminal rows are not ordinary workflow results or cleanup calls. + continue if ( etype in { @@ -4243,6 +4609,7 @@ def _receiver_condition_wait_bindings() -> dict[int, str | None]: elif etype == "SelectionResolved": if isinstance(payload.get("selection_group_id"), str): selection_markers.append(dict(payload)) + selection_marker_indexes.append(event_index) elif etype == "SelectionOperationCancelled": group_id = payload.get("selection_group_id") member_base = _parallel_group_integer(payload, "member_base_sequence") @@ -4258,6 +4625,7 @@ def _receiver_condition_wait_bindings() -> dict[int, str | None]: detail="conflicting cancellation markers target the same durable member", ) cancelled_selection_members[cancellation_key] = cancellation_payload + cancelled_selection_indexes.setdefault(cancellation_key, event_index) elif etype in ("SideEffectRecorded", "ChildRunCompleted"): shape = "side effect" if etype == "SideEffectRecorded" else "child workflow" _append_resolved_result( @@ -4365,6 +4733,9 @@ def _receiver_condition_wait_bindings() -> dict[int, str | None]: external_storage_cache=external_storage_cache, ), condition_wait_id=condition_wait_id, + after_cancellation_delivery=( + cancellation.delivery_index is not None and event_index > cancellation.delivery_index + ), )) elif etype == "UpdateApplied": update_name = payload.get("update_name") @@ -4393,6 +4764,9 @@ def _receiver_condition_wait_bindings() -> dict[int, str | None]: external_storage_cache=external_storage_cache, ), condition_wait_id=condition_wait_id, + after_cancellation_delivery=( + cancellation.delivery_index is not None and event_index > cancellation.delivery_index + ), )) resolved_results, recorded_steps = _reorder_completed_parallel_groups( @@ -4428,6 +4802,8 @@ def _apply_receiver(receiver: _PendingReceiver) -> None: handler(*receiver.args) def _receiver_due(receiver: _PendingReceiver, *, before_consuming_result: bool) -> bool: + if receiver.after_cancellation_delivery and not cancellation_consumed: + return False if receiver.condition_wait_id is not None: return False if before_consuming_result: @@ -4546,7 +4922,7 @@ def _same_logical_condition_wait(opened: Mapping[str, Any], cmd: WaitCondition) terminal_condition_reopen_cmd: WaitCondition | None = None def _parallel_leaf_kind(command: Any) -> str: - if isinstance(command, ScheduleActivity): + if isinstance(command, ScheduleActivity | RecordLocalActivity): return "activity" if isinstance(command, StartTimer): return "timer" @@ -4608,6 +4984,8 @@ def _annotate_parallel_commands( candidates = _unconsumed_recorded_steps() if base_sequence_override is not None: base_sequence = base_sequence_override + elif cancellation.request is not None: + base_sequence = current_call_sequence elif candidates: base_sequence = candidates[0].workflow_sequence else: @@ -4911,6 +5289,8 @@ def _annotate_selection( "SelectionResolved with a positive group base sequence", ["SelectionResolved"], ) + elif cancellation.request is not None: + base_sequence = current_call_sequence elif candidates: base_sequence = candidates[0].workflow_sequence elif any(sequence not in consumed_selection_sequences for sequence in selection_steps): @@ -5034,6 +5414,199 @@ def _consume_terminal_condition_reopens() -> None: break wait_yield_count += 1 + def _cancellation_boundary(command: Any) -> CancellationDelivery | None: + nonlocal authored_sequence, current_call_sequence + if cancellation.request is None and not prepare_local_activities: + return None + current_call_sequence = authored_sequence + kind: str | None = None + span = 1 + operation_sequence: int | None = None + operation_span = 1 + if isinstance(command, ScheduleActivity): + kind = "activity" + elif isinstance(command, RecordLocalActivity): + kind = "local_activity" + elif isinstance(command, StartTimer): + kind = "timer" + elif isinstance(command, StartChildWorkflow): + kind = "child" + elif isinstance(command, WaitCondition): + has_recorded_open = any( + _history_event_type(event) == "ConditionWaitOpened" + and _workflow_sequence(event.get("payload") or {}) == current_call_sequence + for event in events + ) + if has_recorded_open or ( + cancellation.delivery is not None and cancellation.delivery.sequence == current_call_sequence + ) or not command.predicate(): + kind = "condition" + elif isinstance(command, list): + kind = "parallel" + span = len(_parallel_leaves(command)) + elif isinstance(command, SelectGroup): + kind = "parallel" + span = sum(len(_parallel_leaves(operation)) if isinstance(operation, list) else 1 + for _, operation in command.operations) + elif isinstance(command, DurableOperationHandle): + kind = "selection_handle" + operation_sequence = command.base_sequence + operation_span = command.size + if kind != "selection_handle" and (kind is not None or isinstance( + command, RecordSideEffect | RecordVersionMarker | UpsertMemo | UpsertSearchAttributes | NexusServiceCall, + )): + authored_sequence += span + if kind is None or span == 0 or cancellation.request is None: + return None + return CancellationDelivery( + cancellation.request.request_id, current_call_sequence, kind, span, + operation_sequence, operation_span, + ) + + def _contains_local_activity(operation: Any) -> bool: + if isinstance(operation, RecordLocalActivity): + return True + if isinstance(operation, list): + return any(_contains_local_activity(member) for member in operation) + if isinstance(operation, SelectGroup): + return any(_contains_local_activity(member) for _, member in operation.operations) + return False + + def _validate_local_activity_policies(operation: Any) -> None: + if isinstance(operation, RecordLocalActivity) and operation.cancellation_policy is not None: + if ( + not prepare_local_activities + or operation.cancellation_policy not in local_activity_cancellation_policies + ): + raise LocalActivityExecutionAborted( + "prepared_local_activity_cancellation_policy_not_supported: " + f"{operation.cancellation_policy} requires installed Server discovery and a prepared consumer", + ) + elif isinstance(operation, list): + for member in operation: + _validate_local_activity_policies(member) + elif isinstance(operation, SelectGroup): + for _, member in operation.operations: + _validate_local_activity_policies(member) + + def _prepared_call(command: RecordLocalActivity, sequence: int) -> PreparedLocalActivityCall: + cleanup: dict[str, str] | None = None + if cancellation.request is not None: + delivery = cancellation.delivery + context = cancellation.request.context + delivery_id = ( + events[cancellation.delivery_index].get("id") + if cancellation.delivery_index is not None else None + ) + if ( + not cancellation_consumed or ctx._cancellation_shield_depth < 1 + or delivery is None or context is None + or sequence < delivery.sequence + delivery.sequence_span + or not isinstance(delivery_id, str) or not delivery_id.strip() + ): + raise LocalActivityExecutionAborted( + "prepared cleanup requires a shield after canonical cancellation delivery", + ) + cleanup = { + "request_id": context.request_id, "root_request_id": context.root_request_id, + "delivery_history_event_id": delivery_id, + "cleanup_deadline_at": context.to_dict()["cleanup_deadline_at"], + } + started = False + for event in events: + if _workflow_sequence(event.get("payload") or {}) != sequence: + continue + if _history_event_type(event) == "ActivityStarted": + started = True + elif _history_event_type(event) == "ActivityRetryScheduled": + started = False + return PreparedLocalActivityCall(command, sequence, started, cleanup) + + def _prepared_group(commands: list[Any]) -> PreparedLocalActivityGroup | None: + if not 1 <= len(commands) <= 100: + raise LocalActivityExecutionAborted("atomic prepared groups require 1 to 100 authored members") + openings = { + "ActivityScheduled": "activity", "TimerScheduled": "timer", + "ChildWorkflowScheduled": "child workflow", "ChildRunStarted": "child workflow", + } + present: set[int] = set() + for offset, command in enumerate(commands): + sequence = current_call_sequence + offset + if not isinstance(command, ScheduleActivity | RecordLocalActivity | StartTimer | StartChildWorkflow): + raise LocalActivityExecutionAborted("prepared group contains an unsupported admission operation") + for event in events: + payload = event.get("payload") or {} + if _workflow_sequence(payload) != sequence: + continue + kind = _history_event_type(event) + if kind not in openings: + continue + # A complete path is admission authority. Legacy metadata-poor + # history cannot authorize fresh Source callbacks. + details = {**details_by_sequence.get(sequence, {}), **_recorded_step_details(payload)} + if payload.get("parallel_group_path") != command._parallel_group_path: + raise NonDeterministicReplayError( + sequence, "complete authored parallel group path", [kind], + detail="prepared group history changed or omitted its membership", + ) + _assert_step_matches(command, _RecordedStep(sequence, openings[kind], [kind], details)) + present.add(sequence) + if present and len(present) != len(commands): + raise NonDeterministicReplayError( + current_call_sequence, "all parallel members scheduled", [], + detail="prepared group history is missing a declared member", + ) + # Never reinterpret terminal/started history without its atomic opening. + group_sequences = set(range(current_call_sequence, current_call_sequence + len(commands))) + if not present and group_sequences.intersection(event_types_by_sequence): + raise NonDeterministicReplayError( + current_call_sequence, "complete atomic group opening", [], + detail="prepared group history lacks its original scheduling batch", + ) + calls = tuple( + _prepared_call(command, current_call_sequence + offset) + for offset, command in enumerate(commands) + if isinstance(command, RecordLocalActivity) and current_call_sequence + offset not in resolved_sequences + ) + if not calls: + return None + return PreparedLocalActivityGroup( + current_call_sequence, len(commands), () if present else tuple(commands), calls, bool(present), + ) + + def _assert_cancellation_call_matches(command: Any, boundary: CancellationDelivery) -> None: + if isinstance(command, list): + leaves, _ = _annotate_parallel_commands(command, boundary.sequence) + elif isinstance(command, SelectGroup): + leaves, _, _ = _annotate_selection(command, None) + elif isinstance(command, DurableOperationHandle): + _assert_authored_selection_handle(command) + return + else: + leaves = [command] + opening_shapes = { + "ActivityScheduled": "activity", "ActivityStarted": "activity", + "TimerScheduled": "timer", "ChildWorkflowScheduled": "child workflow", + "ChildRunStarted": "child workflow", "ConditionWaitOpened": "condition wait", + } + for offset, leaf in enumerate(leaves): + sequence = boundary.sequence + offset + for event in events: + event_type = _history_event_type(event) + payload = event.get("payload") or {} + if _workflow_sequence(payload) != sequence or event_type not in opening_shapes: + continue + if event_type == "TimerScheduled" and _internal_timeout_timer_kind( + payload, condition_wait_ids_by_sequence, + ) is not None: + continue + _assert_step_matches(leaf, _RecordedStep( + workflow_sequence=sequence, shape=opening_shapes[event_type], + event_types=[event_type], details={ + **details_by_sequence.get(sequence, {}), **_recorded_step_details(payload), + }, + )) + def _terminal_state(value: Any, *, include_pending: bool) -> _ReplayState: _apply_due_receivers() _consume_terminal_condition_reopens() @@ -5044,6 +5617,63 @@ def _terminal_state(value: Any, *, include_pending: bool) -> _ReplayState: _assert_no_unconsumed_history("complete workflow") return _state(commands + [CompleteWorkflow(result=value)]) + def _consume_cancellation(command: Any, boundary: CancellationDelivery) -> _ReplayState | None: + nonlocal cancellation_consumed, authored_sequence, wait_yield_count, advanced_cmd + _assert_cancellation_call_matches(command, boundary) + if boundary.call_kind == "condition": + for event in events: + payload = event.get("payload") or {} + if ( + _history_event_type(event) == "ConditionWaitOpened" + and _workflow_sequence(payload) == boundary.sequence + ): + _apply_condition_wait_receivers(payload.get("condition_wait_id")) + if ( + wait_yield_count < len(recorded_wait_steps) + and recorded_wait_steps[wait_yield_count].workflow_sequence == boundary.sequence + ): + wait_yield_count += 1 + cancellation_consumed = True + ctx._cancel_requested = True + ctx._cancellation_request_id = boundary.request_id + metadata = cancellation.request.context if cancellation.request is not None else None + if metadata is not None: + ctx._observe_cancellation_replay_time( + events[cancellation.delivery_index] if cancellation.delivery_index is not None else None, + ) + ctx._cancellation_context = ctx._bind_cancellation_context(metadata) + _apply_due_receivers() + if isinstance(command, DurableOperationHandle): + authored_sequence += 1 + try: + advanced_cmd = gen.throw(WorkflowCancelled( + "workflow cancellation was requested", request_id=boundary.request_id, + context=ctx._cancellation_context, + )) + return None + except StopIteration as stop: + return _terminal_state(stop.value, include_pending=True) + + def _advance_cancellation_clock(indexes: Iterable[int | None]) -> None: + if ctx._cancellation_context is None: + return + for index in sorted({index for index in indexes if index is not None}): + ctx._observe_cancellation_replay_time(events[index]) + + def _advance_selection_clock(base: int, size: int, failure: BaseException | None) -> None: + indexes = [ + selection_resolution_indexes[sequence] for sequence in range(base, base + size) + if sequence in selection_resolution_indexes + ] + if failure is not None: + failure_index = next( + selection_resolution_indexes[sequence] for sequence in range(base, base + size) + if selection_resolutions.get(sequence) is failure + ) + indexes = [index for index in indexes if index <= failure_index] + _advance_cancellation_clock(indexes) + + replay_token = _ACTIVE_WORKFLOW_REPLAY.set(ctx) try: while True: # Cursor-0 receivers are start-boundary events. Enter run() once @@ -5062,6 +5692,43 @@ def _terminal_state(value: Any, *, include_pending: bool) -> _ReplayState: return _terminal_state(stop.value, include_pending=True) first = False _apply_due_receivers() + _validate_local_activity_policies(cmd) + if prepare_local_activities and isinstance(cmd, list | SelectGroup) and _contains_local_activity(cmd) and ( + not prepare_local_activity_groups or isinstance(cmd, SelectGroup) + ): + raise LocalActivityExecutionAborted( + "prepared local groups require an implemented atomic group consumer", + ) + boundary = _cancellation_boundary(cmd) + if cancellation.delivery is not None and not cancellation_consumed: + cancellation_marker = cancellation.delivery + if cancellation_marker.sequence < current_call_sequence: + raise NonDeterministicReplayError( + cancellation_marker.sequence, "committed cancellation boundary", + ["CooperativeCancellationDelivered"], + detail="workflow passed the committed authored call", + ) + if cancellation_marker.sequence == current_call_sequence: + if boundary != cancellation_marker or ctx._cancellation_shield_depth > 0: + raise NonDeterministicReplayError( + cancellation_marker.sequence, "matching unshielded authored call", + ["CooperativeCancellationDelivered"], + detail="committed cancellation call kind or range changed", + ) + terminal = _consume_cancellation(cmd, boundary) + if terminal is not None: + return terminal + continue + elif boundary is not None and not cancellation_consumed and ctx._cancellation_shield_depth == 0: + eligible_sequence = boundary.operation_sequence or boundary.sequence + eligible_span = ( + boundary.operation_sequence_span + if boundary.operation_sequence is not None else boundary.sequence_span + ) + if cancellation.eligible(eligible_sequence, eligible_span): + _assert_cancellation_call_matches(cmd, boundary) + cancellation_intent = boundary + return _state(pending) if isinstance(cmd, SelectGroup): marker = ( selection_markers[selection_marker_cursor] @@ -5152,6 +5819,7 @@ def _terminal_state(value: Any, *, include_pending: bool) -> _ReplayState: ["SelectionResolved"], detail="the committed winner has no terminal member history", ) + _advance_cancellation_clock([selection_marker_indexes[selection_marker_cursor - 1]]) for sequence in range( winner.base_sequence, winner.base_sequence + winner.size, @@ -5176,14 +5844,40 @@ def _terminal_state(value: Any, *, include_pending: bool) -> _ReplayState: "ordinary parallel list groups do not support WaitCondition; " "use WorkflowContext.select() for durable condition selection" ) - cmd, result_shape = _annotate_parallel_commands(cmd) + prepared_group = prepare_local_activities and _contains_local_activity(cmd) + cmd, result_shape = _annotate_parallel_commands(cmd, current_call_sequence if prepared_group else None) needed = len(cmd) + if prepared_group: + group = _prepared_group(cmd) + if group is not None: + ctx.logger._set_replaying(False) + return _state(pending, prepared_local_activity_group=group) + if any(current_call_sequence + offset not in resolved_sequences for offset in range(needed)): + # Local results are durable. Remote members own their + # remaining work, with no repeat scheduling batch. + return _state(pending) + elif any(isinstance(command, RecordLocalActivity) for command in cmd) and ( + result_cursor + needed > len(resolved_results) + ): + raise LocalActivityExecutionAborted("local groups require negotiated atomic prepared admission") if result_cursor + needed <= len(resolved_results): for offset, child_command in enumerate(cmd): _assert_next_step_matches(child_command, offset) vals = resolved_results[result_cursor : result_cursor + needed] + consumed_steps = recorded_steps[result_cursor : result_cursor + needed] result_cursor += needed failed = _first_yield_failure(vals) + consumed_indexes = [step.history_index for step in consumed_steps] + if failed is not None: + failure_index = next( + step.history_index for step, value in zip(consumed_steps, vals, strict=True) + if value is failed + ) + consumed_indexes = [ + index for index in consumed_indexes + if index is not None and failure_index is not None and index <= failure_index + ] + _advance_cancellation_clock(consumed_indexes) if failed is not None: try: advanced_cmd = gen.throw(failed) @@ -5200,6 +5894,9 @@ def _terminal_state(value: Any, *, include_pending: bool) -> _ReplayState: if isinstance(cmd, DurableOperationHandle): _assert_authored_selection_handle(cmd) if _selection_cancellation_for_handle(cmd) is not None: + _advance_cancellation_clock([ + cancelled_selection_indexes[(cmd.selection_group_id, cmd.base_sequence)], + ]) try: advanced_cmd = gen.throw( DurableOperationCancelled( @@ -5230,6 +5927,7 @@ def _terminal_state(value: Any, *, include_pending: bool) -> _ReplayState: if not resolved: ctx.logger._set_replaying(False) return _state(pending) + _advance_selection_clock(cmd.base_sequence, cmd.size, failure) if failure is not None: try: advanced_cmd = gen.throw(failure) @@ -5374,6 +6072,34 @@ def _terminal_state(value: Any, *, include_pending: bool) -> _ReplayState: resolution: str | None = None opened: dict[str, Any] | None = None if wait_yield_count < len(wait_opened): + if cancellation.request is not None and wait_yield_count < len(recorded_wait_steps): + physical_sequence = recorded_wait_steps[wait_yield_count].workflow_sequence + if physical_sequence > current_call_sequence: + current_call_sequence = physical_sequence + authored_sequence = max(authored_sequence, physical_sequence + 1) + physical_boundary = CancellationDelivery( + cancellation.request.request_id, physical_sequence, "condition", + ) + if ( + cancellation.delivery is not None and not cancellation_consumed + and cancellation.delivery.sequence == physical_sequence + ): + if cancellation.delivery != physical_boundary or ctx._cancellation_shield_depth > 0: + raise NonDeterministicReplayError( + physical_sequence, "matching unshielded reopened condition", + ["CooperativeCancellationDelivered"], + ) + terminal = _consume_cancellation(cmd, physical_boundary) + if terminal is not None: + return terminal + break + if ( + cancellation.delivery is None and not cancellation_consumed + and ctx._cancellation_shield_depth == 0 and cancellation.eligible(physical_sequence) + ): + _assert_cancellation_call_matches(cmd, physical_boundary) + cancellation_intent = physical_boundary + return _state(pending) if wait_yield_count < len(recorded_wait_steps): _assert_step_matches(cmd, recorded_wait_steps[wait_yield_count]) opened = wait_opened[wait_yield_count] @@ -5384,6 +6110,8 @@ def _terminal_state(value: Any, *, include_pending: bool) -> _ReplayState: if isinstance(opened_id, str): resolution = wait_resolutions.get(opened_id) _apply_condition_wait_receivers(opened_id) + if resolution is not None: + _advance_cancellation_clock([condition_wait_terminal_indexes.get(opened_id)]) next_wait_index = wait_yield_count + 1 has_reopened_same_wait = ( opened is not None @@ -5445,6 +6173,8 @@ def _terminal_state(value: Any, *, include_pending: bool) -> _ReplayState: if result_cursor < len(resolved_results): _assert_next_step_matches(cmd) val = resolved_results[result_cursor] + if not isinstance(cmd, RecordLocalActivity) or prepare_local_activities: + _advance_cancellation_clock([recorded_steps[result_cursor].history_index]) result_cursor += 1 if isinstance(val, ActivityFailed | ChildWorkflowFailed): try: @@ -5456,6 +6186,8 @@ def _terminal_state(value: Any, *, include_pending: bool) -> _ReplayState: continue ctx.logger._set_replaying(False) _assert_pending_step_matches(cmd) + if isinstance(cmd, RecordLocalActivity) and prepare_local_activities: + return _state(pending, _prepared_call(cmd, current_call_sequence)) if isinstance(cmd, RecordLocalActivity) and local_activity_executor is not None: local_result = local_activity_executor(cmd) if cmd.outcome is None or cmd.arguments_envelope is None: @@ -5487,3 +6219,5 @@ def _terminal_state(value: Any, *, include_pending: bool) -> _ReplayState: return _state(pending + [_fail_workflow_from_exception(exc)]) except Exception as exc: return _state(pending + [_fail_workflow_from_exception(exc)]) + finally: + _ACTIVE_WORKFLOW_REPLAY.reset(replay_token) diff --git a/tests/fixtures/cooperative-cancellation-context.json b/tests/fixtures/cooperative-cancellation-context.json new file mode 100644 index 0000000..dbe7e89 --- /dev/null +++ b/tests/fixtures/cooperative-cancellation-context.json @@ -0,0 +1,17 @@ +{ + "schema": "durable-workflow.cancellation-context/v1", + "request_id": "request-1", + "root_request_id": "root-1", + "root_workflow_instance_id": "parent-instance", + "root_workflow_run_id": "parent-run", + "parent_request_id": "root-1", + "reason": "maintenance", + "requester": {"type": "operator", "id": "operator-1", "label": "Maintainer"}, + "source": "control_plane", + "requested_at": "2026-10-01T00:00:00.123456Z", + "cleanup_deadline_at": "2026-10-01T00:00:30.123456Z", + "lineage": [ + {"request_id": "root-1", "workflow_instance_id": "parent-instance", "workflow_run_id": "parent-run"}, + {"request_id": "request-1", "workflow_instance_id": "child-instance", "workflow_run_id": "run-1"} + ] +} diff --git a/tests/fixtures/replay_regressions/activity-cancellation-policy-changed.json b/tests/fixtures/replay_regressions/activity-cancellation-policy-changed.json new file mode 100644 index 0000000..4f781ac --- /dev/null +++ b/tests/fixtures/replay_regressions/activity-cancellation-policy-changed.json @@ -0,0 +1,13 @@ +{ + "$schema": "https://raw.githubusercontent.com/durable-workflow/.github/main/regression-corpus/evidence-schema.json", + "fixture_schema": "durable-workflow.replay-regression/v1", + "id": "python-activity-cancellation-policy-changed", + "protocol_version": "1.20", + "bindings": ["python"], + "workflow": {"type": "golden.single-activity", "input": ["Ada"], "payload_codec": "avro"}, + "history": [ + {"event_type": "ActivityScheduled", "payload": {"sequence": 1, "activity": {"type": "golden.greet", "cancellation_policy": "wait_cancellation_completed"}}} + ], + "expected": {"command_sequence": []}, + "expected_replay_error": {"type": "NonDeterministicReplayError", "message_contains": "activity_cancellation_policy_changed", "workflow_sequence": 1} +} diff --git a/tests/fixtures/replay_regressions/cancellation-scope-unqualified-worker-rejected.json b/tests/fixtures/replay_regressions/cancellation-scope-unqualified-worker-rejected.json new file mode 100644 index 0000000..dc917b2 --- /dev/null +++ b/tests/fixtures/replay_regressions/cancellation-scope-unqualified-worker-rejected.json @@ -0,0 +1,32 @@ +{ + "$schema": "https://raw.githubusercontent.com/durable-workflow/.github/main/regression-corpus/evidence-schema.json", + "fixture_schema": "durable-workflow.replay-regression/v1", + "id": "cancellation-scope-unqualified-worker-rejected", + "protocol_version": "1.20", + "bindings": ["python"], + "workflow": { + "type": "golden.single-activity", + "input": ["Ada"], + "payload_codec": "avro" + }, + "history": [ + { + "event_type": "CancellationScopeOpened", + "payload": { + "schema": "durable-workflow.cancellation-scope/v1", + "workflow_run_id": "scope-admission-run", + "scope_id": "scope-one", + "parent_scope_id": "root", + "sequence": 1, + "shield_parent": false + } + } + ], + "expected_replay_error": { + "type": "LocalActivityExecutionAborted", + "message_contains": "cancellation_scope_execution_not_supported.*Python" + }, + "expected": { + "command_sequence": [] + } +} diff --git a/tests/fixtures/replay_regressions/child-cancellation-policy-changed.json b/tests/fixtures/replay_regressions/child-cancellation-policy-changed.json new file mode 100644 index 0000000..9de5b32 --- /dev/null +++ b/tests/fixtures/replay_regressions/child-cancellation-policy-changed.json @@ -0,0 +1,14 @@ +{ + "$schema": "https://raw.githubusercontent.com/durable-workflow/.github/main/regression-corpus/evidence-schema.json", + "fixture_schema": "durable-workflow.replay-regression/v1", + "id": "python-child-cancellation-policy-changed", + "protocol_version": "1.20", + "bindings": ["python"], + "workflow": {"type": "tests.replay.child-policy-author", "input": [], "payload_codec": "avro"}, + "history": [ + {"event_type": "ChildWorkflowScheduled", "payload": {"sequence": 1, "child_workflow_type": "child", "parent_close_policy": "abandon", "cancellation_policy": "wait_cancellation_completed"}}, + {"event_type": "ChildRunCompleted", "payload": {"sequence": 1, "child_workflow_type": "child"}} + ], + "expected": {"command_sequence": []}, + "expected_replay_error": {"type": "NonDeterministicReplayError", "message_contains": "child_workflow_policy_changed", "workflow_sequence": 1} +} diff --git a/tests/fixtures/replay_regressions/cooperative-reopened-condition-cleanup.json b/tests/fixtures/replay_regressions/cooperative-reopened-condition-cleanup.json new file mode 100644 index 0000000..bb3ac12 --- /dev/null +++ b/tests/fixtures/replay_regressions/cooperative-reopened-condition-cleanup.json @@ -0,0 +1,33 @@ +{ + "$schema": "https://raw.githubusercontent.com/durable-workflow/.github/main/regression-corpus/evidence-schema.json", + "fixture_schema": "durable-workflow.replay-regression/v1", + "id": "cooperative-reopened-condition-cleanup", + "protocol_version": "1.20", + "bindings": ["python"], + "workflow": { + "type": "tests.replay.cooperative-reopened-condition-cleanup", + "input": [], + "payload_codec": "avro" + }, + "history": [ + {"event_type": "ConditionWaitOpened", "payload": {"sequence": 1, "condition_key": "forward-wait", "condition_wait_id": "wait-1"}}, + {"event_type": "ConditionWaitSatisfied", "payload": {"sequence": 1, "condition_wait_id": "wait-1"}}, + {"event_type": "ConditionWaitOpened", "payload": {"sequence": 2, "condition_key": "forward-wait", "condition_wait_id": "wait-2"}}, + { + "event_type": "CooperativeCancellationRequested", + "workflow_command_id": "request-1", + "recorded_at": "2026-09-30T12:00:00Z", + "payload": {"workflow_command_id": "request-1", "workflow_run_id": "run-1", "cleanup_deadline_at": "2026-09-30T12:10:00Z"} + }, + { + "event_type": "CooperativeCancellationDelivered", + "workflow_command_id": "request-1", + "payload": {"workflow_command_id": "request-1", "workflow_run_id": "run-1", "sequence": 2, "call_kind": "condition"} + }, + {"event_type": "TimerScheduled", "payload": {"sequence": 3, "timer_kind": "durable_timer", "delay_seconds": 1}}, + {"event_type": "TimerFired", "payload": {"sequence": 3, "timer_kind": "durable_timer"}} + ], + "expected": { + "command_sequence": [{"type": "complete_workflow", "command_type": "CompleteWorkflow", "result": "request-1"}] + } +} diff --git a/tests/fixtures/replay_regressions/prepared-local-cold-results.json b/tests/fixtures/replay_regressions/prepared-local-cold-results.json new file mode 100644 index 0000000..3ad3b83 --- /dev/null +++ b/tests/fixtures/replay_regressions/prepared-local-cold-results.json @@ -0,0 +1,23 @@ +{ + "$schema": "https://raw.githubusercontent.com/durable-workflow/.github/main/regression-corpus/evidence-schema.json", + "fixture_schema": "durable-workflow.replay-regression/v1", + "id": "prepared-local-cold-results", + "protocol_version": "1.20", + "bindings": ["python"], + "workflow": { + "type": "tests.replay.prepared-local-cold-results", + "input": [], + "payload_codec": "avro" + }, + "history": [ + {"event_type": "ActivityScheduled", "payload": {"sequence": 1, "activity_type": "prepared.first", "execution_mode": "local"}}, + {"event_type": "ActivityStarted", "payload": {"sequence": 1, "activity_type": "prepared.first", "execution_mode": "local", "activity_execution_id": "first", "activity_attempt_id": "first-attempt"}}, + {"event_type": "ActivityCompleted", "payload": {"sequence": 1, "activity_type": "prepared.first", "execution_mode": "local", "result": {"codec": "avro", "blob": "wwHioz3/VYAiNwoUZmFzdC12YWx1ZQ=="}}}, + {"event_type": "ActivityScheduled", "payload": {"sequence": 2, "activity_type": "prepared.second", "execution_mode": "local"}}, + {"event_type": "ActivityStarted", "payload": {"sequence": 2, "activity_type": "prepared.second", "execution_mode": "local", "activity_execution_id": "second", "activity_attempt_id": "second-attempt"}}, + {"event_type": "ActivityCompleted", "payload": {"sequence": 2, "activity_type": "prepared.second", "execution_mode": "local", "result": {"codec": "avro", "blob": "wwHioz3/VYAiNwoUZmFzdC12YWx1ZQ=="}}} + ], + "expected": { + "command_sequence": [{"type": "complete_workflow", "command_type": "CompleteWorkflow", "result": ["fast-value", "fast-value"]}] + } +} diff --git a/tests/fixtures/replay_regressions/prepared-local-group-cold-results.json b/tests/fixtures/replay_regressions/prepared-local-group-cold-results.json new file mode 100644 index 0000000..f160ba6 --- /dev/null +++ b/tests/fixtures/replay_regressions/prepared-local-group-cold-results.json @@ -0,0 +1,23 @@ +{ + "$schema": "https://raw.githubusercontent.com/durable-workflow/.github/main/regression-corpus/evidence-schema.json", + "fixture_schema": "durable-workflow.replay-regression/v1", + "id": "prepared-local-group-cold-results", + "protocol_version": "1.20", + "bindings": ["python"], + "workflow": { + "type": "tests.replay.prepared-local-group-cold-results", + "input": [], + "payload_codec": "avro" + }, + "history": [ + {"event_type": "ActivityScheduled", "payload": {"sequence": 1, "activity_type": "prepared.first", "execution_mode": "local", "parallel_group_path": [{"parallel_group_id": "parallel-activities:1:2", "parallel_group_kind": "activity", "parallel_group_base_sequence": 1, "parallel_group_size": 2, "parallel_group_index": 0}]}}, + {"event_type": "ActivityScheduled", "payload": {"sequence": 2, "activity_type": "prepared.second", "execution_mode": "local", "parallel_group_path": [{"parallel_group_id": "parallel-activities:1:2", "parallel_group_kind": "activity", "parallel_group_base_sequence": 1, "parallel_group_size": 2, "parallel_group_index": 1}]}}, + {"event_type": "ActivityStarted", "payload": {"sequence": 1, "activity_type": "prepared.first", "execution_mode": "local", "activity_execution_id": "first", "activity_attempt_id": "first-original-attempt", "parallel_group_path": [{"parallel_group_id": "parallel-activities:1:2", "parallel_group_kind": "activity", "parallel_group_base_sequence": 1, "parallel_group_size": 2, "parallel_group_index": 0}]}}, + {"event_type": "ActivityStarted", "payload": {"sequence": 2, "activity_type": "prepared.second", "execution_mode": "local", "activity_execution_id": "second", "activity_attempt_id": "second-original-attempt", "parallel_group_path": [{"parallel_group_id": "parallel-activities:1:2", "parallel_group_kind": "activity", "parallel_group_base_sequence": 1, "parallel_group_size": 2, "parallel_group_index": 1}]}}, + {"event_type": "ActivityCompleted", "payload": {"sequence": 1, "activity_type": "prepared.first", "execution_mode": "local", "result": {"codec": "avro", "blob": "wwHioz3/VYAiNwoUZmFzdC12YWx1ZQ=="}, "parallel_group_path": [{"parallel_group_id": "parallel-activities:1:2", "parallel_group_kind": "activity", "parallel_group_base_sequence": 1, "parallel_group_size": 2, "parallel_group_index": 0}]}}, + {"event_type": "ActivityCompleted", "payload": {"sequence": 2, "activity_type": "prepared.second", "execution_mode": "local", "result": {"codec": "avro", "blob": "wwHioz3/VYAiNwoUZmFzdC12YWx1ZQ=="}, "parallel_group_path": [{"parallel_group_id": "parallel-activities:1:2", "parallel_group_kind": "activity", "parallel_group_base_sequence": 1, "parallel_group_size": 2, "parallel_group_index": 1}]}} + ], + "expected": { + "command_sequence": [{"type": "complete_workflow", "command_type": "CompleteWorkflow", "result": ["fast-value", "fast-value"]}] + } +} diff --git a/tests/fixtures/scoped-run-cancellation-context.json b/tests/fixtures/scoped-run-cancellation-context.json new file mode 100644 index 0000000..7d9d0a4 --- /dev/null +++ b/tests/fixtures/scoped-run-cancellation-context.json @@ -0,0 +1,164 @@ +{ + "native_commit": "a51395da59ef70945f91f00581ecb6bad8765bdd", + "child": { + "schema": "durable-workflow.cancellation-context/v2", + "request_id": "child-request", + "root_request_id": "root-request", + "root_workflow_instance_id": "root-instance", + "root_workflow_run_id": "root-run", + "parent_request_id": "inner-request", + "reason": "maintenance", + "requester": { + "type": "operator", + "id": "operator-1" + }, + "source": "api", + "requested_at": "2026-10-04T00:00:00.123456Z", + "cleanup_deadline_at": "2026-10-04T00:00:15.123456Z", + "lineage": [ + { + "request_id": "root-request", + "workflow_instance_id": "root-instance", + "workflow_run_id": "root-run" + }, + { + "request_id": "child-request", + "workflow_instance_id": "child-instance", + "workflow_run_id": "child-run" + } + ], + "scope_origin": { + "schema": "durable-workflow.scoped-cancellation-context/v1", + "root_context": { + "schema": "durable-workflow.cancellation-context/v1", + "request_id": "root-request", + "root_request_id": "root-request", + "root_workflow_instance_id": "root-instance", + "root_workflow_run_id": "root-run", + "parent_request_id": null, + "reason": "maintenance", + "requester": { + "type": "operator", + "id": "operator-1" + }, + "source": "api", + "requested_at": "2026-10-04T00:00:00.123456Z", + "cleanup_deadline_at": "2026-10-04T00:00:30.123456Z", + "lineage": [ + { + "request_id": "root-request", + "workflow_instance_id": "root-instance", + "workflow_run_id": "root-run" + } + ] + }, + "lineage": [ + { + "request_id": "root-request", + "workflow_instance_id": "root-instance", + "workflow_run_id": "root-run", + "scope_id": "outer", + "cleanup_deadline_at": "2026-10-04T00:00:30.123456Z" + }, + { + "request_id": "inner-request", + "workflow_instance_id": "root-instance", + "workflow_run_id": "root-run", + "scope_id": "inner", + "cleanup_deadline_at": "2026-10-04T00:00:20.123456Z" + } + ] + }, + "scope_authority_deadline_at": "2026-10-04T00:00:15.123456Z" + }, + "grandchild": { + "schema": "durable-workflow.cancellation-context/v2", + "request_id": "grandchild-request", + "root_request_id": "root-request", + "root_workflow_instance_id": "root-instance", + "root_workflow_run_id": "root-run", + "parent_request_id": "child-scope-request", + "reason": "maintenance", + "requester": { + "type": "operator", + "id": "operator-1" + }, + "source": "api", + "requested_at": "2026-10-04T00:00:00.123456Z", + "cleanup_deadline_at": "2026-10-04T00:00:12.123456Z", + "lineage": [ + { + "request_id": "root-request", + "workflow_instance_id": "root-instance", + "workflow_run_id": "root-run" + }, + { + "request_id": "child-scope-request", + "workflow_instance_id": "child-instance", + "workflow_run_id": "child-run" + }, + { + "request_id": "grandchild-request", + "workflow_instance_id": "grandchild-instance", + "workflow_run_id": "grandchild-run" + } + ], + "scope_origin": { + "schema": "durable-workflow.scoped-cancellation-context/v1", + "root_context": { + "schema": "durable-workflow.cancellation-context/v1", + "request_id": "root-request", + "root_request_id": "root-request", + "root_workflow_instance_id": "root-instance", + "root_workflow_run_id": "root-run", + "parent_request_id": null, + "reason": "maintenance", + "requester": { + "type": "operator", + "id": "operator-1" + }, + "source": "api", + "requested_at": "2026-10-04T00:00:00.123456Z", + "cleanup_deadline_at": "2026-10-04T00:00:30.123456Z", + "lineage": [ + { + "request_id": "root-request", + "workflow_instance_id": "root-instance", + "workflow_run_id": "root-run" + } + ] + }, + "lineage": [ + { + "request_id": "root-request", + "workflow_instance_id": "root-instance", + "workflow_run_id": "root-run", + "scope_id": "outer", + "cleanup_deadline_at": "2026-10-04T00:00:30.123456Z" + }, + { + "request_id": "inner-request", + "workflow_instance_id": "root-instance", + "workflow_run_id": "root-run", + "scope_id": "inner", + "cleanup_deadline_at": "2026-10-04T00:00:20.123456Z" + }, + { + "request_id": "child-request", + "workflow_instance_id": "child-instance", + "workflow_run_id": "child-run", + "scope_id": "root", + "cleanup_deadline_at": "2026-10-04T00:00:15.123456Z" + }, + { + "request_id": "child-scope-request", + "workflow_instance_id": "child-instance", + "workflow_run_id": "child-run", + "scope_id": "child-scope", + "cleanup_deadline_at": "2026-10-04T00:00:12.123456Z" + } + ] + }, + "scope_authority_deadline_at": "2026-10-04T00:00:12.123456Z" + } +} diff --git a/tests/integration/COOPERATIVE.md b/tests/integration/COOPERATIVE.md new file mode 100644 index 0000000..17e71d2 --- /dev/null +++ b/tests/integration/COOPERATIVE.md @@ -0,0 +1,54 @@ +# Cooperative cancellation source qualification + +The published SDK defaults to worker protocol 1.19. These cases explicitly select +the candidate 1.20 protocol and require compatible Server capability discovery. + +Run the `CI` workflow on the candidate branch with `cooperative_qualification=true` +and `server_commit` set to the exact 40-character public Server candidate SHA. +The workflow checks out that source, builds an isolated MySQL/Redis Server stack, +runs the integration suite, retains JUnit and removes its containers, images, +network and payload volume. Against an already isolated candidate stack: + +```sh +export DURABLE_WORKFLOW_SERVER_URL=http://127.0.0.1:8080 +export DURABLE_WORKFLOW_AUTH_TOKEN=test-token +export DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION=1.20 +export DURABLE_WORKFLOW_COOPERATIVE_QUALIFICATION=1 +pytest tests/integration/test_cooperative_cancellation.py -v --junitxml=cooperative-results.xml +``` + +The suite covers canonical delivery and cleanup, lost replies, waiting runs, +local callbacks, cleanup shutdown/replay and deadlines. Its remote cases use the +actual `Worker.run()` pollers and callback dispatcher, rather than only manually +claiming an attempt and sending a heartbeat. Async and synchronous callbacks are +tested with and without user heartbeats, through an actual accepted worker +registration heartbeat. Shutdown expiry fences callbacks before replacement +cleanup. Another case kills the real remote owner process, delivers the request +in a new process and rejects the dead owner's late completion/failure. +A separate case kills an active owner, waits for the real five-minute activity +lease and repair pass, and requires attempt 2 under a distinct owner before +requesting cancellation. It checks the killed attempt cannot publish a result, +failure or heartbeat, renew its lease, or change canonical history, then verifies +one cleanup with the original request identity. + +Readonly status calls never record user progress or renew the activity lease. +They fail closed on refused/invalid ownership, elapsed execution/session bounds +or failed observation. A pending workflow request becomes an activity stop only +when the workflow worker durably delivers cancellation. Workflow capacity must +remain available while an activity is blocked. Here a Python worker provides it +through separate workflow execution and activity thread capacities. + +Synchronous remote handlers use a lazy pool bounded by configured activity +concurrency. A running thread retains its execution slot until it actually +finishes, even after its durable attempt is abandoned. Worker availability and +poll admission reflect that occupied capacity. Its late heartbeats and result +publication are fenced. Python cannot forcibly stop a running thread or a +callable that suppresses cancellation. Such code may continue external effects +until it returns or its process supervisor stops the process. Activities still +need idempotency and reconciliation. + +Only activity-authored heartbeats extend the existing five-minute activity +lease and any user heartbeat deadline. A positive readonly observation does not +reserve ownership for a later completion. The Server independently validates +completion and failure fences. Exact published Server/SDK qualification remains +a separate release gate from these source scenarios. diff --git a/tests/integration/cooperative_worker.py b/tests/integration/cooperative_worker.py new file mode 100644 index 0000000..63f84ab --- /dev/null +++ b/tests/integration/cooperative_worker.py @@ -0,0 +1,63 @@ +"""A real child worker for cooperative process-loss qualification.""" + +from __future__ import annotations + +import asyncio +import json +import os +import sys + +from durable_workflow import Client, activity +from tests.integration.test_cooperative_cancellation import candidate_worker, cooperative_cleanup, poll_claim + + +async def cleanup(request_id: str) -> str: + info = activity.context().info + print(json.dumps({"phase": "cleanup", "request_id": request_id, "worker_id": info.worker_id}), flush=True) + if os.environ.get("DW_COOPERATIVE_FIXTURE_MODE") == "hold": + await asyncio.Event().wait() + return await cooperative_cleanup(request_id) + + +async def blocked_remote() -> None: + info = activity.context().info + print(json.dumps({"phase": "remote-entered", "task_id": info.task_id, + "activity_attempt_id": info.activity_attempt_id, + "lease_owner": info.worker_id, "attempt_number": info.attempt_number, + "callback_pid": os.getpid()}), flush=True) + await asyncio.Event().wait() + + +async def main() -> None: + queue, worker_id, mode = sys.argv[1:] + if mode not in {"hold", "finish", "remote"}: + raise ValueError("unknown qualification mode") + os.environ["DW_COOPERATIVE_FIXTURE_MODE"] = mode + async with Client( + os.environ["DURABLE_WORKFLOW_SERVER_URL"], + token=os.environ.get("DURABLE_WORKFLOW_AUTH_TOKEN", "test-token"), namespace="default", + ) as client: + worker = candidate_worker(client, queue, worker_id=worker_id) + + worker.activities["tests.python-cooperative-cleanup"] = cleanup + if mode == "remote": + worker.activities["tests.python-cooperative-work"] = blocked_remote + await worker.run() + return + await worker._register() + try: + task = await poll_claim(client, worker) + observation = task["cancellation_request"] + print(json.dumps({ + "phase": "claim", "attempt": task["workflow_task_attempt"], "worker_id": worker_id, + "task_id": task["task_id"], "request_id": observation["request_id"], + "cleanup_deadline_at": observation["cleanup_deadline_at"], + }), flush=True) + commands = await worker._run_workflow_task(task) + print(json.dumps({"phase": "finished", "committed": commands is not None}), flush=True) + finally: + await worker.stop() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/tests/integration/prepared_worker.py b/tests/integration/prepared_worker.py new file mode 100644 index 0000000..fb8cbdd --- /dev/null +++ b/tests/integration/prepared_worker.py @@ -0,0 +1,26 @@ +"""A killable owner for prepared callback and canonical cleanup qualification.""" + +from __future__ import annotations + +import asyncio +import os +import sys +from pathlib import Path + +from tests.integration.test_prepared_local_activity import TrackingPreparedClient, prepared_worker + + +async def main() -> None: + queue, owner, mode, trace = sys.argv[1:] + if mode not in {"hold", "finish"}: + raise ValueError("unknown prepared qualification mode") + os.environ["DW_PREPARED_FIXTURE_MODE"] = mode + async with TrackingPreparedClient( + os.environ["DURABLE_WORKFLOW_SERVER_URL"], + token=os.environ.get("DURABLE_WORKFLOW_AUTH_TOKEN", "test-token"), namespace="default", trace_path=Path(trace), + ) as client: + await prepared_worker(client, queue, worker_id=owner).run() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/tests/integration/test_cooperative_cancellation.py b/tests/integration/test_cooperative_cancellation.py new file mode 100644 index 0000000..3bf81f3 --- /dev/null +++ b/tests/integration/test_cooperative_cancellation.py @@ -0,0 +1,864 @@ +"""Connected cooperative qualification, opt-in until the service tuple is published.""" + +from __future__ import annotations + +import asyncio +import json +import os +import sys +import threading +import time +import uuid +from dataclasses import dataclass +from datetime import datetime +from pathlib import Path +from typing import Any + +import pytest + +from durable_workflow import CancellationPolicy, Client, Worker, activity, serializer, workflow +from durable_workflow.client import WorkflowHandle +from durable_workflow.errors import ServerError, WorkflowCancelled +from durable_workflow.worker import _poll_capacity_delay +from durable_workflow.workflow import LocalActivityExecutionAborted + +pytestmark = pytest.mark.usefixtures("cooperative_runtime") + + +@pytest.fixture +async def cooperative_runtime(server_url: str, server_token: str, monkeypatch: pytest.MonkeyPatch) -> None: + if os.environ.get("DURABLE_WORKFLOW_COOPERATIVE_QUALIFICATION") != "1": + pytest.skip("candidate cooperative Server qualification is opt-in") + async with Client(server_url, token=server_token, namespace="default") as client: + info = await client.get_cluster_info() + assert info["worker_protocol"]["server_capabilities"]["cooperative_cancellation"] is True + assert info["worker_protocol"]["version"] == "1.20" + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + + +@workflow.defn(name="tests.python-cooperative-cleanup") +class CooperativeCleanupWorkflow: + def run(self, ctx: Any, kind: str, cancellation_policy: str | None = None) -> Any: + try: + if kind == "local": + yield ctx.local_activity("tests.python-cooperative-work", []) + elif kind == "remote": + yield ctx.schedule_activity( + "tests.python-cooperative-work", [], cancellation_policy=cancellation_policy, + schedule_to_close_timeout=60 if cancellation_policy is not None else None, + ) + else: + yield ctx.start_timer(300) + except WorkflowCancelled as error: + if kind == "clock_timer": + assert error.context is ctx.cancellation_context and error.context is not None + print(json.dumps({"phase": "remaining-delivery", "remaining": error.context.remaining(), + "context": error.context.to_dict()}), flush=True) + with ctx.cancellation_shield(): + yield ctx.local_activity("tests.python-cooperative-cleanup", [error.request_id]) + if kind == "clock_timer": + assert error.context is not None + print(json.dumps({"phase": "remaining-cleanup", "remaining": error.context.remaining(), + "context": error.context.to_dict()}), flush=True) + return error.request_id + return "not cancelled" + + +@activity.defn(name="tests.python-cooperative-work") +async def cooperative_work() -> str: + return "ordinary result" + + +@activity.defn(name="tests.python-cooperative-cleanup") +async def cooperative_cleanup(request_id: str) -> str: + await activity.context().heartbeat({"request_id": request_id}) + return request_id + + +def candidate_worker(client: Client, queue: str, **kwargs: Any) -> Worker: + return Worker( + client, task_queue=queue, worker_id=kwargs.pop("worker_id", f"{queue}-worker"), + workflows=[CooperativeCleanupWorkflow], activities=[cooperative_work, cooperative_cleanup], + capabilities=["cooperative_cancellation"], **kwargs, + ) + + +async def poll_claim(client: Client, worker: Worker) -> dict[str, Any]: + async def poll() -> dict[str, Any]: + while True: + try: + task = await client.poll_workflow_task( + worker_id=worker.worker_id, task_queue=worker.task_queue, timeout=worker._poll_timeout, + ) + except ServerError as error: + delay = _poll_capacity_delay(error, "workflow_task", worker.task_queue) + if delay is None: + raise + print(json.dumps({"phase": "poll-deferral", "delay_seconds": delay}), flush=True) + await asyncio.sleep(delay) + continue + if task is not None: + return task + await asyncio.sleep(0.1) + return await asyncio.wait_for(poll(), timeout=30) + + +async def events(handle: WorkflowHandle) -> list[dict[str, Any]]: + history = await handle.get_history(page_size=1000) + assert history.get("next_page_token") is None + return history.get("events", history.get("history_events", [])) + + +async def assert_cancelled_cleanup(handle: WorkflowHandle, request_id: str) -> list[dict[str, Any]]: + with pytest.raises(WorkflowCancelled): + await handle.result(timeout=10) + history = await events(handle) + kinds = [event["event_type"] for event in history] + cancellation_kinds = ("CooperativeCancellationRequested", "CooperativeCancellationDelivered", "WorkflowCancelled") + for kind in (*cancellation_kinds, "ActivityCompleted"): + assert kinds.count(kind) == 1 + for kind in ("WorkflowCompleted", "WorkflowFailed", "ActivityFailed", "ActivityTimedOut"): + assert kind not in kinds + for event in history: + if event["event_type"] in cancellation_kinds: + assert event["payload"]["workflow_command_id"] == request_id + return history + + +class LostDeliveryAcknowledgmentClient(Client): + delivery_calls = 0 + + async def deliver_workflow_cancellation(self, **kwargs: Any) -> dict[str, Any]: + result = await super().deliver_workflow_cancellation(**kwargs) + self.delivery_calls += 1 + if self.delivery_calls == 1: + raise ServerError(503, {"reason": "qualification_lost_delivery_acknowledgment"}) + return result + + +class ObservedOwnerClient(Client): + def __init__(self, *args: Any, **kwargs: Any) -> None: + super().__init__(*args, **kwargs) + self.owner_heartbeat = asyncio.Event() + + async def heartbeat_worker(self, **kwargs: Any) -> Any: + reply = await super().heartbeat_worker(**kwargs) + if isinstance(reply, dict) and reply.get("acknowledged") is True: + self.owner_heartbeat.set() + return reply + + +@dataclass(frozen=True) +class AsyncRemoteQualification: + marker: str + user_heartbeat: bool = False + duration_seconds: int | None = None + + async def __call__(self) -> str | None: + if self.user_heartbeat: + await activity.context().heartbeat({"qualification": "remote-in-flight"}) + write_remote_marker(self.marker) + if self.duration_seconds is not None: + await asyncio.sleep(self.duration_seconds) + return "independent-completion" + await asyncio.Event().wait() + return None + + +@dataclass(frozen=True) +class SyncRemoteQualification: + marker: str + user_heartbeat: bool = False + + def __call__(self) -> None: + if self.user_heartbeat: + asyncio.run(activity.context().heartbeat({"qualification": "remote-in-flight"})) + write_remote_marker(self.marker) + while True: + time.sleep(1) + + +def write_remote_marker(marker: str) -> None: + info = activity.context().info + pending = Path(marker + ".writing") + pending.write_text(json.dumps({ + "task_id": info.task_id, "activity_attempt_id": info.activity_attempt_id, + "lease_owner": info.worker_id, "callback_pid": os.getpid(), + })) + pending.replace(marker) + + +async def remote_marker(marker: Path) -> dict[str, Any]: + async def ready() -> dict[str, Any]: + while not marker.exists(): + await asyncio.sleep(0.05) + return json.loads(marker.read_text()) + return await asyncio.wait_for(ready(), timeout=20) + + +async def callback_gone(pid: int) -> None: + async def gone() -> None: + while True: + try: + os.kill(pid, 0) + except ProcessLookupError: + return + await asyncio.sleep(0.05) + await asyncio.wait_for(gone(), timeout=10) + + +async def observed_stop_receipt(client: Client, fence: dict[str, str]) -> dict[str, Any]: + async def receipt() -> dict[str, Any]: + while True: + reply = await client.activity_task_status(**fence) + proof = reply.get("cancellation_acknowledgement") + if isinstance(proof, dict) and proof.get("callback_state") == "stopped": + assert reply["heartbeat_recorded"] is False + assert reply["can_continue"] is False + return proof + await asyncio.sleep(0.1) + return await asyncio.wait_for(receipt(), timeout=15) + + +@pytest.mark.parametrize("handler_kind", ["async", "sync"]) +@pytest.mark.parametrize("user_heartbeat,policy", [ + (False, None), (True, None), (False, CancellationPolicy.TRY_CANCEL), + (False, CancellationPolicy.WAIT_CANCELLATION_COMPLETED), +]) +async def test_actual_remote_worker_stops_callbacks_and_reports_original_cancellation( + server_url: str, server_token: str, handler_kind: str, user_heartbeat: bool, policy: CancellationPolicy | None, + tmp_path: Path, +) -> None: + queue = f"py-cooperative-owner-{uuid.uuid4().hex[:8]}" + marker = tmp_path / "remote" + async with ObservedOwnerClient(server_url, token=server_token, namespace="default") as client: + worker = candidate_worker(client, queue, max_concurrent_activity_tasks=1) + worker.activities["tests.python-cooperative-work"] = ( + AsyncRemoteQualification(str(marker), user_heartbeat) if handler_kind == "async" + else SyncRemoteQualification(str(marker), user_heartbeat) + ) + running = asyncio.create_task(worker.run()) + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, + input=["remote", policy.value if policy is not None else None], + ) + entered = await remote_marker(marker) + await asyncio.wait_for(client.owner_heartbeat.wait(), timeout=15) + fence = {key: entered[key] for key in ("task_id", "activity_attempt_id", "lease_owner")} + accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) + original = accepted["cancellation_request"] + await assert_cancelled_cleanup(handle, original["request_id"]) + await callback_gone(entered["callback_pid"]) + if os.environ.get("DURABLE_WORKFLOW_NATIVE_SOURCE_QUALIFICATION") == "1": + proof = await observed_stop_receipt(client, fence) + assert proof["request_id"] == original["request_id"] + assert proof["root_request_id"] == original["request_id"] + assert proof["cleanup_deadline_at"] == original["cleanup_deadline_at"] + assert proof["received_after_deadline"] is False + duplicate = await client.acknowledge_activity_cancellation(**fence, request_id=original["request_id"]) + assert duplicate["duplicate"] is True + assert duplicate["history_event_id"] == proof["history_event_id"] + print(f"Remote callback stop receipt: {json.dumps([entered, original, proof, duplicate])}") + history = await events(handle) + if os.environ.get("DURABLE_WORKFLOW_NATIVE_SOURCE_QUALIFICATION") == "1": + receipt = [event for event in history if event["event_type"] == "ActivityCancellationAcknowledged"] + assert len(receipt) == 1 + # Public history exposes event sequence/payload, not row IDs. + # The readonly status and duplicate transport prove receipt ID. + payload = receipt[0]["payload"] + assert payload["activity_attempt_id"] == fence["activity_attempt_id"] + assert payload["lease_owner"] == fence["lease_owner"] + assert payload["evidence_source"] == "activity_worker" + for key in ("request_id", "root_request_id", "cleanup_deadline_at", "acknowledged_at"): + assert payload[key] == proof[key] + delivery = [event for event in history if event["event_type"] == "CooperativeCancellationDelivered"][0] + assert delivery["payload"]["call_kind"] == "activity" + if policy is not None: + scheduled = [event for event in history if event["event_type"] == "ActivityScheduled"][0] + assert scheduled["payload"]["activity"]["cancellation_policy"] == policy.value + if policy == CancellationPolicy.WAIT_CANCELLATION_COMPLETED: + kinds = [event["event_type"] for event in history] + assert kinds.index("ActivityCancellationAcknowledged") < kinds.index("CooperativeCancellationDelivered") + assert len([event for event in history if event["event_type"] == "ActivityCancelled"]) == 1 + progress = [event for event in history if event["event_type"] == "ActivityHeartbeatRecorded" + and event["payload"].get("activity_type") == "tests.python-cooperative-work"] + assert bool(progress) is user_heartbeat + with pytest.raises(ServerError) as completion: + await client.complete_activity_task(**fence, result="late") + assert completion.value.status == 409 + with pytest.raises(ServerError) as failure: + await client.fail_activity_task(**fence, message="late", failure_type="LateQualification") + assert failure.value.status == 409 + assert await events(handle) == history + print(f"Remote stopped cancellation history: {json.dumps(history)}") + finally: + await worker.stop() + await asyncio.wait_for(running, timeout=10) + + +async def test_bounded_remote_abandon_completes_after_parent_cancellation( + server_url: str, server_token: str, tmp_path: Path, +) -> None: + queue = f"py-cooperative-abandon-{uuid.uuid4().hex[:8]}" + marker = tmp_path / "remote" + async with ObservedOwnerClient(server_url, token=server_token, namespace="default") as client: + worker = candidate_worker(client, queue, max_concurrent_activity_tasks=1) + worker.activities["tests.python-cooperative-work"] = AsyncRemoteQualification(str(marker), duration_seconds=25) + running = asyncio.create_task(worker.run()) + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, + input=["remote", CancellationPolicy.ABANDON.value], + ) + entered = await remote_marker(marker) + await asyncio.wait_for(client.owner_heartbeat.wait(), timeout=15) + fence = {key: entered[key] for key in ("task_id", "activity_attempt_id", "lease_owner")} + accepted = await handle.request_cancellation(cleanup_timeout_seconds=30) + duplicate = await handle.request_cancellation(cleanup_timeout_seconds=300) + assert duplicate["duplicate"] is True + assert duplicate["cancellation_request"] == accepted["cancellation_request"] + history = await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + os.kill(entered["callback_pid"], 0) + kinds = [event["event_type"] for event in history] + assert "ActivityCancelled" not in kinds + assert "ActivityCancellationAcknowledged" not in kinds + scheduled = [event for event in history if event["event_type"] == "ActivityScheduled"][0] + assert scheduled["payload"]["activity"]["cancellation_policy"] == "abandon" + total_deadline = scheduled["payload"]["activity"]["schedule_to_close_deadline_at"] + assert isinstance(total_deadline, str) + status = await client.activity_task_status(**fence) + assert status["can_continue"] is True + until = time.monotonic() + 35 + while True: + history = await events(handle) + completed = [event for event in history if event["event_type"] == "ActivityCompleted" + and event["payload"].get("activity_type") == "tests.python-cooperative-work"] + if completed or time.monotonic() >= until: + break + await asyncio.sleep(0.1) + assert len(completed) == 1 + assert serializer.decode_envelope(completed[0]["payload"]["result"]) == "independent-completion" + assert completed[0]["payload"]["activity"]["schedule_to_close_deadline_at"] == total_deadline + assert completed[0]["payload"]["activity_attempt_id"] == fence["activity_attempt_id"] + kinds = [event["event_type"] for event in history] + assert kinds.count("WorkflowCancelled") == 1 + assert "WorkflowCompleted" not in kinds + assert "ActivityCancellationAcknowledged" not in kinds + with pytest.raises(ServerError) as completion: + await client.complete_activity_task(**fence, result="stale") + assert completion.value.status == 409 + with pytest.raises(ServerError) as failure: + await client.fail_activity_task(**fence, message="stale", failure_type="LateQualification") + assert failure.value.status == 409 + assert await events(handle) == history + print(f"Bounded remote Abandon: {json.dumps([entered, accepted, completed])}") + finally: + await worker.stop() + await asyncio.wait_for(running, timeout=10) + + +@pytest.mark.parametrize("handler_kind", ["async", "sync"]) +async def test_actual_remote_shutdown_joins_callback_before_replacement_cleanup( + server_url: str, server_token: str, handler_kind: str, tmp_path: Path, +) -> None: + queue = f"py-cooperative-owner-stop-{uuid.uuid4().hex[:8]}" + marker = tmp_path / "remote" + async with Client(server_url, token=server_token, namespace="default") as client: + worker = candidate_worker(client, queue, shutdown_timeout=0.1) + worker.activities["tests.python-cooperative-work"] = ( + AsyncRemoteQualification(str(marker)) if handler_kind == "async" else SyncRemoteQualification(str(marker)) + ) + running = asyncio.create_task(worker.run()) + replacement = candidate_worker(client, queue, worker_id=f"{queue}-replacement") + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["remote"], + ) + entered = await remote_marker(marker) + before = await events(handle) + await asyncio.wait_for(worker.stop(), timeout=10) + await asyncio.wait_for(running, timeout=10) + await callback_gone(entered["callback_pid"]) + assert await events(handle) == before + accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) + await replacement._register() + assert await replacement._run_workflow_task(await poll_claim(client, replacement)) is not None + await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + print(f"Stopped remote callback before replacement: {json.dumps(entered)}") + finally: + await replacement.stop() + await worker.stop() + await asyncio.wait_for(running, timeout=10) + + +async def test_waiting_timer_is_cancelled_by_canonical_delivery( + server_url: str, server_token: str, +) -> None: + queue = f"py-cooperative-waiting-{uuid.uuid4().hex[:8]}" + async with Client(server_url, token=server_token, namespace="default") as client: + worker = candidate_worker(client, queue) + await worker._register() + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["timer"], + ) + task = await poll_claim(client, worker) + assert [command["type"] for command in await worker._run_workflow_task(task) or []] == ["start_timer"] + before = await events(handle) + assert [event["event_type"] for event in before].count("TimerScheduled") == 1 + accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) + resumed = await poll_claim(client, worker) + assert await worker._run_workflow_task(resumed) is not None + history = await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + kinds = [event["event_type"] for event in history] + assert kinds.count("TimerCancelled") == 1 + assert "TimerFired" not in kinds + finally: + await worker.stop() + + +async def test_leased_remote_activity_cannot_complete_after_delivery( + server_url: str, server_token: str, +) -> None: + queue = f"py-cooperative-remote-{uuid.uuid4().hex[:8]}" + cleanup_entered = asyncio.Event() + release_cleanup = asyncio.Event() + + async def cleanup(request_id: str) -> str: + cleanup_entered.set() + await release_cleanup.wait() + return await cooperative_cleanup(request_id) + + async with Client(server_url, token=server_token, namespace="default") as client: + worker = candidate_worker(client, queue) + await worker._register() + # This fixture invokes only local cleanup and claims remote work by API. + # Its in-process closure does not qualify callback process supervision. + worker.activities["tests.python-cooperative-cleanup"] = cleanup + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["remote"], + ) + task = await poll_claim(client, worker) + assert [command["type"] for command in await worker._run_workflow_task(task) or []] == ["schedule_activity"] + remote = await client.poll_activity_task(worker_id=worker.worker_id, task_queue=queue, timeout=5) + assert remote is not None + accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) + resumed = await poll_claim(client, worker) + execution = worker._track(worker._run_workflow_task(resumed)) + await asyncio.wait_for(cleanup_entered.wait(), timeout=10) + late_fence = { + "task_id": remote["task_id"], "activity_attempt_id": remote["activity_attempt_id"], + "lease_owner": worker.worker_id, + } + heartbeat = await client.heartbeat_activity_task(**late_fence) + assert heartbeat["cancel_requested"] is True + assert heartbeat["can_continue"] is False + assert heartbeat["heartbeat_recorded"] is False + release_cleanup.set() + assert await asyncio.wait_for(execution, timeout=10) is not None + before = await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + with pytest.raises(ServerError) as late_result: + await client.complete_activity_task(**late_fence, result="late remote result") + assert late_result.value.status == 409 + assert late_result.value.reason() == "run_cancelled" + with pytest.raises(ServerError) as late_failure: + await client.fail_activity_task( + **late_fence, message="qualification late failure", failure_type="qualification-late-failure", + ) + assert late_failure.value.status == 409 + assert late_failure.value.reason() == "run_cancelled" + assert await events(handle) == before + finally: + release_cleanup.set() + await worker.stop() + + +@pytest.mark.parametrize("lost_ack", [False, True]) +async def test_request_before_claim_retains_identity_and_proves_delivery( + server_url: str, server_token: str, lost_ack: bool, +) -> None: + queue = f"py-cooperative-request-{uuid.uuid4().hex[:8]}" + client_type = LostDeliveryAcknowledgmentClient if lost_ack else Client + async with client_type(server_url, token=server_token, namespace="default") as client: + worker = candidate_worker(client, queue) + await worker._register() + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["timer"], + ) + accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) + repeated = await handle.request_cancellation(cleanup_timeout_seconds=300) + assert accepted["duplicate"] is False + assert repeated["duplicate"] is True + assert repeated["cancellation_request"] == accepted["cancellation_request"] + task = await poll_claim(client, worker) + assert task["cancellation_request"] == accepted["cancellation_request"] + commands = await worker._run_workflow_task(task) + assert commands is not None + assert [command["type"] for command in commands] == ["record_local_activity", "complete_workflow"] + await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + if lost_ack: + assert isinstance(client, LostDeliveryAcknowledgmentClient) + assert client.delivery_calls == 1 + finally: + await worker.stop() + + +@pytest.mark.parametrize("handler_kind", ["async", "sync"]) +async def test_active_local_work_observes_request_and_discards_late_result( + server_url: str, server_token: str, handler_kind: str, +) -> None: + queue = f"py-cooperative-active-{uuid.uuid4().hex[:8]}" + entered = asyncio.Event() + discarded = asyncio.Event() + cleanup_entered = asyncio.Event() + release_thread = threading.Event() + loop = asyncio.get_running_loop() + + async def blocked_work() -> object: + entered.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + discarded.set() + return object() + + def synchronous_work() -> object: + loop.call_soon_threadsafe(entered.set) + assert release_thread.wait(timeout=20) + loop.call_soon_threadsafe(discarded.set) + return object() + + async def cleanup(request_id: str) -> str: + cleanup_entered.set() + return await cooperative_cleanup(request_id) + + async with Client(server_url, token=server_token, namespace="default") as client: + worker = candidate_worker(client, queue) + await worker._register() + # Exercise the existing local replay fence without polling remote work. + # Physical local callback supervision remains a separate qualification. + worker.activities["tests.python-cooperative-work"] = ( + blocked_work if handler_kind == "async" else synchronous_work + ) + worker.activities["tests.python-cooperative-cleanup"] = cleanup + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["local"], + ) + task = await poll_claim(client, worker) + execution = worker._track(worker._run_workflow_task(task)) + await asyncio.wait_for(entered.wait(), timeout=10) + accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) + await asyncio.wait_for(cleanup_entered.wait(), timeout=10) + release_thread.set() + await asyncio.wait_for(discarded.wait(), timeout=10) + assert await asyncio.wait_for(execution, timeout=10) is not None + history = await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + completed = [event for event in history if event["event_type"] == "ActivityCompleted"] + assert completed[0]["payload"]["activity_type"] == "tests.python-cooperative-cleanup" + finally: + release_thread.set() + await worker.stop() + + +@pytest.mark.parametrize("shutdown_kind", ["drain", "timeout"]) +async def test_shutdown_during_shielded_cleanup_reuses_canonical_delivery( + server_url: str, server_token: str, shutdown_kind: str, +) -> None: + queue = f"py-cooperative-shutdown-{uuid.uuid4().hex[:8]}" + entered = asyncio.Event() + release = asyncio.Event() + late_fenced = asyncio.Event() + + async def cleanup(request_id: str) -> object: + entered.set() + try: + await release.wait() + except asyncio.CancelledError: + with pytest.raises(LocalActivityExecutionAborted): + await activity.context().heartbeat() + late_fenced.set() + return object() + return await cooperative_cleanup(request_id) + + async with Client(server_url, token=server_token, namespace="default") as client: + worker = candidate_worker(client, queue, shutdown_timeout=5 if shutdown_kind == "drain" else 0.1) + await worker._register() + # Local-only replay fixture. It does not poll remote activity callbacks. + worker.activities["tests.python-cooperative-cleanup"] = cleanup + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["timer"], + ) + accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) + task = await poll_claim(client, worker) + execution = worker._track(worker._run_workflow_task(task)) + await asyncio.wait_for(entered.wait(), timeout=10) + before = await events(handle) + stop = asyncio.create_task(worker.stop()) + if shutdown_kind == "drain": + await asyncio.sleep(0) + release.set() + await asyncio.wait_for(stop, timeout=10) + if shutdown_kind == "timeout": + await asyncio.wait_for(late_fenced.wait(), timeout=10) + assert execution.cancelled() or execution.result() is None + after = await events(handle) + assert after[:len(before)] == before + assert [event["event_type"] for event in after[len(before):]] == ["RepairRequested"] + assert after[-1]["payload"]["command"]["request_method"] == "DELETE" + assert after[-1]["payload"]["command"]["request_path"].endswith(worker.worker_id) + replacement = candidate_worker(client, queue, worker_id=f"{queue}-replacement") + await replacement._register() + try: + reclaimed = await poll_claim(client, replacement) + assert reclaimed["workflow_task_attempt"] > task["workflow_task_attempt"] + assert reclaimed["lease_owner"] == replacement.worker_id + original = accepted["cancellation_request"] + assert reclaimed["cancellation_request"]["request_id"] == original["request_id"] + assert reclaimed["cancellation_request"]["cleanup_deadline_at"] == original["cleanup_deadline_at"] + assert await replacement._run_workflow_task(reclaimed) is not None + finally: + await replacement.stop() + else: + assert await execution is not None + await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + finally: + release.set() + await worker.stop() + + +async def native_process(queue: str, worker_id: str, mode: str) -> asyncio.subprocess.Process: + return await asyncio.create_subprocess_exec( + sys.executable, "-m", "tests.integration.cooperative_worker", queue, worker_id, mode, + stdout=asyncio.subprocess.PIPE, stderr=asyncio.subprocess.PIPE, + ) + + +async def process_event(process: asyncio.subprocess.Process, phase: str, timeout: float = 40) -> dict[str, Any]: + async def read() -> dict[str, Any]: + assert process.stdout is not None + while line := await process.stdout.readline(): + value = json.loads(line) + if value["phase"] == phase: + return value + assert process.stderr is not None + pytest.fail(f"worker exited before {phase}: {(await process.stderr.read()).decode()}") + return await asyncio.wait_for(read(), timeout=timeout) + + +async def test_sigkill_activity_owner_reclaims_attempt_before_cooperative_cleanup( + server_url: str, server_token: str, +) -> None: + queue = f"py-cooperative-activity-reclaim-{uuid.uuid4().hex[:8]}" + processes: list[asyncio.subprocess.Process] = [] + async with Client(server_url, token=server_token, namespace="default") as client: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["remote"], + ) + try: + owner = await native_process(queue, f"{queue}-killed", "remote") + processes.append(owner) + original = await process_event(owner, "remote-entered") + assert original["attempt_number"] == 1 + fence = {key: original[key] for key in ("task_id", "activity_attempt_id", "lease_owner")} + leased = await client.activity_task_status(**fence) + assert leased["can_continue"] is True + assert leased["attempt_status"] == "running" + expires_at = datetime.fromisoformat(leased["lease_expires_at"].replace("Z", "+00:00")).timestamp() + assert expires_at > time.time(), "kill an actually current lease" + print(f"SIGKILL original activity: {json.dumps([original, leased])}") + owner.kill() + assert await asyncio.wait_for(owner.wait(), timeout=10) == -9 + await callback_gone(original["callback_pid"]) + + successor = await native_process(queue, f"{queue}-successor", "remote") + processes.append(successor) + # Wait for Native's actual five-minute lease and normal repair. + # No clock, storage row or production lease is changed. + reclaimed = await process_event(successor, "remote-entered", timeout=330) + assert reclaimed["task_id"] == original["task_id"] + assert reclaimed["activity_attempt_id"] != original["activity_attempt_id"] + assert reclaimed["lease_owner"] != original["lease_owner"] + assert reclaimed["attempt_number"] == 2 + assert time.time() >= expires_at, "reclaim precedes actual lease expiry" + print(f"SIGKILL successor activity: {json.dumps(reclaimed)}") + + closed = await client.activity_task_status(**fence) + assert closed["attempt_status"] == "expired" + assert closed["can_continue"] is False + before = await events(handle) + with pytest.raises(ServerError) as completion: + await client.complete_activity_task(**fence, result="late") + assert completion.value.status == 409 + with pytest.raises(ServerError) as failure: + await client.fail_activity_task(**fence, message="late", failure_type="LateQualification") + assert failure.value.status == 409 + heartbeat = await client.heartbeat_activity_task(**fence) + assert heartbeat["can_continue"] is False + assert heartbeat["heartbeat_recorded"] is False + assert heartbeat["cancel_requested"] is False + assert heartbeat["reason"] == "attempt_closed" + assert heartbeat["lease_expires_at"] == closed["lease_expires_at"] + assert heartbeat["last_heartbeat_at"] == closed["last_heartbeat_at"] + assert await client.activity_task_status(**fence) == closed + assert await events(handle) == before, "dead attempt changed canonical history" + + accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) + history = await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + await callback_gone(reclaimed["callback_pid"]) + assert len([event for event in history if event["event_type"] == "ActivityCancelled"]) == 1 + print(f"SIGKILL reclaimed cancellation history: {json.dumps(history)}") + finally: + for process in processes: + if process.returncode is None: + process.kill() + await asyncio.wait_for(process.wait(), timeout=10) + + +async def test_killed_remote_owner_cannot_publish_after_cold_workflow_delivery( + server_url: str, server_token: str, +) -> None: + queue = f"py-cooperative-remote-process-{uuid.uuid4().hex[:8]}" + processes: list[asyncio.subprocess.Process] = [] + async with Client(server_url, token=server_token, namespace="default") as client: + seed = candidate_worker(client, queue, worker_id=f"{queue}-seed") + await seed._register() + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["remote"], + ) + owner = await native_process(queue, f"{queue}-killed", "remote") + processes.append(owner) + claimed = await process_event(owner, "remote-entered") + before = await events(handle) + owner.kill() + assert await asyncio.wait_for(owner.wait(), timeout=10) == -9 + await callback_gone(claimed["callback_pid"]) + assert await events(handle) == before + accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) + replacement = await native_process(queue, f"{queue}-replacement", "finish") + processes.append(replacement) + resumed = await process_event(replacement, "cleanup") + assert resumed["request_id"] == accepted["cancellation_request"]["request_id"] + assert (await process_event(replacement, "finished"))["committed"] is True + assert await asyncio.wait_for(replacement.wait(), timeout=10) == 0 + history = await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + assert len([event for event in history if event["event_type"] == "ActivityCancelled"]) == 1 + fence = {key: claimed[key] for key in ("task_id", "activity_attempt_id", "lease_owner")} + with pytest.raises(ServerError) as completion: + await client.complete_activity_task(**fence, result="late") + assert completion.value.status == 409 + with pytest.raises(ServerError) as failure: + await client.fail_activity_task(**fence, message="late", failure_type="LateQualification") + assert failure.value.status == 409 + assert await events(handle) == history + finally: + for process in processes: + if process.returncode is None: + process.kill() + await process.wait() + await seed.stop() + + +async def test_killed_process_reclaims_cleanup_in_a_new_process( + server_url: str, server_token: str, +) -> None: + queue = f"py-cooperative-process-{uuid.uuid4().hex[:8]}" + processes: list[asyncio.subprocess.Process] = [] + async with Client(server_url, token=server_token, namespace="default") as client: + seed = candidate_worker(client, queue, worker_id=f"{queue}-seed") + await seed._register() + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, + input=["clock_timer"], + ) + accepted = await handle.request_cancellation(cleanup_timeout_seconds=120) + original = accepted["cancellation_request"] + first = await native_process(queue, f"{queue}-killed", "hold") + processes.append(first) + claim = await process_event(first, "claim") + original_remaining = await process_event(first, "remaining-delivery") + cleanup = await process_event(first, "cleanup") + assert cleanup["request_id"] == original["request_id"] + before = await events(handle) + first.kill() + assert await asyncio.wait_for(first.wait(), timeout=10) == -9 + assert await events(handle) == before + replacement = await native_process(queue, f"{queue}-replacement", "finish") + processes.append(replacement) + reclaimed = await process_event(replacement, "claim") + assert reclaimed["attempt"] > claim["attempt"] + assert reclaimed["worker_id"] != claim["worker_id"] + assert reclaimed["request_id"] == original["request_id"] + assert reclaimed["cleanup_deadline_at"] == original["cleanup_deadline_at"] + replacement_remaining = await process_event(replacement, "remaining-delivery") + assert replacement_remaining == original_remaining + resumed = await process_event(replacement, "cleanup") + assert resumed["request_id"] == original["request_id"] + completed_remaining = await process_event(replacement, "remaining-cleanup") + assert completed_remaining["context"] == original_remaining["context"] + assert completed_remaining["remaining"] == original_remaining["remaining"] > 0 + assert (await process_event(replacement, "finished"))["committed"] is True + assert await asyncio.wait_for(replacement.wait(), timeout=10) == 0 + history = await assert_cancelled_cleanup(handle, original["request_id"]) + deadline = datetime.fromisoformat(original["cleanup_deadline_at"].replace("Z", "+00:00")) + delivered = next(event for event in history if event["event_type"] == "CooperativeCancellationDelivered") + recorded = datetime.fromisoformat(delivered["timestamp"].replace("Z", "+00:00")) + assert original_remaining["remaining"] == pytest.approx((deadline - recorded).total_seconds(), abs=1e-6) + remaining_observations = { + "original": original_remaining, "replacement": replacement_remaining, "completed": completed_remaining, + } + print(f"SIGKILL replay remaining-time observations: {json.dumps(remaining_observations)}") + finally: + for process in processes: + if process.returncode is None: + process.kill() + await process.wait() + await seed.stop() + + +@pytest.mark.parametrize("close_kind", ["deadline", "terminate"]) +async def test_cleanup_deadline_and_termination_fence_in_flight_local_result( + server_url: str, server_token: str, close_kind: str, +) -> None: + queue = f"py-cooperative-close-{uuid.uuid4().hex[:8]}" + entered = asyncio.Event() + + async def cleanup(request_id: str) -> object: + entered.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + return object() + + async with Client(server_url, token=server_token, namespace="default") as client: + worker = candidate_worker(client, queue) + await worker._register() + # Local-only replay fixture. It does not poll remote activity callbacks. + worker.activities["tests.python-cooperative-cleanup"] = cleanup + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["timer"], + ) + await handle.request_cancellation(cleanup_timeout_seconds=2 if close_kind == "deadline" else 60) + task = await poll_claim(client, worker) + execution = worker._track(worker._run_workflow_task(task)) + await asyncio.wait_for(entered.wait(), timeout=10) + if close_kind == "terminate": + await handle.terminate(reason="qualification termination during cleanup") + assert await asyncio.wait_for(execution, timeout=15) is None + history = await events(handle) + kinds = [event["event_type"] for event in history] + assert kinds.count("WorkflowCancelled" if close_kind == "deadline" else "WorkflowTerminated") == 1 + assert kinds.count("CooperativeCancellationRequested") == 1 + assert kinds.count("CooperativeCancellationDelivered") == 1 + assert "ActivityCompleted" not in kinds + assert "ActivityFailed" not in kinds + assert "WorkflowCompleted" not in kinds + finally: + await worker.stop() diff --git a/tests/integration/test_prepared_local_activity.py b/tests/integration/test_prepared_local_activity.py new file mode 100644 index 0000000..9c0e391 --- /dev/null +++ b/tests/integration/test_prepared_local_activity.py @@ -0,0 +1,359 @@ +"""Opt-in exact-source proof of Python's durable sequential local consumer.""" + +from __future__ import annotations + +import asyncio +import json +import os +import sys +import uuid +from datetime import datetime +from pathlib import Path +from typing import Any + +import pytest + +from durable_workflow import Client, Worker, activity, workflow +from durable_workflow.errors import ServerError, WorkflowCancelled +from durable_workflow.workflow import replay +from tests.integration.test_cooperative_cancellation import callback_gone, events, poll_claim, remote_marker + +pytestmark = pytest.mark.usefixtures("prepared_runtime") + + +@pytest.fixture +async def prepared_runtime(server_url: str, server_token: str, monkeypatch: pytest.MonkeyPatch) -> None: + if os.environ.get("DURABLE_WORKFLOW_NATIVE_SOURCE_QUALIFICATION") != "1": + pytest.skip("prepared local consumer requires exact-source Native qualification") + async with Client(server_url, token=server_token, namespace="default") as client: + info = await client.get_cluster_info() + assert info["worker_protocol"]["server_capabilities"]["prepared_local_activities"] is True + assert info["worker_protocol"]["server_capabilities"]["prepared_local_activity_groups"] is True + assert info["worker_protocol"]["server_capabilities"]["prepared_local_activity_cancellation_policies"] == [ + "try_cancel", "wait_cancellation_completed", + ] + assert info["worker_protocol"]["version"] == "1.20" + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + + +@workflow.defn(name="tests.python-prepared-sequential") +class PreparedSequentialWorkflow: + def run(self, ctx: Any, marker: str) -> Any: + yield ctx.upsert_memo({"before": "durable local admission"}) + first = yield ctx.local_activity("tests.python-prepared-local", [marker, "first"], heartbeat_timeout=10) + second = yield ctx.local_activity("tests.python-prepared-local", [marker, "second"], heartbeat_timeout=10) + return [first, second] + + +@activity.defn(name="tests.python-prepared-local") +async def prepared_local(marker: str, phase: str) -> dict[str, Any]: + info = activity.context().info + Path(marker + "." + phase).write_text(json.dumps({ + "callback_pid": os.getpid(), "attempt_id": info.activity_attempt_id, + })) + await activity.context().heartbeat({"phase": phase}) + return {"phase": phase, "bytes": b"\x00\xff", "attempt": info.activity_attempt_id} + + +@workflow.defn(name="tests.python-prepared-cancellation") +class PreparedCancellationWorkflow: + def run(self, ctx: Any, marker: str, group: bool = False, policy: str | None = None) -> Any: + try: + if group: + yield [ctx.local_activity("tests.python-prepared-blocked", [marker, "work-0", None], + cancellation_policy=policy), + ctx.local_activity("tests.python-prepared-blocked", [marker, "work-1", None], + cancellation_policy=policy)] + else: + yield ctx.local_activity("tests.python-prepared-blocked", [marker, "work", None], + cancellation_policy=policy) + except WorkflowCancelled as error: + assert error.context is ctx.cancellation_context and error.context is not None + delivery_path = Path(marker + ".remaining-delivery-" + str(os.getpid()) + ".json") + pending = delivery_path.with_suffix(".writing") + pending.write_text(json.dumps({"context": error.context.to_dict(), "remaining": error.context.remaining()})) + pending.replace(delivery_path) + with ctx.cancellation_shield(): + if group: + yield [ctx.local_activity( + "tests.python-prepared-blocked", [marker, "cleanup-" + str(index), error.request_id], + retry_policy={"max_attempts": 2, "backoff_seconds": [0]}, + cancellation_policy=policy, + ) for index in range(2)] + else: + yield ctx.local_activity( + "tests.python-prepared-blocked", [marker, "cleanup", error.request_id], + retry_policy={"max_attempts": 2, "backoff_seconds": [0]}, + cancellation_policy=policy, + ) + final_path = Path(marker + ".remaining-final.json") + pending = final_path.with_suffix(".writing") + pending.write_text(json.dumps({"context": error.context.to_dict(), "remaining": error.context.remaining()})) + pending.replace(final_path) + return error.request_id + return "not cancelled" + + +@activity.defn(name="tests.python-prepared-blocked") +async def prepared_blocked(marker: str, phase: str, request_id: str | None) -> dict[str, Any]: + info = activity.context().info + path = Path(marker + "." + phase + "." + info.worker_id) + pending = path.with_suffix(path.suffix + ".writing") + pending.write_text(json.dumps({"phase": phase, "callback_pid": os.getpid(), "request_id": request_id, + "activity_attempt_id": info.activity_attempt_id})) + pending.replace(path) + if phase.startswith("work") or os.environ.get("DW_PREPARED_FIXTURE_MODE") == "hold": + # Intentionally no application heartbeat. + await asyncio.Event().wait() + return {"request_id": request_id, "bytes": b"\x00\xff"} + + +@workflow.defn(name="tests.python-prepared-group") +class PreparedGroupWorkflow: + def run(self, ctx: Any, marker: str) -> Any: + yield ctx.upsert_memo({"before": "atomic local group"}) + first = ctx.local_activity("tests.python-prepared-peer", [marker, "first", "second"], heartbeat_timeout=10) + second = ctx.local_activity("tests.python-prepared-peer", [marker, "second", "first"], heartbeat_timeout=10) + return (yield [first, [second, ctx.start_timer(1)]]) + + +@activity.defn(name="tests.python-prepared-peer") +async def prepared_peer(marker: str, phase: str, peer: str) -> dict[str, Any]: + info = activity.context().info + Path(marker + "." + phase).write_text(json.dumps({ + "callback_pid": os.getpid(), "attempt_id": info.activity_attempt_id, + })) + await remote_marker(Path(marker + "." + peer)) + await activity.context().heartbeat({"phase": phase}) + return {"phase": phase, "bytes": b"\x00\xff", "attempt": info.activity_attempt_id} + + +def prepared_worker(client: Client, queue: str, **kwargs: Any) -> Worker: + return Worker( + client, task_queue=queue, worker_id=kwargs.pop("worker_id", queue + "-owner"), + workflows=[PreparedSequentialWorkflow, PreparedCancellationWorkflow, PreparedGroupWorkflow], + activities=[prepared_local, prepared_blocked, prepared_peer], + capabilities=["cooperative_cancellation", "prepared_local_activities", "prepared_local_activity_groups"], + **kwargs, + ) + + +class TrackingPreparedClient(Client): + def __init__(self, *args: Any, trace_path: Path, **kwargs: Any) -> None: + super().__init__(*args, **kwargs) + self.trace_path = trace_path + + async def prepared_local_activity_operation(self, **kwargs: Any) -> dict[str, Any]: + try: + receipt = await super().prepared_local_activity_operation(**kwargs) + except ServerError as error: + with self.trace_path.open("a") as stream: + stream.write(json.dumps({"operation": kwargs["operation"], "sequence": kwargs["body"].get("sequence"), + "refused": error.reason(), "status": error.status, + "response": error.body}) + "\n") + raise + if kwargs["operation"] != "control" or receipt.get("active") is False: + with self.trace_path.open("a") as stream: + stream.write(json.dumps({"operation": kwargs["operation"], "receipt": receipt}) + "\n") + return receipt + + +async def test_prepared_prefix_two_callbacks_and_cold_replay_use_canonical_history( + server_url: str, server_token: str, tmp_path: Path, +) -> None: + queue = "py-prepared-sequential-" + uuid.uuid4().hex[:8] + marker = str(tmp_path / "callback") + async with Client(server_url, token=server_token, namespace="default") as client: + worker = prepared_worker(client, queue) + await worker._register() + try: + handle = await client.start_workflow(workflow_type="tests.python-prepared-sequential", task_queue=queue, + workflow_id=queue, input=[marker]) + task = await poll_claim(client, worker) + commands = await worker._run_workflow_task(task) + assert commands is not None and [command["type"] for command in commands] == ["complete_workflow"] + result = await handle.result(timeout=10) + assert [item["phase"] for item in result] == ["first", "second"] + assert all(item["bytes"] == b"\x00\xff" for item in result) + assert result[0]["attempt"] != result[1]["attempt"] + history = await events(handle) + kinds = [event["event_type"] for event in history] + assert kinds.count("ActivityStarted") == kinds.count("ActivityCompleted") == 2 + assert kinds.count("ActivityHeartbeatRecorded") == 2 + assert sorted(event["payload"]["progress"]["details"]["phase"] for event in history + if event["event_type"] == "ActivityHeartbeatRecorded") == ["first", "second"] + assert kinds.count("MemoUpserted") == kinds.count("WorkflowCompleted") == 1 + outcome = replay(PreparedSequentialWorkflow, history, [marker], run_id=handle.run_id or "", + prepare_local_activities=True) + assert outcome.prepared_local_activity is None and outcome.commands[0].result == result # type: ignore[union-attr] + for phase in ("first", "second"): + proof = json.loads(Path(marker + "." + phase).read_text()) + await callback_gone(proof["callback_pid"]) + print("prepared sequential history: " + json.dumps(history)) + finally: + await worker.stop() + + +async def test_prepared_nested_mixed_group_starts_peers_concurrently_and_cold_replays_results_in_position( + server_url: str, server_token: str, tmp_path: Path, +) -> None: + queue = "py-prepared-group-" + uuid.uuid4().hex[:8] + marker = str(tmp_path / "callback") + trace_path = tmp_path / "group-receipts.jsonl" + async with TrackingPreparedClient( + server_url, token=server_token, namespace="default", trace_path=trace_path, + ) as client: + worker = prepared_worker(client, queue) + owner = asyncio.create_task(worker.run()) + try: + handle = await client.start_workflow(workflow_type="tests.python-prepared-group", task_queue=queue, + workflow_id=queue, input=[marker]) + result = await handle.result(timeout=15) + assert result[0]["phase"] == "first" and result[1][0]["phase"] == "second" + assert result[0]["bytes"] == result[1][0]["bytes"] == b"\x00\xff" + assert result[0]["attempt"] != result[1][0]["attempt"] + history = await events(handle) + kinds = [event["event_type"] for event in history] + assert kinds.count("ActivityScheduled") == kinds.count("ActivityStarted") == 2 + assert kinds.count("ActivityCompleted") == 2 + assert kinds.count("ActivityHeartbeatRecorded") == 2 + assert sorted(event["payload"]["progress"]["details"]["phase"] for event in history + if event["event_type"] == "ActivityHeartbeatRecorded") == ["first", "second"] + assert kinds.count("TimerScheduled") == kinds.count("TimerFired") == 1 + assert kinds.count("MemoUpserted") == kinds.count("WorkflowCompleted") == 1 + completed = replay(PreparedGroupWorkflow, history, [marker], run_id=handle.run_id or "", + prepare_local_activities=True, prepare_local_activity_groups=True) + assert completed.prepared_local_activity_group is None + assert completed.commands[0].result == result # type: ignore[union-attr] + for phase in ("first", "second"): + await callback_gone(json.loads(Path(marker + "." + phase).read_text())["callback_pid"]) + print("prepared nested mixed group history: " + json.dumps(history)) + finally: + await worker.stop() + await asyncio.wait_for(owner, timeout=10) + if trace_path.exists(): + print("prepared nested mixed group receipts: " + trace_path.read_text()) + print("prepared nested mixed group final history: " + json.dumps(await events(handle))) + + +@pytest.mark.parametrize("group", [False, True], ids=["sequential", "atomic-group"]) +@pytest.mark.parametrize("policy", [None, "try_cancel", "wait_cancellation_completed"], + ids=["historical-omission", "try-cancel", "wait-cancellation-completed"]) +async def test_prepared_callback_stop_cleanup_sigkill_and_cold_recovery_keep_original_30_second_deadline( + server_url: str, server_token: str, tmp_path: Path, group: bool, policy: str | None, +) -> None: + queue = "py-prepared-cancel-" + uuid.uuid4().hex[:8] + marker = str(tmp_path / "callback") + trace_path = tmp_path / "receipts.jsonl" + processes: list[asyncio.subprocess.Process] = [] + logs: list[Any] = [] + + async def owner(name: str, mode: str) -> asyncio.subprocess.Process: + log = (tmp_path / (name + ".log")).open("wb") + logs.append(log) + process = await asyncio.create_subprocess_exec( + sys.executable, "-m", "tests.integration.prepared_worker", queue, name, mode, str(trace_path), + stdout=log, stderr=log, + ) + processes.append(process) + return process + + async with Client(server_url, token=server_token, namespace="default") as client: + try: + handle = await client.start_workflow(workflow_type="tests.python-prepared-cancellation", task_queue=queue, + workflow_id=queue, input=[marker, group, policy]) + first = await owner(queue + "-first", "hold") + work_phases = ["work-0", "work-1"] if group else ["work"] + cleanup_phases = ["cleanup-0", "cleanup-1"] if group else ["cleanup"] + work = [await remote_marker(Path(marker + "." + phase + "." + queue + "-first")) + for phase in work_phases] + accepted = await handle.request_cancellation(cleanup_timeout_seconds=30) + original = accepted["cancellation_request"] + requested_history = await events(handle) + original_context = next(event["payload"]["cancellation"] for event in requested_history + if event["event_type"] == "CooperativeCancellationRequested") + cleanup = [await remote_marker(Path(marker + "." + phase + "." + queue + "-first")) + for phase in cleanup_phases] + assert all(item["request_id"] == original["request_id"] for item in cleanup) + original_remaining = await remote_marker(Path(marker + ".remaining-delivery-" + str(first.pid) + ".json")) + assert original_remaining["context"] == original_context + for item in work: + await callback_gone(item["callback_pid"]) + first.kill() + await asyncio.wait_for(first.wait(), timeout=10) + for item in cleanup: + await callback_gone(item["callback_pid"]) + duplicate = await handle.request_cancellation(cleanup_timeout_seconds=90) + for field in ("request_id", "requested_at", "cleanup_deadline_at"): + assert duplicate["cancellation_request"][field] == original[field] + successor = await owner(queue + "-replacement", "finish") + deadline = datetime.fromisoformat(original["cleanup_deadline_at"].replace("Z", "+00:00")) + timeout = (deadline - datetime.now(deadline.tzinfo)).total_seconds() + assert timeout > 0 + with pytest.raises(WorkflowCancelled): + await handle.result(timeout=timeout) + history = await events(handle) + replacement_remaining = await remote_marker( + Path(marker + ".remaining-delivery-" + str(successor.pid) + ".json"), + ) + completed_remaining = await remote_marker(Path(marker + ".remaining-final.json")) + assert replacement_remaining == original_remaining + assert completed_remaining["context"] == original_context + assert 0 < completed_remaining["remaining"] < original_remaining["remaining"] + delivered = next(event for event in history if event["event_type"] == "CooperativeCancellationDelivered") + completed = max(datetime.fromisoformat(event["timestamp"].replace("Z", "+00:00")) + for event in history if event["event_type"] == "ActivityCompleted") + delivered_at = datetime.fromisoformat(delivered["timestamp"].replace("Z", "+00:00")) + assert original_remaining["remaining"] == pytest.approx((deadline - delivered_at).total_seconds(), abs=1e-6) + assert completed_remaining["remaining"] == pytest.approx((deadline - completed).total_seconds(), abs=1e-6) + kinds = [event["event_type"] for event in history] + assert kinds.count("CooperativeCancellationDelivered") == kinds.count("WorkflowCancelled") == 1 + assert kinds.count("ActivityCancellationAcknowledged") == len(work) + assert kinds.count("ActivityCancelled") == len(work) + assert "ActivityHeartbeatRecorded" not in kinds + scheduled = [event for event in history if event["event_type"] == "ActivityScheduled"] + assert len(scheduled) == len(work) + len(cleanup) + assert all( + event["payload"]["activity"].get("cancellation_policy", "try_cancel") == (policy or "try_cancel") + for event in scheduled + ) + if policy == "wait_cancellation_completed": + delivery_index = kinds.index("CooperativeCancellationDelivered") + assert all(index < delivery_index for index, kind in enumerate(kinds) + if kind == "ActivityCancellationAcknowledged") + terminal = next(event for event in history if event["event_type"] == "WorkflowCancelled") + assert datetime.fromisoformat(terminal["timestamp"].replace("Z", "+00:00")) < deadline + receipts = [json.loads(line) for line in trace_path.read_text().splitlines()] + admissions = [entry["receipt"] for entry in receipts + if entry["operation"] == "prepare" and entry["receipt"].get("cancellation_cleanup")] + assert len(admissions) == 2 * len(cleanup) + assert len({item["activity_attempt_id"] for item in admissions}) == len(admissions) + assert all(item["cancellation_cleanup"] == admissions[0]["cancellation_cleanup"] for item in admissions) + authority = admissions[0]["cancellation_cleanup"] + assert authority["request_id"] == original["request_id"] + assert authority["root_request_id"] == original_context["root_request_id"] + assert datetime.fromisoformat(authority["cleanup_deadline_at"].replace("Z", "+00:00")) == deadline + recoveries = [entry["receipt"] for entry in receipts if entry["operation"] == "recover"] + assert len(recoveries) == len(cleanup) + assert all(item["callback_stop_state"] == "unknown" for item in recoveries) + replacement = [await remote_marker(Path(marker + "." + phase + "." + queue + "-replacement")) + for phase in cleanup_phases] + for item in replacement: + await callback_gone(item["callback_pid"]) + print("prepared physically joined callbacks: " + json.dumps({ + "work": work, "killed_cleanup": cleanup, "replacement_cleanup": replacement, + })) + print("prepared remaining-time observations: " + json.dumps({ + "original": original_remaining, "replacement": replacement_remaining, "completed": completed_remaining, + })) + print("prepared cleanup recovery receipts: " + json.dumps(receipts)) + print("prepared cleanup recovery history: " + json.dumps(history)) + finally: + for process in processes: + if process.returncode is None: + process.kill() + await asyncio.wait_for(process.wait(), timeout=10) + for log in logs: + log.close() + for path in tmp_path.glob("*.log"): + print(path.name + ": " + path.read_text()) diff --git a/tests/test_activity_cancellation_policies.py b/tests/test_activity_cancellation_policies.py new file mode 100644 index 0000000..f08dfdb --- /dev/null +++ b/tests/test_activity_cancellation_policies.py @@ -0,0 +1,136 @@ +from __future__ import annotations + +from typing import Any + +import pytest + +from durable_workflow import CancellationPolicy, ParentClosePolicy, serializer +from durable_workflow import workflow as workflow_module +from durable_workflow.errors import NonDeterministicReplayError, WorkflowCancelled +from durable_workflow.worker import Worker +from durable_workflow.workflow import ( + CompleteWorkflow, + ScheduleActivity, + WorkflowContext, + commands_to_server_commands, + replay, +) +from tests.test_cooperative_cancellation import marker, request +from tests.test_cooperative_cancellation_worker import ClaimServer, claimed_task + + +def workflow(options: dict[str, Any], mode: str = "sequential") -> type: + class PolicyWorkflow: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + first = ctx.schedule_activity("work", [], **options) + try: + if mode == "parallel": + return (yield [first]) + if mode == "selection": + return (yield ctx.select({"work": first})) + return (yield first) + except WorkflowCancelled as cancelled: + return cancelled.request_id + + return PolicyWorkflow + + +def event(kind: str, policy: Any = None) -> dict[str, Any]: + activity = {"type": "work"} + if policy is not None: + activity["cancellation_policy"] = policy + return {"event_type": kind, "payload": {"sequence": 1, "activity": activity}} + + +def test_changed_activity_policy_is_rejected_during_replay() -> None: + with pytest.raises(NonDeterministicReplayError, match="activity_cancellation_policy_changed"): + replay(workflow({}), [event("ActivityScheduled", "wait_cancellation_completed")], []) + + +@pytest.mark.parametrize("policy", list(CancellationPolicy)) +def test_typed_and_string_policies_match_both_encoders_and_replay(policy: CancellationPolicy) -> None: + command = ScheduleActivity("work", [], schedule_to_close_timeout=60, cancellation_policy=policy) + direct = command.to_server_command("queue") + assert direct == commands_to_server_commands([command], "queue")[0] + assert direct == ScheduleActivity( + "work", [], schedule_to_close_timeout=60, cancellation_policy=policy.value, + ).to_server_command("queue") + assert type(direct["cancellation_policy"]) is str + history = [event("ActivityScheduled", policy.value), event("ActivityStarted"), event("ActivityCompleted")] + history[-1]["payload"]["result"] = serializer.envelope("recorded") + assert replay(workflow({"cancellation_policy": policy, "schedule_to_close_timeout": 60}), history, []).commands == [ + CompleteWorkflow("recorded"), + ] + + +def test_omitted_policy_preserves_wire_and_historical_try_default() -> None: + assert "cancellation_policy" not in ScheduleActivity("work", []).to_server_command("queue") + history = [event("ActivityCompleted")] + history[0]["payload"]["result"] = serializer.envelope("recorded") + assert replay(workflow({}), history, []).commands == [CompleteWorkflow("recorded")] + assert replay(workflow({"cancellation_policy": CancellationPolicy.TRY_CANCEL}), history, []).commands == [ + CompleteWorkflow("recorded"), + ] + + +@pytest.mark.parametrize("mode", ["sequential", "parallel", "selection"]) +@pytest.mark.parametrize("delivered", [False, True]) +def test_changed_policy_fails_group_matching_and_cancellation_delivery(mode: str, delivered: bool) -> None: + options = {"cancellation_policy": CancellationPolicy.WAIT_CANCELLATION_COMPLETED} + commands = commands_to_server_commands(replay(workflow(options, mode), [], []).commands, "queue") + history = [{"event_type": "ActivityScheduled", "payload": {"sequence": 1, **command}} for command in commands] + if delivered: + history += [request(), marker(1, "activity" if mode == "sequential" else "parallel", sequence_span=1)] + assert replay(workflow(options, mode), history, [], run_id="run-1").commands == [CompleteWorkflow("request-1")] + with pytest.raises(NonDeterministicReplayError, match="activity_cancellation_policy_changed"): + replay(workflow({"cancellation_policy": CancellationPolicy.TRY_CANCEL}, mode), history, [], run_id="run-1") + + +@pytest.mark.parametrize("value", ["unknown", "", True, 1, [], ParentClosePolicy.ABANDON]) +def test_invalid_policy_is_refused_before_suspension(value: Any) -> None: + with pytest.raises(ValueError, match="cancellation_policy"): + ScheduleActivity("work", [], cancellation_policy=value) + + +@pytest.mark.parametrize("timeout", [None, 0, -1, 1.5, "60", True]) +def test_abandon_requires_finite_total_timeout(timeout: Any) -> None: + with pytest.raises(ValueError, match="schedule_to_close_timeout"): + ScheduleActivity("work", [], cancellation_policy=CancellationPolicy.ABANDON, schedule_to_close_timeout=timeout) + + +@pytest.mark.parametrize("kind", ["ActivityScheduled", "ActivityCompleted"]) +@pytest.mark.parametrize("value", [None, "unknown", []]) +def test_malformed_canonical_history_is_rejected(kind: str, value: Any) -> None: + item = event(kind) + item["payload"]["activity"]["cancellation_policy"] = value + with pytest.raises(NonDeterministicReplayError, match="invalid_activity_cancellation_policy_history"): + replay(workflow({}), [item], []) + + +def test_conflicting_history_is_rejected() -> None: + with pytest.raises(NonDeterministicReplayError, match="activity_cancellation_policy_history_conflict"): + replay(workflow({}), [event("ActivityScheduled", "try_cancel"), event("ActivityStarted", "abandon")], []) + + +@pytest.mark.parametrize("protocol", ["1.19", "1.20"]) +@pytest.mark.parametrize("policy", list(CancellationPolicy)) +async def test_worker_without_opt_in_refuses_explicit_policy( + monkeypatch: pytest.MonkeyPatch, protocol: str, policy: CancellationPolicy, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", protocol) + cls = workflow_module.defn(name="activity-policy-worker")(workflow({ + "cancellation_policy": policy, "schedule_to_close_timeout": 60, + })) + server = ClaimServer(history=[]) + worker = Worker(server.client, task_queue="queue", worker_id="cooperative-worker", workflows=[cls]) + await worker._register() + result = await worker._run_workflow_task({ + **claimed_task(observed=False), "workflow_type": "activity-policy-worker", "arguments": serializer.envelope([]), + }) + assert result is None + server.client.complete_workflow_task.assert_not_awaited() + failure = server.client.fail_workflow_task.await_args.kwargs + assert failure["failure_type"] == "RuntimeCapabilityUnsupported" + assert "activity_cancellation_policy_not_supported" in failure["message"] + assert "cooperative-worker" in failure["message"] + assert "1.20" in failure["message"] diff --git a/tests/test_activity_process.py b/tests/test_activity_process.py new file mode 100644 index 0000000..014b7f4 --- /dev/null +++ b/tests/test_activity_process.py @@ -0,0 +1,263 @@ +from __future__ import annotations + +import asyncio +import ctypes +import os +import signal +import subprocess +import sys +import time +from pathlib import Path +from typing import Any + +import pytest + +from durable_workflow import activity +from durable_workflow._activity_process import ( + CallbackInvocation, + CallbackProcessLost, + SupervisedCallback, +) +from durable_workflow.activity import ActivityInfo +from durable_workflow.errors import NonRetryableError +from durable_workflow.interceptors import ActivityHandler, ActivityInterceptorContext, PassthroughWorkerInterceptor + + +def invocation(handler: Any, *args: Any, interceptors: tuple[Any, ...] = ()) -> CallbackInvocation: + return CallbackInvocation( + handler=handler, args=args, + info=ActivityInfo("task", "process-test", "attempt", 1, "queue", "owner"), + task={"task_id": "task", "activity_attempt_id": "attempt"}, interceptors=interceptors, + ) + + +def typed_result(value: bytes) -> dict[str, Any]: + return {"bytes": value, "task": activity.context().info.task_id, "nested": [1, None, {"x": True}]} + + +async def authored_heartbeat() -> bytes: + await activity.context().heartbeat({"progress": b"typed"}) + return b"result" + + +def authored_sync_heartbeat() -> bytes: + return asyncio.run(authored_heartbeat()) + + +class ProcessInterceptor(PassthroughWorkerInterceptor): + async def execute_activity(self, context: ActivityInterceptorContext, next: ActivityHandler) -> Any: + assert context.worker_id == "owner" + assert context.task["activity_attempt_id"] == "attempt" + return {"intercepted": await next(context)} + + +class ProcessFailure(NonRetryableError): + code = 42 + + +def failing() -> None: + raise ProcessFailure("original failure") + + +async def blocked_without_python_progress(marker: str) -> None: + signal.signal(signal.SIGTERM, signal.SIG_IGN) + Path(marker).write_text(str(os.getpid())) + # PyDLL retains the callback interpreter's GIL during the C call. A Python + # thread or signal callback in this interpreter cannot supervise this work. + ctypes.PyDLL(None).sleep(60) + + +def wait_for_release(marker: str) -> None: + Path(marker).write_text(str(os.getpid())) + while not Path(marker + ".release").exists(): + time.sleep(0.02) + + +async def wait_for_file(path: Path) -> None: + async def ready() -> None: + while not path.exists(): + await asyncio.sleep(0.02) + await asyncio.wait_for(ready(), timeout=5.0) + + +async def wait_for_exit(pid: int) -> None: + async def gone() -> None: + while True: + try: + os.kill(pid, 0) + except ProcessLookupError: + return + await asyncio.sleep(0.02) + await asyncio.wait_for(gone(), timeout=7.0) + + +async def no_heartbeat(details: dict[str, Any] | None) -> None: + raise AssertionError("callback emitted an unauthored heartbeat") + + +async def test_spawn_preserves_typed_result_context_and_interceptors() -> None: + callback = SupervisedCallback(invocation(typed_result, b"\x00\xff", interceptors=(ProcessInterceptor(),))) + try: + await callback.start() + outcome = await callback.result(no_heartbeat) + assert outcome.failure is None + assert outcome.value == {"intercepted": {"bytes": b"\x00\xff", "task": "task", + "nested": [1, None, {"x": True}]}} + assert callback.stopped is True + assert not Path(callback.directory).exists() + finally: + await callback.close() + + +@pytest.mark.parametrize("handler", [authored_heartbeat, authored_sync_heartbeat]) +async def test_only_authored_heartbeat_crosses_to_the_owner(handler: Any) -> None: + observed: list[Any] = [] + owner_loop = asyncio.get_running_loop() + + async def heartbeat(details: dict[str, Any] | None) -> None: + assert asyncio.get_running_loop() is owner_loop + observed.append(details) + + callback = SupervisedCallback(invocation(handler)) + try: + await callback.start() + outcome = await callback.result(heartbeat) + assert outcome.value == b"result" + assert outcome.failure is None + assert observed == [{"progress": b"typed"}] + assert callback.stopped is True + finally: + await callback.close() + + +async def test_spawn_preserves_application_failure_metadata() -> None: + callback = SupervisedCallback(invocation(failing)) + try: + await callback.start() + outcome = await callback.result(no_heartbeat) + assert outcome.failure is not None + assert outcome.failure.message == "original failure" + assert outcome.failure.failure_type == "ProcessFailure" + assert outcome.failure.failure_class == "tests.test_activity_process.ProcessFailure" + assert outcome.failure.failure_code == 42 + assert outcome.failure.non_retryable is True + assert "raise ProcessFailure" in outcome.failure.stack_trace + assert callback.stopped is True + finally: + await callback.close() + + +@pytest.mark.skipif(os.name != "posix", reason="POSIX process signals") +async def test_stop_kills_and_joins_a_gil_blocked_callback_without_heartbeats(tmp_path: Path) -> None: + marker = tmp_path / "callback" + callback = SupervisedCallback(invocation(blocked_without_python_progress, str(marker))) + try: + await callback.start() + await wait_for_file(marker) + pid = int(marker.read_text()) + started = time.monotonic() + await callback.stop() + assert time.monotonic() - started < 3.0 + assert callback.stopped is True + await wait_for_exit(pid) + assert not Path(callback.directory).exists() + finally: + await callback.close() + + +@pytest.mark.skipif(os.name != "posix", reason="POSIX process signals") +async def test_owner_sigkill_leaves_supervisor_to_stop_and_reap_callback(tmp_path: Path) -> None: + marker = tmp_path / "callback" + supervisor_marker = tmp_path / "supervisor" + script = tmp_path / "owner.py" + script.write_text("""import asyncio +import sys +from pathlib import Path +from durable_workflow._activity_process import SupervisedCallback +from tests.test_activity_process import blocked_without_python_progress, invocation +async def run(): + callback = SupervisedCallback(invocation(blocked_without_python_progress, sys.argv[1])) + await callback.start() + Path(sys.argv[2]).write_text(str(callback._supervisor.pid) + '\\n' + callback.directory) + await asyncio.Event().wait() +if __name__ == '__main__': + asyncio.run(run()) +""") + environment = {**os.environ, "PYTHONPATH": os.pathsep.join(sys.path)} + owner = subprocess.Popen([sys.executable, str(script), str(marker), str(supervisor_marker)], env=environment) + try: + await wait_for_file(marker) + await wait_for_file(supervisor_marker) + supervisor_pid, directory = supervisor_marker.read_text().splitlines() + callback_pid = int(marker.read_text()) + os.kill(owner.pid, signal.SIGKILL) + await asyncio.to_thread(owner.wait, 5.0) + assert owner.returncode == -signal.SIGKILL + await wait_for_exit(callback_pid) + await wait_for_exit(int(supervisor_pid)) + assert not Path(directory).exists() + finally: + if owner.poll() is None: + owner.kill() + await asyncio.to_thread(owner.wait, 5.0) + + +@pytest.mark.skipif(os.name != "posix", reason="POSIX process signals") +async def test_dead_supervisor_with_live_callback_never_proves_stop(tmp_path: Path) -> None: + marker = tmp_path / "callback" + callback = SupervisedCallback(invocation(wait_for_release, str(marker))) + try: + await callback.start() + await wait_for_file(marker) + pid = int(marker.read_text()) + callback._supervisor.kill() + with pytest.raises((CallbackProcessLost, BrokenPipeError, ConnectionResetError)): + await callback.stop() + assert callback.stopped is False + os.kill(pid, 0) # A dead supervisor alone left the callback alive. + finally: + Path(str(marker) + ".release").touch() + if marker.exists(): + await wait_for_exit(int(marker.read_text())) + await callback.close() + + +def test_nonimportable_handler_is_rejected_before_spawn() -> None: + def callback() -> None: + pass + with pytest.raises(ValueError, match="spawn-compatible.*importable handler"): + SupervisedCallback(invocation(callback)) + + +async def test_cancellation_during_confirmed_result_join_still_reaps_and_proves_stop( + monkeypatch: pytest.MonkeyPatch, +) -> None: + callback = SupervisedCallback(invocation(typed_result, b"typed")) + entering_join = asyncio.Event() + release_join = asyncio.Event() + close = SupervisedCallback.close + + async def paused_close(self: SupervisedCallback) -> None: + entering_join.set() + await release_join.wait() + await close(self) + + monkeypatch.setattr(SupervisedCallback, "close", paused_close) + result: asyncio.Task[Any] | None = None + try: + await callback.start() + result = asyncio.create_task(callback.result(no_heartbeat)) + await asyncio.wait_for(entering_join.wait(), timeout=5) + result.cancel() + await asyncio.sleep(0) + assert not callback.stopped + release_join.set() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(result, timeout=7) + assert callback.stopped and callback._joined and callback._exitcode == 0 + finally: + release_join.set() + if result is not None: + await asyncio.gather(result, return_exceptions=True) + if not callback.stopped: + await callback.stop() diff --git a/tests/test_cancellation_context.py b/tests/test_cancellation_context.py new file mode 100644 index 0000000..2efe3cd --- /dev/null +++ b/tests/test_cancellation_context.py @@ -0,0 +1,203 @@ +from __future__ import annotations + +import json +from dataclasses import FrozenInstanceError +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +import pytest + +from durable_workflow import CancellationContext +from durable_workflow._cooperative_cancellation import read_cancellation_history +from durable_workflow.errors import NonDeterministicReplayError, WorkflowCancelled +from durable_workflow.workflow import CompleteWorkflow, FailWorkflow, ScheduleActivity, WorkflowContext, replay +from tests.test_cooperative_cancellation import completed_activity + + +def snapshot() -> dict[str, Any]: + return json.loads((Path(__file__).parent / "fixtures/cooperative-cancellation-context.json").read_text()) + + +def request() -> dict[str, Any]: + return { + "event_type": "CooperativeCancellationRequested", "recorded_at": "2026-10-01T00:00:05Z", + "payload": { + "workflow_command_id": "request-1", "workflow_instance_id": "child-instance", "workflow_run_id": "run-1", + "reason": "maintenance", "cleanup_deadline_at": "2026-10-01T00:00:30.123456Z", "cancellation": snapshot(), + }, + } + + +def delivery() -> dict[str, Any]: + return { + "event_type": "CooperativeCancellationDelivered", + "payload": { + "workflow_command_id": "request-1", "workflow_run_id": "run-1", "sequence": 1, + "call_kind": "timer", "cancellation": snapshot(), + }, + } + + +def observation() -> dict[str, Any]: + return { + "request_id": "request-1", "requested_at": "2026-10-01T00:00:05Z", + "cleanup_deadline_at": "2026-10-01T00:00:30.123456Z", "history_refresh_page_token": "opaque-first-page", + "cancellation": {**snapshot(), "reason": "untrusted observation"}, + } + + +def test_context_and_nested_metadata_are_immutable() -> None: + original = snapshot() + context = CancellationContext.from_dict(original) + original["reason"] = "changed" + original["requester"]["id"] = "changed" + original["lineage"][0]["request_id"] = "changed" + detached = context.to_dict() + detached["requester"]["id"] = "changed" + detached["lineage"].reverse() + assert context.to_dict() == snapshot() + assert context.request_id == "request-1" + assert context.root_request_id == context.parent_request_id == "root-1" + assert context.deadline == datetime(2026, 10, 1, 0, 0, 30, 123456, timezone.utc) + assert context.requested_at == datetime(2026, 10, 1, 0, 0, 0, 123456, timezone.utc) + with pytest.raises(FrozenInstanceError): + context.reason = "changed" # type: ignore[misc] + with pytest.raises(TypeError): + context.requester["id"] = "changed" # type: ignore[index] + with pytest.raises(FrozenInstanceError): + context.lineage[0].request_id = "changed" # type: ignore[misc] + + +def test_timezone_and_object_key_order_do_not_change_the_context() -> None: + value = snapshot() + value["requested_at"] = "2026-09-30T20:00:00.123456-04:00" + value["cleanup_deadline_at"] = "2026-09-30T20:00:30.123456-04:00" + value["requester"] = dict(reversed(list(value["requester"].items()))) + value["lineage"] = [dict(reversed(list(entry.items()))) for entry in value["lineage"]] + assert CancellationContext.from_dict(value) == CancellationContext.from_dict(snapshot()) + + +@pytest.mark.parametrize("field,value", [ + ("schema", "unknown"), ("request_id", ""), ("source", " "), ("parent_request_id", "wrong"), + ("root_request_id", "wrong"), ("root_workflow_run_id", "wrong"), ("reason", []), ("requester", []), + ("lineage", []), ("requested_at", "2026-02-30T00:00:00Z"), + ("cleanup_deadline_at", "2026-10-01T00:00:00.123456Z"), +]) +def test_invalid_context_is_rejected(field: str, value: Any) -> None: + with pytest.raises(ValueError): + CancellationContext.from_dict({**snapshot(), field: value}) + + +@pytest.mark.parametrize("field", ["request_id", "workflow_run_id"]) +def test_lineage_cannot_cycle(field: str) -> None: + value = snapshot() + value["lineage"][1][field] = value["lineage"][0][field] + with pytest.raises(ValueError, match="cycle"): + CancellationContext.from_dict(value) + + +def test_lineage_order_and_requester_metadata_are_validated() -> None: + value = snapshot() + value["lineage"].reverse() + with pytest.raises(ValueError, match="identities"): + CancellationContext.from_dict(value) + value = snapshot() + value["requester"]["authorization"] = "unsupported" + with pytest.raises(ValueError, match="unsupported metadata"): + CancellationContext.from_dict(value) + + +def test_canonical_child_keeps_root_time_after_expired_local_admission() -> None: + event = request() + event["recorded_at"] = "2026-10-01T00:00:31Z" + state = read_cancellation_history([event], run_id="run-1") + assert state.request is not None + assert state.request.requested_at == "2026-10-01T00:00:00.123456Z" + assert state.request.cleanup_deadline_at == "2026-10-01T00:00:30.123456Z" + assert state.request.context == CancellationContext.from_dict(snapshot()) + assert state.delivery is None + + +def test_observation_keeps_refresh_route_and_history_supplies_context() -> None: + observed = read_cancellation_history([], run_id="run-1", observation=observation()) + assert observed.request is not None and observed.request.context is None + state = read_cancellation_history([request(), delivery()], run_id="run-1", observation=observation()) + assert state.request is not None + assert state.request.context == CancellationContext.from_dict(snapshot()) + assert state.request.history_refresh_page_token == "opaque-first-page" + assert state.request.requested_at == "2026-10-01T00:00:00.123456Z" + + +@pytest.mark.parametrize("field,value", [ + ("workflow_command_id", "other"), ("workflow_instance_id", "other"), + ("cleanup_deadline_at", "2026-10-01T00:00:35Z"), ("reason", "changed"), ("cancellation", None), +]) +def test_context_must_match_canonical_local_request(field: str, value: Any) -> None: + event = request() + event["payload"][field] = value + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history([event], run_id="run-1") + + +def test_context_must_name_local_run_and_not_postdate_admission() -> None: + event = request() + event["payload"]["cancellation"]["lineage"][1]["workflow_run_id"] = "other" + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history([event], run_id="run-1") + event = request() + event["recorded_at"] = "2026-09-30T23:59:59Z" + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history([event], run_id="run-1") + + +def test_delivery_cannot_change_accepted_context() -> None: + event = delivery() + event["payload"]["cancellation"]["reason"] = "changed" + with pytest.raises(NonDeterministicReplayError, match="changes the canonical"): + read_cancellation_history([request(), event], run_id="run-1") + + +class ContextCleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + assert ctx.cancellation_context is None + try: + yield ctx.start_timer(10) + except WorkflowCancelled as cancelled: + assert cancelled.context is not None + assert cancelled.context is ctx.cancellation_context + with ctx.cancellation_shield(): + assert cancelled.context is ctx.cancellation_context + ctx.throw_if_cancellation_requested() + yield ctx.schedule_activity("cleanup", []) + return cancelled.context.to_dict() + return None + + +def test_cold_replay_exposes_same_context_only_at_delivery_and_cleanup() -> None: + history = [request(), delivery()] + first = replay(ContextCleanup, history, [], run_id="run-1") + assert len(first.commands) == 1 and isinstance(first.commands[0], ScheduleActivity) + assert first.commands[0].activity_type == "cleanup" + history.append(completed_activity(2, "cleanup", "cleaned")) + for _restart in range(2): + assert replay(ContextCleanup, history, [], run_id="run-1").commands == [CompleteWorkflow(snapshot())] + + +class ExplicitContextCheck: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(10) + except WorkflowCancelled: + try: + ctx.throw_if_cancellation_requested() + except WorkflowCancelled as cancelled: + assert cancelled.request_id == "request-1" + assert cancelled.context == CancellationContext.from_dict(snapshot()) + raise + + +def test_explicit_check_keeps_delivered_context() -> None: + outcome = replay(ExplicitContextCheck, [request(), delivery()], [], run_id="run-1") + assert len(outcome.commands) == 1 and isinstance(outcome.commands[0], FailWorkflow) + assert outcome.commands[0].exception_type == "WorkflowCancelled" diff --git a/tests/test_cancellation_remaining_time.py b/tests/test_cancellation_remaining_time.py new file mode 100644 index 0000000..db7d5c1 --- /dev/null +++ b/tests/test_cancellation_remaining_time.py @@ -0,0 +1,265 @@ +from __future__ import annotations + +from copy import deepcopy +from datetime import datetime, timezone +from typing import Any + +import pytest + +from durable_workflow import CancellationContext, serializer +from durable_workflow.errors import ActivityFailed, DurableOperationCancelled, WorkflowCancelled +from durable_workflow.workflow import ( + CompleteWorkflow, + RecordLocalActivity, + RecordSideEffect, + ScheduleActivity, + WorkflowContext, + replay, +) +from tests.test_cancellation_context import delivery, request, snapshot +from tests.test_cooperative_cancellation import completed_activity +from tests.test_durable_selection import _activity_completed, _activity_scheduled, _winner_marker + + +def history() -> list[dict[str, Any]]: + return [ + {"event_type": "WorkflowStarted", "payload": {"timestamp": "2025-01-01T00:00:00Z"}}, + request(), + {**delivery(), "timestamp": "2026-10-01T00:00:08Z"}, + ] + + +def test_detached_metadata_has_no_remaining_time_clock() -> None: + context = CancellationContext.from_dict(snapshot()) + with pytest.raises(RuntimeError, match="active workflow replay"): + context.remaining() + + +def test_first_execution_and_cold_replay_keep_the_same_budget_through_synchronous_results() -> None: + observations: list[tuple[float, float]] = [] + contexts: list[CancellationContext] = [] + calls: list[str] = [] + + class Cleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled as error: + assert error.context is ctx.cancellation_context + assert error.context is not None + contexts.append(error.context) + before = error.context.remaining() + with ctx.cancellation_shield(): + yield ctx.side_effect(lambda: calls.append("once") or "recorded") + observations.append((before, error.context.remaining())) + yield ctx.schedule_activity("cleanup", []) + assert ctx.now() == datetime(2025, 1, 1, tzinfo=timezone.utc) + return {"remaining": error.context.remaining(), "context": error.context.to_dict()} + + events = history() + first = replay(Cleanup, events, [], run_id="run-1") + assert isinstance(first.commands[0], RecordSideEffect) + assert isinstance(first.commands[1], ScheduleActivity) + events.extend([ + {"event_type": "SideEffectRecorded", "timestamp": "2026-10-01T00:00:28Z", "payload": { + "sequence": 2, "result": serializer.envelope("recorded"), + }}, + {**completed_activity(3, "cleanup", "cleaned"), "timestamp": "2026-10-01T00:00:25Z"}, + {"event_type": "WorkflowTaskCompleted", "timestamp": "2026-10-01T00:00:29Z", "payload": {}}, + ]) + for _restart in range(2): + assert replay(Cleanup, events, [], run_id="run-1").commands == [ + CompleteWorkflow({"remaining": 5.123456, "context": snapshot()}), + ] + assert observations == [(22.123456, 22.123456)] * 3 + assert calls == ["once"] + for context in contexts: + with pytest.raises(RuntimeError, match="active workflow replay"): + context.remaining() + + +def test_completed_parallel_group_uses_consumed_members_without_regressing_on_clock_skew() -> None: + class Cleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled as error: + assert error.context is not None + with ctx.cancellation_shield(): + yield [ctx.schedule_activity("one", []), ctx.schedule_activity("two", [])] + return error.context.remaining() + + events = history() + [ + {**completed_activity(2, "one", "one"), "timestamp": "2026-10-01T00:00:25Z"}, + {**completed_activity(3, "two", "two"), "timestamp": "2026-10-01T00:00:20Z"}, + ] + assert replay(Cleanup, events, [], run_id="run-1").commands == [CompleteWorkflow(5.123456)] + + +def test_parallel_failure_does_not_consume_a_future_sibling_completion() -> None: + class Cleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled as error: + assert error.context is not None + with ctx.cancellation_shield(): + try: + yield [ctx.schedule_activity("one", []), ctx.schedule_activity("two", [])] + except ActivityFailed: + return error.context.remaining() + + events = history() + [ + {"event_type": "ActivityFailed", "timestamp": "2026-10-01T00:00:12Z", "payload": { + "sequence": 2, "activity_type": "one", "message": "failed", "exception_class": "RuntimeError", + }}, + {**completed_activity(3, "two", "two"), "timestamp": "2026-10-01T00:00:29Z"}, + ] + assert replay(Cleanup, events, [], run_id="run-1").commands == [CompleteWorkflow(18.123456)] + + +@pytest.mark.parametrize("timestamp", [None, "invalid", "2026-10-01T00:00:08", "2026-02-30T00:00:08Z"]) +def test_missing_or_invalid_delivery_time_refuses_a_host_clock(timestamp: str | None) -> None: + class Cleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled as error: + assert error.context is not None + with pytest.raises(RuntimeError, match="recorded timestamp"): + error.context.remaining() + return "refused" + + events = history() + events[-1]["timestamp"] = timestamp + assert replay(Cleanup, events, [], run_id="run-1").commands == [CompleteWorkflow("refused")] + + +def test_recorded_expiry_clamps_remaining_to_zero_and_accepts_timezone_offsets() -> None: + class Cleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled as error: + assert error.context is not None + with ctx.cancellation_shield(): + yield ctx.start_timer(10) + return error.context.remaining() + + events = history() + [{"event_type": "TimerFired", "recorded_at": "2026-09-30T20:00:31-04:00", "payload": { + "sequence": 2, "timer_kind": "durable_timer", "duration_ms": 10000, + }}] + assert replay(Cleanup, events, [], run_id="run-1").commands == [CompleteWorkflow(0.0)] + + +def selection_event(event: dict[str, Any], timestamp: str) -> dict[str, Any]: + shifted = deepcopy(event) + shifted["timestamp"] = f"2026-10-01T00:00:{timestamp}Z" + payload = shifted["payload"] + for item in [payload, *payload.get("parallel_group_path", [])]: + for key in ("sequence", "parallel_group_base_sequence", "selection_member_base_sequence", + "selection_group_base_sequence", "member_base_sequence"): + if key in item: + item[key] += 1 + for key in ("parallel_group_id", "selection_group_id"): + if key in item: + item[key] = "select-calls:2:2" + return shifted + + +def selection_history() -> list[dict[str, Any]]: + return history() + [ + selection_event(_activity_scheduled(0, "slow"), "09"), + selection_event(_activity_scheduled(1, "fast"), "09"), + selection_event(_activity_completed(1, "fast", "fast"), "10"), + selection_event(_winner_marker(), "12"), + ] + + +def test_selection_advances_at_the_winner_and_awaited_handle_without_future_loser_lookahead() -> None: + class Cleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled as error: + assert error.context is not None + with ctx.cancellation_shield(): + selected = yield ctx.select({ + "slow": ctx.schedule_activity("slow-activity", []), + "fast": ctx.schedule_activity("fast-activity", []), + }) + at_winner = error.context.remaining() + yield selected.winner.await_result() + at_old_result = error.context.remaining() + yield selected.handles["slow"].await_result() + return [at_winner, at_old_result, error.context.remaining()] + + events = selection_history() + [selection_event(_activity_completed(0, "slow", "slow"), "20")] + assert replay(Cleanup, events, [], run_id="run-1").commands == [ + CompleteWorkflow([18.123456, 18.123456, 10.123456]), + ] + + +def test_cancelled_handle_uses_the_first_durable_receipt_despite_identical_redelivery() -> None: + class Cleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled as error: + assert error.context is not None + with ctx.cancellation_shield(): + selected = yield ctx.select({ + "slow": ctx.schedule_activity("slow-activity", []), + "fast": ctx.schedule_activity("fast-activity", []), + }) + yield selected.handles["slow"].cancel() + try: + yield selected.handles["slow"].await_result() + except DurableOperationCancelled: + return error.context.remaining() + + cancelled = {"event_type": "SelectionOperationCancelled", "timestamp": "2026-10-01T00:00:15Z", + "payload": {"selection_group_id": "select-calls:2:2", "member_key": "slow", "member_index": 0, + "member_base_sequence": 2, "member_size": 1, "operation_kind": "activity", + "operation_identity": "activity-slow"}} + events = selection_history() + [cancelled, {**cancelled, "timestamp": "2026-10-01T00:00:28Z"}] + assert replay(Cleanup, events, [], run_id="run-1").commands == [CompleteWorkflow(15.123456)] + + +def test_inline_local_result_persistence_preserves_the_first_execution_budget() -> None: + observations: list[float] = [] + calls: list[str] = [] + + class Cleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled as error: + assert error.context is not None + with ctx.cancellation_shield(): + yield ctx.local_activity("inline", []) + observations.append(error.context.remaining()) + yield ctx.schedule_activity("cleanup", []) + return error.context.remaining() + + def execute(command: RecordLocalActivity) -> str: + calls.append("once") + command.arguments_envelope = serializer.envelope(command.arguments) + command.result_envelope = serializer.envelope("local") + command.outcome = {"outcome": "completed", "attempts": []} + return "local" + + events = history() + first = replay(Cleanup, events, [], run_id="run-1", local_activity_executor=execute) + assert isinstance(first.commands[0], RecordLocalActivity) + events.extend([ + {"event_type": "ActivityCompleted", "timestamp": "2026-10-01T00:00:28Z", "payload": { + "sequence": 2, "activity_type": "inline", "execution_mode": "local", "result": serializer.envelope("local"), + }}, + {**completed_activity(3, "cleanup", "cleaned"), "timestamp": "2026-10-01T00:00:25Z"}, + ]) + assert replay(Cleanup, events, [], run_id="run-1", local_activity_executor=execute).commands == [ + CompleteWorkflow(5.123456), + ] + assert observations == [22.123456, 22.123456] + assert calls == ["once"] diff --git a/tests/test_cancellation_scope_admission.py b/tests/test_cancellation_scope_admission.py new file mode 100644 index 0000000..6445fe8 --- /dev/null +++ b/tests/test_cancellation_scope_admission.py @@ -0,0 +1,107 @@ +from __future__ import annotations + +from copy import deepcopy +from typing import Any +from unittest.mock import AsyncMock + +import pytest + +from durable_workflow import serializer, workflow +from durable_workflow.client import Client +from durable_workflow.errors import QueryFailed +from durable_workflow.worker import Worker +from durable_workflow.workflow import LocalActivityExecutionAborted + + +@workflow.defn(name="scope-admission-probe") +class ScopeProbe: + calls: list[str] = [] + + def __init__(self) -> None: + self.calls.append("constructed") + + def run(self, ctx: workflow.WorkflowContext): # type: ignore[no-untyped-def] + self.calls.append("run") + yield ctx.start_timer(1) + return "done" + + @workflow.query("state") + def state(self) -> str: + self.calls.append("query") + return "state" + + +@pytest.fixture(autouse=True) +def reset_calls() -> None: + ScopeProbe.calls = [] + + +def scoped_histories() -> list[dict[str, Any]]: + histories: list[dict[str, Any]] = [ + {"event_type": name, "payload": {"sequence": 1, "scope_id": "scope-one"}} + for name in ( + "CancellationScopeOpened", + "CancellationScopeRequested", + "CancellationScopeDelivered", + "CancellationScopeRequestConflicted", + ) + ] + for location in (None, "activity", "timer", "child_workflow"): + value = {"cancellation_scope_id": "scope-one"} + histories.append({"event_type": "TimerScheduled", "payload": value if location is None else {location: value}}) + for malformed in (None, True, 1, "", [], {}): + histories.append({"event_type": "TimerScheduled", "payload": {"cancellation_scope_id": malformed}}) + return histories + + +@pytest.mark.parametrize("event", scoped_histories()) +def test_unqualified_scope_replay_stops_before_constructing_or_running_workflow(event: dict[str, Any]) -> None: + with pytest.raises(LocalActivityExecutionAborted, match="cancellation_scope_execution_not_supported.*Python"): + workflow.replay(ScopeProbe, [event], [], run_id="run-one", payload_codec="avro") + assert ScopeProbe.calls == [] + + +def test_query_replay_does_not_enter_scope_code_or_its_query_handler() -> None: + with pytest.raises(QueryFailed, match="cancellation_scope_execution_not_supported.*Python"): + workflow.query_state(ScopeProbe, [scoped_histories()[0]], [], "state", run_id="run-one", payload_codec="avro") + assert ScopeProbe.calls == [] + + +@pytest.mark.parametrize("explicit_root", [False, True]) +def test_historical_omission_and_explicit_root_preserve_existing_replay(explicit_root: bool) -> None: + payload: dict[str, Any] = {"sequence": 1, "duration_seconds": 1} + if explicit_root: + payload["cancellation_scope_id"] = "root" + history = [{"event_type": name, "payload": deepcopy(payload)} for name in ("TimerScheduled", "TimerFired")] + outcome = workflow.replay(ScopeProbe, history, [], run_id="run-one", payload_codec="avro") + assert [type(command).__name__ for command in outcome.commands] == ["CompleteWorkflow"] + assert ScopeProbe.calls == ["constructed", "run"] + + +def test_scope_named_application_values_are_preserved_as_data() -> None: + @workflow.defn(name="scope-data-probe") + class DataProbe: + def run(self, ctx: workflow.WorkflowContext): # type: ignore[no-untyped-def] + return (yield ctx.side_effect(lambda: pytest.fail("Recorded application data must replay."))) + + value = {"cancellation_scope_id": "application-owned", "activity": {"cancellation_scope_id": None}} + event = {"event_type": "SideEffectRecorded", "payload": { + "sequence": 1, "result": serializer.envelope(value, codec="avro"), + }} + outcome = workflow.replay(DataProbe, [event], [], run_id="run-one", payload_codec="avro") + assert outcome.commands[0].result == value + + +async def test_worker_abandons_unsupported_scope_without_publishing_a_terminal_failure( + caplog: pytest.LogCaptureFixture, +) -> None: + client = AsyncMock(spec=Client) + worker = Worker(client, task_queue="scope", worker_id="scope-worker", workflows=[ScopeProbe]) + task = {"task_id": "scope-task", "workflow_id": "scope-instance", "run_id": "run-one", + "workflow_type": "scope-admission-probe", "workflow_task_attempt": 2, "payload_codec": "avro", + "arguments": serializer.envelope([], codec="avro"), "history_events": [scoped_histories()[0]]} + assert await worker._run_workflow_task_core(task) is None + assert ScopeProbe.calls == [] + client.complete_workflow_task.assert_not_awaited() + client.fail_workflow_task.assert_not_awaited() + assert "cancellation_scope_execution_not_supported" in caplog.text diff --git a/tests/test_child_workflow_policies.py b/tests/test_child_workflow_policies.py new file mode 100644 index 0000000..8fb1250 --- /dev/null +++ b/tests/test_child_workflow_policies.py @@ -0,0 +1,261 @@ +from __future__ import annotations + +from typing import Any + +import pytest + +from durable_workflow import CancellationPolicy, ParentClosePolicy, serializer +from durable_workflow import workflow as workflow_module +from durable_workflow.errors import NonDeterministicReplayError, WorkflowCancelled +from durable_workflow.worker import Worker +from durable_workflow.workflow import ( + CompleteWorkflow, + StartChildWorkflow, + WorkflowContext, + commands_to_server_commands, + replay, +) +from tests.test_cooperative_cancellation import marker, request +from tests.test_cooperative_cancellation_worker import ClaimServer, claimed_task + + +def options() -> dict[str, Any]: + return { + "parent_close_policy": ParentClosePolicy.REQUEST_CANCELLATION, + "cancellation_policy": CancellationPolicy.WAIT_CANCELLATION_COMPLETED, + } + + +def workflow(mode: str, policies: dict[str, Any]) -> type: + class PolicyWorkflow: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + first = ctx.start_child_workflow("child", ["argument"], **policies) + try: + if mode == "parallel": + return (yield [first]) + if mode == "selection": + return (yield ctx.select({ + "first": first, + "second": ctx.start_child_workflow("other-child", [], **policies), + })) + return (yield first) + except WorkflowCancelled as cancelled: + return cancelled.request_id + + return PolicyWorkflow + + +def scheduled(commands: list[Any]) -> list[dict[str, Any]]: + return [ + {"event_type": "ChildWorkflowScheduled", "payload": { + "sequence": index + 1, "child_workflow_type": command["workflow_type"], + "child_workflow_run_id": f"child-run-{index + 1}", **command, + }} + for index, command in enumerate(commands_to_server_commands(commands, "queue")) + ] + + +@pytest.mark.parametrize("parent", list(ParentClosePolicy)) +@pytest.mark.parametrize("operation", list(CancellationPolicy)) +def test_typed_and_string_policies_round_trip_through_both_encoders_and_cold_replay( + parent: ParentClosePolicy, operation: CancellationPolicy, +) -> None: + child = StartChildWorkflow("child", ["argument"], parent_close_policy=parent, cancellation_policy=operation) + direct = child.to_server_command("queue") + assert direct == commands_to_server_commands([child], "queue")[0] + assert direct == StartChildWorkflow( + "child", ["argument"], parent_close_policy=parent.value, cancellation_policy=operation.value, + ).to_server_command("queue") + assert type(direct["parent_close_policy"]) is str + assert type(direct["cancellation_policy"]) is str + assert serializer.decode_envelope(direct["arguments"]) == ["argument"] + history = scheduled([child]) + [{"event_type": "ChildRunCompleted", "payload": { + "sequence": 1, "child_workflow_type": "child", "result": serializer.envelope("recorded-result"), + }}] + result = replay(workflow("sequential", { + "parent_close_policy": parent, "cancellation_policy": operation, + }), history, []) + assert result.commands == [CompleteWorkflow("recorded-result")] + + +@pytest.mark.parametrize("field", ["parent_close_policy", "cancellation_policy"]) +@pytest.mark.parametrize("value", ["unknown", "", True, 1, [], {"type": "abandon"}]) +def test_invalid_options_are_rejected_before_scheduling(field: str, value: Any) -> None: + with pytest.raises(ValueError, match=field): + StartChildWorkflow("child", **{field: value}) + + +def test_policy_enum_types_are_not_interchangeable() -> None: + with pytest.raises(ValueError, match="parent_close_policy"): + StartChildWorkflow("child", parent_close_policy=CancellationPolicy.ABANDON) + with pytest.raises(ValueError, match="cancellation_policy"): + StartChildWorkflow("child", cancellation_policy=ParentClosePolicy.ABANDON) + + +def test_existing_positional_constructor_and_omitted_wire_fields_are_preserved() -> None: + child = StartChildWorkflow("child", ["argument"], "queue", "request_cancel", None, 60, 30) + assert child.execution_timeout_seconds == 60 + assert child.run_timeout_seconds == 30 + assert child.cancellation_policy is None + assert "cancellation_policy" not in child.to_server_command("queue") + omitted = StartChildWorkflow("child").to_server_command("queue") + assert "parent_close_policy" not in omitted + assert "cancellation_policy" not in omitted + result = replay(workflow("sequential", { + "parent_close_policy": ParentClosePolicy.ABANDON, "cancellation_policy": CancellationPolicy.ABANDON, + }), scheduled([StartChildWorkflow("child", ["argument"])]), []) + assert len(result.commands) == 1 + assert isinstance(result.commands[0], StartChildWorkflow) + assert result.commands[0].cancellation_policy == CancellationPolicy.ABANDON + + +@pytest.mark.parametrize("mode", ["sequential", "parallel", "selection"]) +@pytest.mark.parametrize("field", ["parent_close_policy", "cancellation_policy"]) +@pytest.mark.parametrize("delivered", [False, True]) +def test_changed_policies_fail_cold_replay_and_cancellation_delivery(mode: str, field: str, delivered: bool) -> None: + policies = options() + history = scheduled(replay(workflow(mode, policies), [], []).commands) + if delivered: + history += [request(), marker(1, "child" if mode == "sequential" else "parallel", sequence_span=len(history))] + assert replay(workflow(mode, policies), history, [], run_id="run-1").commands == [CompleteWorkflow("request-1")] + else: + cold = replay(workflow(mode, policies), history, []) + if mode == "selection": + assert cold.commands == [] + else: + assert scheduled(cold.commands) == history + policies[field] = "abandon" + with pytest.raises(NonDeterministicReplayError, match="child_workflow_policy_changed") as captured: + replay(workflow(mode, policies), history, [], run_id="run-1") + assert captured.value.workflow_sequence == 1 + + +@pytest.mark.parametrize("mode", ["sequential", "parallel", "selection"]) +@pytest.mark.parametrize("field", ["parent_close_policy", "cancellation_policy"]) +def test_historical_defaults_cannot_change_to_cooperative_policies(mode: str, field: str) -> None: + history = scheduled(replay(workflow(mode, {}), [], []).commands) + with pytest.raises(NonDeterministicReplayError, match="child_workflow_policy_changed"): + replay(workflow(mode, {field: options()[field]}), history, []) + + +@pytest.mark.parametrize("event", [ + "ChildRunStarted", "ChildRunCompleted", "ChildRunFailed", "ChildRunCancelled", "ChildRunTerminated", +]) +@pytest.mark.parametrize("policy,reason", [ + ("try_cancel", "child_workflow_policy_history_conflict"), + ({"type": "try_cancel"}, "invalid_child_workflow_policy_history"), +]) +def test_invalid_and_conflicting_child_history_is_rejected(event: str, policy: Any, reason: str) -> None: + history = scheduled(replay(workflow("sequential", options()), [], []).commands) + history.append({"event_type": event, "payload": {"sequence": 1, "cancellation_policy": policy}}) + with pytest.raises(NonDeterministicReplayError, match=reason): + replay(workflow("sequential", options()), history, []) + + +def test_later_start_and_terminal_events_preserve_the_scheduled_snapshot() -> None: + history = scheduled(replay(workflow("sequential", options()), [], []).commands) + history += [ + {"event_type": "ChildRunStarted", "payload": {"sequence": 1, "child_workflow_type": "child"}}, + {"event_type": "ChildRunCompleted", "payload": {"sequence": 1, "result": serializer.envelope("done")}}, + ] + assert replay(workflow("sequential", options()), history, []).commands == [CompleteWorkflow("done")] + + +def test_selection_handle_keeps_the_loser_policy_after_the_winner_and_during_cancellation() -> None: + def authoring(second_policy: CancellationPolicy) -> type: + class AwaitSecond: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + selected = yield ctx.select({ + "first": ctx.start_child_workflow("child", ["argument"], **options()), + "second": ctx.start_child_workflow("other-child", [], **{ + **options(), "cancellation_policy": second_policy, + }), + }) + try: + return (yield selected.handles["second"].await_result()) + except WorkflowCancelled as cancelled: + return cancelled.request_id + + return AwaitSecond + + original = authoring(CancellationPolicy.WAIT_CANCELLATION_COMPLETED) + history = scheduled(replay(original, [], [], run_id="run-1").commands) + history += [ + {"id": "child-completed", "event_type": "ChildRunCompleted", "payload": { + **history[0]["payload"], "result": serializer.envelope("first-result"), + }}, + {"event_type": "SelectionResolved", "payload": { + "selection_group_id": "select-calls:1:2", "selection_group_base_sequence": 1, "selection_group_size": 2, + "member_key": "first", "member_index": 0, "member_base_sequence": 1, "member_size": 1, + "operation_kind": "child", "operation_identity": "child-run-1", "outcome": "completed", + "resolution_event_id": "child-completed", "resolution_event_type": "ChildRunCompleted", + }}, + ] + assert replay(original, history, [], run_id="run-1").commands == [] + for delivered in (False, True): + current = list(history) + if delivered: + current += [request(), marker(3, "selection_handle", operation_sequence=2, operation_sequence_span=1)] + assert replay(original, current, [], run_id="run-1").commands == [CompleteWorkflow("request-1")] + with pytest.raises(NonDeterministicReplayError, match="child_workflow_policy_changed") as captured: + replay(authoring(CancellationPolicy.TRY_CANCEL), current, [], run_id="run-1") + assert captured.value.workflow_sequence == 2 + + +@pytest.mark.parametrize("protocol", ["1.19", "1.20"]) +@pytest.mark.parametrize("policies", [ + {"parent_close_policy": ParentClosePolicy.REQUEST_CANCELLATION}, + {"cancellation_policy": CancellationPolicy.TRY_CANCEL}, + {"cancellation_policy": CancellationPolicy.WAIT_CANCELLATION_COMPLETED}, +]) +async def test_worker_without_opt_in_refuses_cooperative_child_commands( + monkeypatch: pytest.MonkeyPatch, protocol: str, policies: dict[str, Any], +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", protocol) + cls = workflow_module.defn(name="child-policy-worker")(workflow("sequential", policies)) + server = ClaimServer(history=[]) + worker = Worker(server.client, task_queue="queue", worker_id="cooperative-worker", workflows=[cls]) + await worker._register() + result = await worker._run_workflow_task({ + **claimed_task(observed=False), "workflow_type": "child-policy-worker", "arguments": serializer.envelope([]), + }) + assert result is None + server.client.complete_workflow_task.assert_not_awaited() + failure = server.client.fail_workflow_task.await_args.kwargs + assert failure["failure_type"] == "RuntimeCapabilityUnsupported" + assert "child_cancellation_policy_not_supported" in failure["message"] + assert "cooperative-worker" in failure["message"] + assert "1.20" in failure["message"] + + +@pytest.mark.parametrize("parent", [ + ParentClosePolicy.ABANDON, ParentClosePolicy.REQUEST_CANCEL, ParentClosePolicy.TERMINATE, +]) +async def test_legacy_child_policies_remain_available_without_cooperation( + monkeypatch: pytest.MonkeyPatch, parent: ParentClosePolicy, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.19") + cls = workflow_module.defn(name="child-policy-worker")(workflow("sequential", { + "parent_close_policy": parent, "cancellation_policy": CancellationPolicy.ABANDON, + })) + server = ClaimServer(history=[]) + worker = Worker(server.client, task_queue="queue", worker_id="cooperative-worker", workflows=[cls]) + await worker._register() + result = await worker._run_workflow_task({ + **claimed_task(observed=False), "workflow_type": "child-policy-worker", "arguments": serializer.envelope([]), + }) + assert result is not None and result[0]["parent_close_policy"] == parent.value + server.client.fail_workflow_task.assert_not_awaited() + + +async def test_capable_worker_transmits_both_typed_child_policies(monkeypatch: pytest.MonkeyPatch) -> None: + cls = workflow_module.defn(name="child-policy-worker")(workflow("sequential", options())) + server = ClaimServer(history=[]) + worker = await server.worker(monkeypatch, workflows=[cls]) + result = await worker._run_workflow_task({ + **claimed_task(observed=False), "workflow_type": "child-policy-worker", "arguments": serializer.envelope([]), + }) + assert result is not None + assert result[0]["parent_close_policy"] == "request_cancellation" + assert result[0]["cancellation_policy"] == "wait_cancellation_completed" + server.client.fail_workflow_task.assert_not_awaited() diff --git a/tests/test_ci_checkout.py b/tests/test_ci_checkout.py index 5ed297a..66ec059 100644 --- a/tests/test_ci_checkout.py +++ b/tests/test_ci_checkout.py @@ -22,6 +22,7 @@ [ ("cli", "https://github.com/durable-workflow/cli.git"), ("server", "https://github.com/durable-workflow/server.git"), + ("workflow", "https://github.com/durable-workflow/workflow.git"), ], ) def test_public_checkout_uses_github_authority_on_every_runner( @@ -72,3 +73,48 @@ def test_ci_workflow_uses_portable_public_checkouts() -> None: line.strip() for line in workflow.splitlines() if line.strip().startswith("repository: durable-workflow/") ] assert public_repository_inputs == ["repository: durable-workflow/.github"] + + +@pytest.mark.parametrize("commit", ["main", "short", "A" * 40, "--upload-pack=command", "0" * 40 + ";command"]) +def test_candidate_checkout_rejects_non_sha_before_running_git(tmp_path: Path, commit: str) -> None: + capture = tmp_path / "called" + fake_git = tmp_path / "git" + fake_git.write_text('#!/bin/sh\ntouch "$GIT_CAPTURE"\n') + fake_git.chmod(0o755) + environment = {**os.environ, "PATH": str(tmp_path), "GIT_CAPTURE": str(capture)} + result = subprocess.run( + [sys.executable, str(CHECKOUT_SCRIPT), "server", str(tmp_path / "server"), "--commit", commit], + env=environment, capture_output=True, text=True, + ) + assert result.returncode != 0 + assert not capture.exists() + + +@pytest.mark.parametrize("matches", [True, False]) +@pytest.mark.parametrize("repository", ["server", "workflow"]) +def test_candidate_checkout_verifies_the_requested_public_commit( + tmp_path: Path, matches: bool, repository: str, +) -> None: + commit = "a" * 40 + resolved = commit if matches else "b" * 40 + capture = tmp_path / "git-calls" + fake_git = tmp_path / "git" + fake_git.write_text( + '#!/bin/sh\nprintf "%s\\n" "$*" >> "$GIT_CAPTURE"\n' + 'case "$*" in *rev-parse*) printf "%s\\n" "$RESOLVED_COMMIT" ;; esac\n' + ) + fake_git.chmod(0o755) + environment = { + **os.environ, "PATH": str(tmp_path), "GIT_CAPTURE": str(capture), "RESOLVED_COMMIT": resolved, + } + result = subprocess.run( + [sys.executable, str(CHECKOUT_SCRIPT), repository, str(tmp_path / repository), "--commit", commit], + env=environment, capture_output=True, text=True, + ) + assert (result.returncode == 0) is matches + calls = capture.read_text().splitlines() + assert f"https://github.com/durable-workflow/{repository}.git" in calls[0] + assert any(f"fetch --depth=1 origin {commit}" in call for call in calls) + assert all("credential.helper=" in call for call in calls) + if matches: + assert f"Integration {repository} source commit: {commit}" in result.stdout diff --git a/tests/test_cooperative_cancellation.py b/tests/test_cooperative_cancellation.py new file mode 100644 index 0000000..4516421 --- /dev/null +++ b/tests/test_cooperative_cancellation.py @@ -0,0 +1,425 @@ +from __future__ import annotations + +from copy import deepcopy +from typing import Any + +import pytest + +from durable_workflow import serializer +from durable_workflow._cooperative_cancellation import CancellationDelivery, read_cancellation_history +from durable_workflow.errors import NonDeterministicReplayError, WorkflowCancelled +from durable_workflow.workflow import CompleteWorkflow, ScheduleActivity, StartTimer, WorkflowContext, replay +from tests.test_durable_selection import _activity_completed, _activity_scheduled, _winner_marker + + +def observation() -> dict[str, Any]: + return { + "request_id": "request-1", + "requested_at": "2026-09-30T12:00:00Z", + "cleanup_deadline_at": "2026-09-30T12:10:00Z", + "history_refresh_page_token": "opaque-first-page", + } + + +def request() -> dict[str, Any]: + return { + "event_type": "CooperativeCancellationRequested", + "workflow_command_id": "request-1", + "recorded_at": "2026-09-30T12:00:00.000120Z", + "payload": { + "workflow_command_id": "request-1", + "workflow_run_id": "run-1", + "cleanup_deadline_at": "2026-09-30T12:10:00Z", + }, + } + + +def marker(sequence: int = 2, call_kind: str = "activity", **fields: Any) -> dict[str, Any]: + return { + "event_type": "CooperativeCancellationDelivered", + "workflow_command_id": "request-1", + "payload": { + "workflow_command_id": "request-1", + "workflow_run_id": "run-1", + "sequence": sequence, + "call_kind": call_kind, + **fields, + }, + } + + +def test_observation_retains_original_identity_without_becoming_delivery() -> None: + state = read_cancellation_history([request()], run_id="run-1", observation=observation()) + assert state.request is not None + assert state.request.requested_at == observation()["requested_at"] + assert state.request.history_refresh_page_token == "opaque-first-page" + assert state.delivery is None + assert state.eligible(1) + + +def test_prior_result_remains_resolved_but_later_result_is_eligible() -> None: + result = {"event_type": "ActivityCompleted", "payload": {"sequence": 1}} + later = {"event_type": "TimerFired", "payload": {"sequence": 2}} + state = read_cancellation_history([result, request(), later]) + assert not state.eligible(1) + assert state.eligible(2) + assert state.eligible(1, 2) + + +def test_prior_parallel_failure_wins_over_a_later_request() -> None: + state = read_cancellation_history( + [ + {"event_type": "ActivityFailed", "payload": {"sequence": 1}}, + request(), + ] + ) + assert not state.eligible(1, 2) + + +def test_cold_history_preserves_one_canonical_delivery_and_its_range() -> None: + state = read_cancellation_history([request(), marker(2, "parallel", sequence_span=3)], run_id="run-1") + assert state.delivery == CancellationDelivery("request-1", 2, "parallel", 3) + assert [state.delivery.interrupts(sequence) for sequence in range(1, 6)] == [False, True, True, True, False] + assert not state.eligible(5) + + +def test_selection_handle_targets_its_earlier_operation_range() -> None: + state = read_cancellation_history( + [ + request(), + marker(5, "selection_handle", operation_sequence=2, operation_sequence_span=2), + ] + ) + assert state.delivery is not None + assert state.delivery.interrupts(2) + assert state.delivery.interrupts(3) + assert not state.delivery.interrupts(5) + + +@pytest.mark.parametrize( + "fields", + [ + {"sequence": True}, + {"sequence": 0}, + {"sequence": -1}, + {"sequence": 2**63 - 1}, + {"call_kind": "side_effect"}, + {"sequence_span": 2}, + {"sequence_span": 0}, + {"call_kind": "parallel", "sequence_span": 1001}, + {"operation_sequence": 1}, + {"operation_sequence_span": 2}, + {"call_kind": "selection_handle"}, + {"call_kind": "selection_handle", "operation_sequence": 2}, + {"call_kind": "selection_handle", "operation_sequence": 1, "operation_sequence_span": 2}, + ], +) +def test_invalid_marker_ranges_fail_closed(fields: dict[str, Any]) -> None: + event = marker() + event["payload"].update(fields) + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history([request(), event]) + + +@pytest.mark.parametrize( + "history", + [ + [marker()], + [marker(), request()], + [request(), request()], + [request(), marker(), marker()], + ], +) +def test_missing_reordered_or_duplicate_canonical_markers_fail_closed(history: list[dict[str, Any]]) -> None: + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history(history) + + +@pytest.mark.parametrize("field", ["request_id", "cleanup_deadline_at"]) +def test_observation_cannot_replace_original_request(field: str) -> None: + observed = observation() + observed[field] = "different" if field == "request_id" else "2026-09-30T12:20:00Z" + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history([request()], observation=observed) + + +@pytest.mark.parametrize( + "field,value", + [ + ("request_id", ""), + ("requested_at", "2026-09-30T12:00:00"), + ("cleanup_deadline_at", "2026-09-30T11:00:00Z"), + ("history_refresh_page_token", ""), + ], +) +def test_invalid_observation_fails_closed(field: str, value: Any) -> None: + observed = observation() + observed[field] = value + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history([], observation=observed) + + +def test_delivery_run_and_request_identity_are_checked() -> None: + for field, value in [("workflow_run_id", "other-run"), ("workflow_command_id", "other-request")]: + event = deepcopy(marker()) + event["payload"][field] = value + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history([request(), event], run_id="run-1") + + +def completed_activity(sequence: int, name: str, result: Any) -> dict[str, Any]: + return { + "event_type": "ActivityCompleted", + "payload": { + "sequence": sequence, + "activity_type": name, + "result": serializer.envelope(result), + }, + } + + +class PriorResultCleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + assert not ctx.is_cancellation_requested + ctx.throw_if_cancellation_requested() + previous = yield ctx.schedule_activity("previous", []) + try: + yield ctx.start_timer(30) + except WorkflowCancelled as exc: + assert ctx.is_cancellation_requested + with ctx.cancellation_shield(): + ctx.throw_if_cancellation_requested() + cleanup = yield ctx.schedule_activity("cleanup", [previous]) + return {"previous": previous, "cleanup": cleanup, "request_id": exc.request_id} + return "not cancelled" + + +def test_observation_and_request_do_not_throw_before_the_authored_call() -> None: + prior = completed_activity(1, "previous", "already committed") + for history in [[prior], [prior, request()]]: + outcome = replay(PriorResultCleanup, history, [], run_id="run-1", cancellation_request=observation()) + assert outcome.commands == [] + assert outcome.cancellation_delivery == CancellationDelivery("request-1", 2, "timer") + + +def test_cold_replay_reaches_same_delivery_and_durable_cleanup_boundary() -> None: + history = [ + completed_activity(1, "previous", "already committed"), + {"event_type": "TimerScheduled", "payload": {"sequence": 2, "timer_kind": "durable_timer"}}, + request(), + {"event_type": "TimerCancelled", "payload": {"sequence": 2, "timer_kind": "durable_timer"}}, + marker(2, "timer"), + ] + for _restart in range(2): + outcome = replay(PriorResultCleanup, history, [], run_id="run-1") + assert outcome.cancellation_delivery is None + assert len(outcome.commands) == 1 + assert isinstance(outcome.commands[0], ScheduleActivity) + assert outcome.commands[0].activity_type == "cleanup" + assert outcome.commands[0].arguments == ["already committed"] + + history.extend( + [ + {"event_type": "ActivityScheduled", "payload": {"sequence": 3, "activity_type": "cleanup"}}, + completed_activity(3, "cleanup", "cleanup committed"), + ] + ) + for _restart in range(2): + outcome = replay(PriorResultCleanup, history, [], run_id="run-1") + assert len(outcome.commands) == 1 + assert isinstance(outcome.commands[0], CompleteWorkflow) + assert outcome.commands[0].result == { + "previous": "already committed", + "cleanup": "cleanup committed", + "request_id": "request-1", + } + + +class ScalarCleanup: + def run(self, ctx: WorkflowContext, kind: str): # type: ignore[no-untyped-def] + call = { + "activity": lambda: ctx.schedule_activity("forward", []), + "local_activity": lambda: ctx.local_activity("forward", []), + "timer": lambda: ctx.start_timer(30), + "child": lambda: ctx.start_child_workflow("forward-child", []), + "condition": lambda: ctx.wait_condition(lambda: False, key="forward-wait"), + }[kind]() + try: + yield call + except WorkflowCancelled as exc: + with ctx.cancellation_shield(): + yield ctx.start_timer(1) + return exc.request_id + return "not cancelled" + + +@pytest.mark.parametrize("kind", ["activity", "local_activity", "timer", "child", "condition"]) +def test_scalar_marker_interrupts_at_call_and_cleanup_does_not_redeliver(kind: str) -> None: + calls: list[Any] = [] + history = [request(), marker(1, kind)] + outcome = replay(ScalarCleanup, history, [kind], run_id="run-1", local_activity_executor=calls.append) + assert len(outcome.commands) == 1 + assert isinstance(outcome.commands[0], StartTimer) + assert outcome.commands[0].delay_seconds == 1 + assert calls == [] + outcome = replay( + ScalarCleanup, + history + + [ + {"event_type": "TimerFired", "payload": {"sequence": 2, "timer_kind": "durable_timer"}}, + ], + [kind], + run_id="run-1", + ) + assert outcome.commands == [CompleteWorkflow("request-1")] + + +def test_pending_local_call_is_delivered_before_executing_the_callable() -> None: + calls: list[Any] = [] + outcome = replay( + ScalarCleanup, [request()], ["local_activity"], run_id="run-1", local_activity_executor=calls.append + ) + assert calls == [] + assert outcome.commands == [] + assert outcome.cancellation_delivery == CancellationDelivery("request-1", 1, "local_activity") + + +class ParallelCleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield [ctx.schedule_activity("first", []), [ctx.start_timer(30), ctx.start_child_workflow("child", [])]] + except WorkflowCancelled: + with ctx.cancellation_shield(): + yield ctx.schedule_activity("cleanup", []) + return "cleaned" + + +def test_parallel_delivery_preserves_flat_span_and_cleanup_cursor_on_restart() -> None: + observed = replay(ParallelCleanup, [request()], [], run_id="run-1") + assert observed.cancellation_delivery == CancellationDelivery("request-1", 1, "parallel", 3) + history = [request(), marker(1, "parallel", sequence_span=3), completed_activity(4, "cleanup", "done")] + assert replay(ParallelCleanup, history, [], run_id="run-1").commands == [CompleteWorkflow("cleaned")] + + +@pytest.mark.parametrize( + "history,kind", + [ + ([request(), marker(1, "timer")], "activity"), + ([{"event_type": "TimerFired", "payload": {"sequence": 1}}, request(), marker(2, "activity")], "timer"), + ([request(), marker(1, "parallel", sequence_span=2)], "timer"), + ], +) +def test_changed_or_unreached_delivery_boundary_is_non_deterministic(history: list[dict[str, Any]], kind: str) -> None: + with pytest.raises(NonDeterministicReplayError): + replay(ScalarCleanup, history, [kind], run_id="run-1") + + +class ShieldBeforeRequest: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + with ctx.cancellation_shield(): + yield ctx.start_timer(1) + try: + yield ctx.start_timer(2) + except WorkflowCancelled as exc: + return exc.request_id + + +def test_shield_defers_first_delivery_without_resetting_request_identity() -> None: + first = replay(ShieldBeforeRequest, [request()], [], run_id="run-1") + assert first.commands == [StartTimer(1)] + assert first.cancellation_delivery is None + history = [request(), {"event_type": "TimerFired", "payload": {"sequence": 1}}] + second = replay(ShieldBeforeRequest, history, [], run_id="run-1") + assert second.cancellation_delivery == CancellationDelivery("request-1", 2, "timer") + cold = replay(ShieldBeforeRequest, history + [marker(2, "timer")], [], run_id="run-1") + assert cold.commands == [CompleteWorkflow("request-1")] + + +class SelectionCleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + selected = yield ctx.select({ + "slow": ctx.schedule_activity("slow-activity", []), + "fast": ctx.schedule_activity("fast-activity", []), + }) + yield selected.handles["slow"].await_result() + except WorkflowCancelled as exc: + with ctx.cancellation_shield(): + yield ctx.start_timer(1) + return exc.request_id + + +def selection_history() -> list[dict[str, Any]]: + return [ + _activity_scheduled(0, "slow"), _activity_scheduled(1, "fast"), + _activity_completed(1, "fast", "winner"), _winner_marker(), + ] + + +def test_resolved_selection_replays_before_pending_loser_delivery() -> None: + outcome = replay(SelectionCleanup, selection_history() + [request()], [], run_id="run-1") + assert outcome.commands == [] + assert outcome.cancellation_delivery == CancellationDelivery("request-1", 3, "selection_handle", 1, 1) + + +def test_selection_handle_cold_delivery_preserves_opening_identity_and_cleanup_cursor() -> None: + history = selection_history() + [request(), marker(3, "selection_handle", operation_sequence=1)] + for _restart in range(2): + outcome = replay(SelectionCleanup, history, [], run_id="run-1") + assert outcome.commands == [StartTimer(1)] + history.append({"event_type": "TimerFired", "payload": {"sequence": 4}}) + assert replay(SelectionCleanup, history, [], run_id="run-1").commands == [CompleteWorkflow("request-1")] + + +def test_first_selection_call_delivers_as_one_parallel_boundary() -> None: + outcome = replay(SelectionCleanup, [request()], [], run_id="run-1") + assert outcome.cancellation_delivery == CancellationDelivery("request-1", 1, "parallel", 2) + cold = replay(SelectionCleanup, [request(), marker(1, "parallel", sequence_span=2)], [], run_id="run-1") + assert cold.commands == [StartTimer(1)] + + +def test_marker_cannot_replace_a_result_committed_before_request() -> None: + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history([completed_activity(1, "previous", "result"), request(), marker(1)]) + + +class ReopenedConditionCleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.wait_condition(lambda: False, key="forward-wait") + except WorkflowCancelled as exc: + with ctx.cancellation_shield(): + yield ctx.schedule_activity("cleanup", []) + return exc.request_id + + +def reopened_condition_history() -> list[dict[str, Any]]: + return [ + {"event_type": "ConditionWaitOpened", "payload": { + "sequence": 1, "condition_key": "forward-wait", "condition_wait_id": "wait-1", + }}, + {"event_type": "ConditionWaitSatisfied", "payload": { + "sequence": 1, "condition_wait_id": "wait-1", + }}, + {"event_type": "ConditionWaitOpened", "payload": { + "sequence": 2, "condition_key": "forward-wait", "condition_wait_id": "wait-2", + }}, + request(), + ] + + +def test_request_during_a_false_condition_reopen_uses_the_actual_wait_occurrence() -> None: + outcome = replay(ReopenedConditionCleanup, reopened_condition_history(), [], run_id="run-1") + assert outcome.commands == [] + assert outcome.cancellation_delivery == CancellationDelivery("request-1", 2, "condition") + + +def test_cold_delivery_during_a_false_reopen_consumes_the_interrupted_wait_once() -> None: + history = reopened_condition_history() + [marker(2, "condition")] + for _restart in range(2): + outcome = replay(ReopenedConditionCleanup, history, [], run_id="run-1") + assert len(outcome.commands) == 1 + assert isinstance(outcome.commands[0], ScheduleActivity) + assert outcome.commands[0].activity_type == "cleanup" + history.append(completed_activity(3, "cleanup", "done")) + assert replay(ReopenedConditionCleanup, history, [], run_id="run-1").commands == [CompleteWorkflow("request-1")] diff --git a/tests/test_cooperative_cancellation_client.py b/tests/test_cooperative_cancellation_client.py new file mode 100644 index 0000000..11a3f81 --- /dev/null +++ b/tests/test_cooperative_cancellation_client.py @@ -0,0 +1,503 @@ +from __future__ import annotations + +import asyncio +import time +from collections.abc import AsyncIterator +from typing import Any +from unittest.mock import AsyncMock, patch + +import httpx +import pytest +import pytest_asyncio + +from durable_workflow.client import Client, WorkflowHandle +from durable_workflow.errors import RuntimeCapabilityUnsupported, RuntimeDiscoveryUnavailable, ServerError +from durable_workflow.retry_policy import TransportRetryPolicy + + +def request_response(*, duplicate: bool = False, run_id: str = "run-1") -> dict[str, Any]: + return { + "accepted": True, "duplicate": duplicate, "workflow_id": "wf/1", "run_id": run_id, + "cancellation_request": { + "request_id": "original-request", "requested_at": "2026-09-30T12:00:00Z", + "cleanup_deadline_at": "2026-09-30T12:10:00Z", "history_refresh_page_token": "opaque-first-page", + "delivery_sequence": None, "delivered_at": None, + }, + } + + +def delivery_response(**overrides: Any) -> dict[str, Any]: + return { + "delivered": True, "task_id": "task/1", "workflow_run_id": "run-1", + "request_id": "original-request", "sequence": 3, "call_kind": "activity", "sequence_span": 1, + "operation_sequence": None, "operation_sequence_span": 1, "reason": None, **overrides, + } + + +def response(body: Any, status: int = 200) -> httpx.Response: + return httpx.Response(status, json=body, request=httpx.Request("POST", "http://test")) + + +def pending_delivery_response(**overrides: Any) -> dict[str, Any]: + return { + "delivered": False, "task_id": "task/1", "workflow_run_id": "run-1", + "reason": "cancellation_waiting_for_child", "claim_released": True, + "request_id": None, "sequence": None, "call_kind": None, "sequence_span": None, + "operation_sequence": None, "operation_sequence_span": None, **overrides, + } + + +def activity_stop_response(**overrides: Any) -> dict[str, Any]: + return { + "task_id": "task/1", "activity_attempt_id": "attempt-1", "lease_owner": "worker-1", + "request_id": "original-request", "acknowledged": True, "duplicate": False, + "reason": None, "heartbeat_recorded": False, "history_event_id": "original-event", **overrides, + } + + +@pytest_asyncio.fixture +async def client() -> AsyncIterator[Client]: + async with Client( + "http://localhost:8080", control_token="control-token", worker_token="worker-token", namespace="ns1", + retry_policy=TransportRetryPolicy(initial_backoff_seconds=0, jitter=False), + ) as value: + value._cluster_info = {"worker_protocol": { + "version": "1.20", "server_capabilities": {"cooperative_cancellation": True}, + }} + yield value + + +@pytest.mark.asyncio +@pytest.mark.parametrize("run_id", [None, "run/1"]) +async def test_request_uses_control_credential_and_exact_run_route(client: Client, run_id: str | None) -> None: + expected = request_response(run_id=run_id or "run-1") + with patch.object(client._http, "request", new_callable=AsyncMock, return_value=response(expected, 202)) as send: + result = await client.request_workflow_cancellation( + "wf/1", run_id=run_id, reason="stop", cleanup_timeout_seconds=30, + ) + suffix = "/runs/run%2F1" if run_id is not None else "" + assert send.call_args.args[:2] == ("POST", f"/api/workflows/wf%2F1{suffix}/request-cancellation") + assert send.call_args.kwargs["headers"]["Authorization"] == "Bearer control-token" + assert send.call_args.kwargs["headers"]["X-Namespace"] == "ns1" + assert send.call_args.kwargs["headers"]["X-Durable-Workflow-Control-Plane-Version"] == "2" + assert send.call_args.kwargs["json"] == {"reason": "stop", "cleanup_timeout_seconds": 30} + assert result == expected + + +@pytest.mark.asyncio +async def test_duplicate_returns_original_server_identity_and_deadline(client: Client) -> None: + first = request_response() + duplicate = request_response(duplicate=True) + with patch.object(client._http, "request", new_callable=AsyncMock, + side_effect=[response(first, 202), response(duplicate)]) as send: + accepted = await client.request_workflow_cancellation("wf/1") + repeated = await client.request_workflow_cancellation("wf/1", cleanup_timeout_seconds=3600) + assert accepted["cancellation_request"] == repeated["cancellation_request"] + assert repeated["duplicate"] is True + assert send.await_args_list[0].kwargs["json"] == {} + assert send.await_args_list[1].kwargs["json"] == {"cleanup_timeout_seconds": 3600} + + +@pytest.mark.asyncio +async def test_run_bound_handle_preserves_selection(client: Client) -> None: + with patch.object(client, "request_workflow_cancellation", new_callable=AsyncMock, + return_value=request_response()) as send: + result = await WorkflowHandle(client, "wf/1", "run-1").request_cancellation( + reason="stop", cleanup_timeout_seconds=20, + ) + send.assert_awaited_once_with("wf/1", run_id="run-1", reason="stop", cleanup_timeout_seconds=20) + assert result["cancellation_request"]["request_id"] == "original-request" + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("info", "error"), [ + ({"worker_protocol": {"server_capabilities": {"cooperative_cancellation": False}}}, RuntimeCapabilityUnsupported), + ({}, RuntimeDiscoveryUnavailable), + ({"worker_protocol": {"version": "1.20", "server_capabilities": {"cooperative_cancellation": 1}}}, + RuntimeDiscoveryUnavailable), +]) +async def test_unsupported_or_missing_discovery_sends_no_request(client: Client, info: dict, error: type) -> None: + client._cluster_info = info + with patch.object(client._http, "request", new_callable=AsyncMock) as send, pytest.raises(error): + await client.request_workflow_cancellation("wf/1") + send.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("version", [None, True, "1.19", "2.20", "one", "1.20", "1." + "2" * 5000]) +async def test_capability_without_compatible_protocol_is_not_admitted(client: Client, version: Any) -> None: + client._cluster_info["worker_protocol"]["version"] = version + with ( + patch.object(client._http, "request", new_callable=AsyncMock) as send, + pytest.raises(RuntimeDiscoveryUnavailable), + ): + await client.request_workflow_cancellation("wf/1") + send.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_first_request_discovers_before_mutation(client: Client) -> None: + info = client._cluster_info + client._cluster_info = None + with patch.object(client._http, "request", new_callable=AsyncMock, + side_effect=[response(info), response(request_response(), 202)]) as send: + await client.request_workflow_cancellation("wf/1") + assert send.await_args_list[0].args[:2] == ("GET", "/api/cluster/info") + assert send.await_args_list[0].kwargs["headers"]["Authorization"] == "Bearer worker-token" + assert send.await_args_list[1].kwargs["headers"]["Authorization"] == "Bearer control-token" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("timeout", [True, 0, -1, 3601, 1.5]) +async def test_invalid_cleanup_timeout_is_rejected_before_io(client: Client, timeout: Any) -> None: + with patch.object(client._http, "request", new_callable=AsyncMock) as send, pytest.raises(ValueError): + await client.request_workflow_cancellation("wf/1", cleanup_timeout_seconds=timeout) + send.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", [ + {"workflow_id": "other"}, {"run_id": "other"}, {"accepted": 1}, {"duplicate": None}, + {"cancellation_request": None}, + {"cancellation_request": {"request_id": "original-request"}}, + {"cancellation_request": { + **request_response()["cancellation_request"], "cleanup_deadline_at": "2026-09-30T11:00:00Z", + }}, +]) +async def test_malformed_or_mismatched_request_ack_is_rejected(client: Client, change: dict) -> None: + with ( + patch.object(client._http, "request", new_callable=AsyncMock, + return_value=response({**request_response(), **change}, 202)), + pytest.raises(ServerError) as error, + ): + await client.request_workflow_cancellation("wf/1", run_id="run-1") + assert error.value.reason() == "invalid_cooperative_cancellation_response" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("selection", [False, True]) +async def test_delivery_sends_owner_attempt_and_authored_boundary( + client: Client, monkeypatch: pytest.MonkeyPatch, selection: bool, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + options: dict[str, Any] = {"call_kind": "activity"} + if selection: + options.update(call_kind="selection_handle", operation_sequence=1, operation_sequence_span=2) + with patch.object(client._http, "request", new_callable=AsyncMock, + return_value=response(delivery_response(**options))) as send: + result = await client.deliver_workflow_cancellation( + task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, + request_id="original-request", sequence=3, **options, + ) + assert result["delivered"] is True + assert send.call_args.args[:2] == ("POST", "/api/worker/workflow-tasks/task%2F1/deliver-cancellation") + assert send.call_args.kwargs["headers"]["Authorization"] == "Bearer worker-token" + assert send.call_args.kwargs["headers"]["X-Durable-Workflow-Protocol-Version"] == "1.20" + assert send.call_args.kwargs["json"] == { + "lease_owner": "worker-1", "workflow_task_attempt": 2, "request_id": "original-request", + "sequence": 3, "call_kind": "activity", "sequence_span": 1, **options, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind,reason", [ + (kind, "cancellation_waiting_for_child") for kind in ("child", "parallel", "selection_handle") +] + [ + (kind, "cancellation_waiting_for_activity") + for kind in ("activity", "local_activity", "parallel", "selection_handle") +]) +async def test_pending_cancellation_requires_explicit_claim_release( + client: Client, monkeypatch: pytest.MonkeyPatch, kind: str, reason: str, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + options = {"operation_sequence": 1} if kind == "selection_handle" else {} + pending = pending_delivery_response(reason=reason) + with patch.object(client._http, "request", new_callable=AsyncMock, return_value=response(pending)): + result = await client.deliver_workflow_cancellation( + task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, + request_id="original-request", sequence=3, call_kind=kind, **options, + ) + assert result == pending + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ["child", "activity"]) +@pytest.mark.parametrize("change", [ + {"claim_released": False}, {"claim_released": "true"}, {"claim_released": None}, + {"task_id": "other"}, {"reason": "other"}, {"reason": []}, {"request_id": "original-request"}, + {"sequence": 3}, {"call_kind": "child"}, {"sequence_span": 1}, + {"operation_sequence": 1}, {"operation_sequence_span": 1}, {"delivered": 0}, +]) +async def test_malformed_pending_cancellation_ack_is_rejected( + client: Client, monkeypatch: pytest.MonkeyPatch, change: dict[str, Any], kind: str, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + pending = pending_delivery_response(reason="cancellation_waiting_for_" + kind) | change + with ( + patch.object(client._http, "request", new_callable=AsyncMock, + return_value=response(pending)), + pytest.raises(ServerError) as error, + ): + await client.deliver_workflow_cancellation( + task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, + request_id="original-request", sequence=3, call_kind=kind, + ) + assert error.value.reason() == "invalid_cooperative_cancellation_delivery" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind,reason", [ + ("timer", "cancellation_waiting_for_child"), ("timer", "cancellation_waiting_for_activity"), + ("child", "cancellation_waiting_for_activity"), ("activity", "cancellation_waiting_for_child"), +]) +async def test_pending_reply_cannot_release_an_unrelated_claim( + client: Client, monkeypatch: pytest.MonkeyPatch, kind: str, reason: str, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + with ( + patch.object(client._http, "request", new_callable=AsyncMock, + return_value=response(pending_delivery_response(reason=reason))), + pytest.raises(ServerError), + ): + await client.deliver_workflow_cancellation( + task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, + request_id="original-request", sequence=3, call_kind=kind, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("version", ["1.19", "2.20", "malformed"]) +async def test_delivery_requires_explicit_compatible_protocol( + client: Client, monkeypatch: pytest.MonkeyPatch, version: str, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", version) + with patch.object(client._http, "request", new_callable=AsyncMock) as send, pytest.raises(ValueError): + await client.deliver_workflow_cancellation( + task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, + request_id="original-request", sequence=3, call_kind="activity", + ) + send.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", [ + {"workflow_task_attempt": True}, {"sequence": True}, {"call_kind": "made_up"}, + {"sequence_span": 2}, {"call_kind": "selection_handle", "operation_sequence": 3}, + {"operation_sequence": 1}, {"lease_owner": ""}, +]) +async def test_invalid_delivery_is_rejected_before_io( + client: Client, monkeypatch: pytest.MonkeyPatch, change: dict, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + options = {"task_id": "task/1", "lease_owner": "worker-1", "workflow_task_attempt": 2, + "request_id": "original-request", "sequence": 3, "call_kind": "activity", **change} + with patch.object(client._http, "request", new_callable=AsyncMock) as send, pytest.raises(ValueError): + await client.deliver_workflow_cancellation(**options) + send.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", [ + {"task_id": "other"}, {"request_id": "other"}, {"sequence": 4}, {"sequence": True}, + {"sequence_span": True}, {"operation_sequence_span": True}, {"call_kind": "timer"}, {"delivered": 1}, +]) +async def test_delivery_ack_must_match_exact_canonical_boundary( + client: Client, monkeypatch: pytest.MonkeyPatch, change: dict, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + with patch.object(client._http, "request", new_callable=AsyncMock, + return_value=response(delivery_response(**change))), pytest.raises(ServerError) as error: + await client.deliver_workflow_cancellation( + task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, + request_id="original-request", sequence=3, call_kind="activity", + ) + assert error.value.reason() == "invalid_cooperative_cancellation_delivery" + + +@pytest.mark.asyncio +async def test_delivery_retries_identical_request_after_acknowledgment_loss( + client: Client, monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + with patch.object(client._http, "request", new_callable=AsyncMock, + side_effect=[httpx.ReadTimeout("acknowledgment lost"), response(delivery_response())]) as send: + await client.deliver_workflow_cancellation( + task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, + request_id="original-request", sequence=3, call_kind="activity", + ) + assert send.await_count == 2 + assert send.await_args_list[0] == send.await_args_list[1] + + +@pytest.mark.asyncio +async def test_active_claim_refusal_does_not_fall_back_to_immediate_cancel(client: Client) -> None: + reason = "active_claim_cancellation_not_supported" + with ( + patch.object(client._http, "request", new_callable=AsyncMock, + return_value=response({"reason": reason}, 409)) as send, + pytest.raises(ServerError) as error, + ): + await client.request_workflow_cancellation("wf/1") + assert error.value.reason() == reason + assert send.await_count == 1 + assert send.call_args.args[1].endswith("/request-cancellation") + + +@pytest.mark.asyncio +async def test_activity_stop_receipt_uses_original_identity_and_worker_credential( + client: Client, monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + with patch.object(client._http, "request", new_callable=AsyncMock, + side_effect=[response(activity_stop_response()), + response(activity_stop_response(duplicate=True))]) as send: + original = await client.acknowledge_activity_cancellation( + task_id="task/1", activity_attempt_id="attempt-1", lease_owner="worker-1", request_id="original-request", + ) + duplicate = await client.acknowledge_activity_cancellation( + task_id="task/1", activity_attempt_id="attempt-1", lease_owner="worker-1", request_id="original-request", + ) + assert original["history_event_id"] == duplicate["history_event_id"] == "original-event" + assert duplicate["duplicate"] is True + for call in send.await_args_list: + assert call.args[:2] == ("POST", "/api/worker/activity-tasks/task%2F1/acknowledge-cancellation") + assert call.kwargs["headers"]["Authorization"] == "Bearer worker-token" + assert call.kwargs["headers"]["X-Namespace"] == "ns1" + assert call.kwargs["headers"]["X-Durable-Workflow-Protocol-Version"] == "1.20" + assert call.kwargs["json"] == { + "activity_attempt_id": "attempt-1", "lease_owner": "worker-1", "request_id": "original-request", + } + assert call.kwargs["timeout"] == 5.0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", [ + {"task_id": "other"}, {"activity_attempt_id": "other"}, {"lease_owner": "other"}, {"request_id": "other"}, + {"acknowledged": 1}, {"acknowledged": False}, {"duplicate": 1}, {"reason": "refused"}, + {"heartbeat_recorded": True}, {"heartbeat_recorded": 0}, {"history_event_id": None}, {"history_event_id": " "}, +]) +async def test_activity_stop_receipt_rejects_mismatched_or_unproved_response( + client: Client, monkeypatch: pytest.MonkeyPatch, change: dict[str, Any], +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + with ( + patch.object(client._http, "request", new_callable=AsyncMock, + return_value=response(activity_stop_response(**change))), + pytest.raises(ServerError) as error, + ): + await client.acknowledge_activity_cancellation( + task_id="task/1", activity_attempt_id="attempt-1", lease_owner="worker-1", request_id="original-request", + ) + assert error.value.reason() == "invalid_activity_cancellation_acknowledgement" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("field", ["task_id", "activity_attempt_id", "lease_owner", "request_id"]) +@pytest.mark.parametrize("identity", [None, True, "", " ", "x" * 256, "é" * 128]) +async def test_activity_stop_receipt_rejects_invalid_identity_before_io( + client: Client, monkeypatch: pytest.MonkeyPatch, field: str, identity: Any, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + options = {"task_id": "task/1", "activity_attempt_id": "attempt-1", "lease_owner": "worker-1", + "request_id": "original-request", field: identity} + with patch.object(client._http, "request", new_callable=AsyncMock) as send, pytest.raises(ValueError): + await client.acknowledge_activity_cancellation(**options) + send.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("version", [None, "1.19", "2.20", "invalid"]) +async def test_activity_stop_receipt_requires_explicit_worker_opt_in( + client: Client, monkeypatch: pytest.MonkeyPatch, version: str | None, +) -> None: + if version is None: + monkeypatch.delenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", raising=False) + else: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", version) + with patch.object(client._http, "request", new_callable=AsyncMock) as send, pytest.raises(ValueError): + await client.acknowledge_activity_cancellation( + task_id="task/1", activity_attempt_id="attempt-1", lease_owner="worker-1", request_id="original-request", + ) + send.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("status", "reason"), [(409, "lease_owner_mismatch"), (503, "storage_fenced")]) +async def test_activity_stop_receipt_preserves_server_refusal( + client: Client, monkeypatch: pytest.MonkeyPatch, status: int, reason: str, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + with ( + patch.object(client._http, "request", new_callable=AsyncMock, + return_value=response({"reason": reason, "request_admitted": False}, status)) as send, + pytest.raises(ServerError) as error, + ): + await client.acknowledge_activity_cancellation( + task_id="task/1", activity_attempt_id="attempt-1", lease_owner="worker-1", request_id="original-request", + ) + assert error.value.reason() == reason + assert send.await_count == (1 if status == 409 else 3) + assert all(call == send.await_args_list[0] for call in send.await_args_list) + + +@pytest.mark.asyncio +async def test_activity_stop_receipt_retries_original_identity_after_response_loss( + client: Client, monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + with patch.object(client._http, "request", new_callable=AsyncMock, + side_effect=[httpx.ReadTimeout("receipt response lost"), + response(activity_stop_response(duplicate=True))]) as send: + result = await client.acknowledge_activity_cancellation( + task_id="task/1", activity_attempt_id="attempt-1", lease_owner="worker-1", request_id="original-request", + ) + assert result["duplicate"] is True + assert result["history_event_id"] == "original-event" + assert send.await_count == 2 + assert send.await_args_list[0] == send.await_args_list[1] + + +@pytest.mark.asyncio +async def test_activity_stop_receipt_transport_has_one_total_budget( + client: Client, monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + stopped = asyncio.Event() + + async def blocked(*args: Any, **kwargs: Any) -> httpx.Response: + try: + await asyncio.Event().wait() + finally: + stopped.set() + raise AssertionError("blocked transport resumed") + + started = time.monotonic() + with ( + patch.object(client._http, "request", new_callable=AsyncMock, side_effect=blocked) as send, + pytest.raises(asyncio.TimeoutError), + ): + await client.acknowledge_activity_cancellation( + task_id="task/1", activity_attempt_id="attempt-1", lease_owner="worker-1", request_id="original-request", + ) + assert time.monotonic() - started < 6.0 + assert stopped.is_set() + assert send.await_count == 1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("reason", ["lease_owner_mismatch", "workflow_task_attempt_mismatch", "lease_expired"]) +async def test_delivery_refusal_preserves_the_server_lease_reason( + client: Client, monkeypatch: pytest.MonkeyPatch, reason: str, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + with ( + patch.object(client._http, "request", new_callable=AsyncMock, + return_value=response({"delivered": False, "reason": reason}, 409)) as send, + pytest.raises(ServerError) as error, + ): + await client.deliver_workflow_cancellation( + task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, + request_id="original-request", sequence=3, call_kind="activity", + ) + assert error.value.reason() == reason + assert send.await_count == 1 diff --git a/tests/test_cooperative_cancellation_worker.py b/tests/test_cooperative_cancellation_worker.py new file mode 100644 index 0000000..692a337 --- /dev/null +++ b/tests/test_cooperative_cancellation_worker.py @@ -0,0 +1,611 @@ +from __future__ import annotations + +import asyncio +import threading +from concurrent.futures import ThreadPoolExecutor +from copy import deepcopy +from typing import Any +from unittest.mock import AsyncMock + +import pytest + +from durable_workflow import activity, serializer, workflow +from durable_workflow.client import Client +from durable_workflow.errors import ServerError, WorkflowCancelled +from durable_workflow.worker import Worker +from durable_workflow.workflow import LocalActivityExecutionAborted +from tests.test_cooperative_cancellation import marker, observation, request +from tests.test_worker import compatible_cluster_info + + +@workflow.defn(name="worker-cancellation") +class CancellationWorkflow: + def run(self, ctx: workflow.WorkflowContext, kind: str): # type: ignore[no-untyped-def] + ctx.throw_if_cancellation_requested() + try: + if kind == "local_activity": + yield ctx.local_activity("work", []) + elif kind == "activity": + yield ctx.schedule_activity("work", []) + elif kind == "child": + yield ctx.start_child_workflow("child", []) + elif kind == "parallel": + yield [ctx.start_timer(20), ctx.schedule_activity("work", [])] + else: + yield ctx.start_timer(30) + except WorkflowCancelled as error: + with ctx.cancellation_shield(): + yield ctx.start_timer(1) + return error.request_id + return "not cancelled" + + +@workflow.defn(name="worker-cancellation-prior") +class PriorCommandsWorkflow: + def run(self, ctx: workflow.WorkflowContext): # type: ignore[no-untyped-def] + yield ctx.upsert_memo({"stage": "before-boundary"}) + yield ctx.start_timer(30) + + +@workflow.defn(name="worker-cancellation-local-cleanup") +class LocalCleanupWorkflow: + def run(self, ctx: workflow.WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(30) + except WorkflowCancelled as error: + with ctx.cancellation_shield(): + yield ctx.local_activity("cleanup", [error.request_id]) + return error.request_id + + +def claimed_task(kind: str = "timer", *, observed: bool = True) -> dict[str, Any]: + return { + "task_id": "task-1", "workflow_id": "workflow-1", "run_id": "run-1", + "workflow_type": "worker-cancellation", "workflow_task_attempt": 4, + "payload_codec": "avro", "arguments": serializer.envelope([kind], codec="avro"), + "history_events": [], **({"cancellation_request": observation()} if observed else {}), + } + + +def lease_ack(*, observed: bool = False) -> dict[str, Any]: + return { + "task_id": "task-1", "lease_owner": "cooperative-worker", "workflow_task_attempt": 4, + "renewed": True, **({"cancellation_request": observation()} if observed else {}), + } + + +class ClaimServer: + def __init__( + self, *, history: list[dict[str, Any]] | None = None, + lease_owner: str = "cooperative-worker", workflow_task_attempt: int = 4, + task_id: str = "task-1", + ) -> None: + self.client = AsyncMock(spec=Client) + self.lease_owner = lease_owner + self.workflow_task_attempt = workflow_task_attempt + self.task_id = task_id + self.history = list(history if history is not None else [request()]) + self.trace: list[str] = [] + self.delivery_error: Exception | None = None + self.commit_delivery = True + self.pending_deliveries = 0 + self.client.get_cluster_info.return_value = compatible_cluster_info(worker_protocol={ + "version": "1.20", "server_capabilities": { + "query_tasks": True, "long_poll_timeout": 30, "cooperative_cancellation": True, + "workflow_memo_updates": True, + }, + }) + self.client.register_worker.return_value = {"registered": True} + self.client.heartbeat_workflow_task.return_value = { + **lease_ack(), "lease_owner": lease_owner, "workflow_task_attempt": workflow_task_attempt, + } + self.client.workflow_task_history.side_effect = self.page + self.client.deliver_workflow_cancellation.side_effect = self.deliver + self.client.complete_workflow_task.side_effect = self.complete + + async def page(self, **kwargs: Any) -> dict[str, Any]: + assert kwargs == { + "task_id": self.task_id, "next_history_page_token": "opaque-first-page", + "lease_owner": self.lease_owner, "workflow_task_attempt": self.workflow_task_attempt, + } + self.trace.append("history") + return {"history_events": deepcopy(self.history), "next_history_page_token": None} + + async def deliver(self, **kwargs: Any) -> dict[str, Any]: + self.trace.append("delivery") + if self.pending_deliveries: + self.pending_deliveries -= 1 + return { + "delivered": False, "task_id": self.task_id, + "reason": "cancellation_waiting_for_child", "claim_released": True, + } + if self.commit_delivery: + self.history.append(marker( + kwargs["sequence"], kwargs["call_kind"], sequence_span=kwargs["sequence_span"], + operation_sequence=kwargs["operation_sequence"], + operation_sequence_span=kwargs["operation_sequence_span"], + )) + if self.delivery_error: + raise self.delivery_error + return {"delivered": True} + + async def complete(self, **kwargs: Any) -> dict[str, Any]: + self.trace.append("completion") + return {"outcome": "completed"} + + async def worker(self, monkeypatch: pytest.MonkeyPatch, **kwargs: Any) -> Worker: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + # These transport/replay fixtures use local callbacks that share mock + # state with this process. They never poll or advertise remote handlers. + local_callbacks = kwargs.pop("activities", []) + worker = Worker( + self.client, task_queue="queue", worker_id=self.lease_owner, + workflows=kwargs.pop("workflows", [CancellationWorkflow]), + capabilities=["cooperative_cancellation"], **kwargs, + ) + await worker._register() + worker.activities.update({callback.__activity_name__: callback for callback in local_callbacks}) + return worker + + +@pytest.mark.parametrize("kind,span", [("timer", 1), ("activity", 1), ("local_activity", 1), ("parallel", 2)]) +async def test_claim_commits_delivery_and_reloads_before_cleanup( + monkeypatch: pytest.MonkeyPatch, kind: str, span: int, +) -> None: + server = ClaimServer() + worker = await server.worker(monkeypatch) + commands = await worker._run_workflow_task(claimed_task(kind)) + assert server.trace == ["history", "delivery", "history", "completion"] + assert commands is not None and len(commands) == 1 + assert commands[0]["type"] == "start_timer" + assert commands[0]["delay_seconds"] == 1 + assert server.client.deliver_workflow_cancellation.await_args.kwargs == { + "task_id": "task-1", "lease_owner": "cooperative-worker", "workflow_task_attempt": 4, + "request_id": "request-1", "sequence": 1, "call_kind": kind, "sequence_span": span, + "operation_sequence": None, "operation_sequence_span": 1, + } + server.client.fail_workflow_task.assert_not_awaited() + + +async def test_pending_child_returns_to_other_work_and_replays_on_a_new_claim( + monkeypatch: pytest.MonkeyPatch, +) -> None: + @workflow.defn(name="other-work") + class OtherWorkflow: + def run(self, ctx: workflow.WorkflowContext, kind: str) -> str: + return kind + + server = ClaimServer() + server.pending_deliveries = 1 + worker = await server.worker(monkeypatch, workflows=[CancellationWorkflow, OtherWorkflow]) + assert await worker._run_workflow_task(claimed_task("child")) is None + assert server.trace == ["history", "delivery"] + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + + other = claimed_task(observed=False) + other.update(task_id="other-task", workflow_id="other-workflow", run_id="other-run", workflow_type="other-work") + commands = await worker._run_workflow_task(other) + assert commands is not None and commands[0]["type"] == "complete_workflow" + + server.task_id = "parent-resume-task" + server.workflow_task_attempt = 1 + resumed = claimed_task("child") + resumed.update(task_id=server.task_id, workflow_task_attempt=1) + original_request = deepcopy(resumed["cancellation_request"]) + commands = await worker._run_workflow_task(resumed) + assert commands is not None and commands[0]["type"] == "start_timer" + assert commands[0]["delay_seconds"] == 1 + assert server.trace == ["history", "delivery", "completion", "history", "delivery", "history", "completion"] + assert resumed["cancellation_request"] == original_request + assert server.client.deliver_workflow_cancellation.await_args.kwargs["workflow_task_attempt"] == 1 + assert server.client.deliver_workflow_cancellation.await_args.kwargs["request_id"] == "request-1" + server.client.fail_workflow_task.assert_not_awaited() + + +async def test_lost_delivery_ack_is_resolved_only_from_canonical_history(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer() + server.delivery_error = TimeoutError("ack lost after commit") + worker = await server.worker(monkeypatch) + commands = await worker._run_workflow_task(claimed_task()) + assert commands is not None and commands[0]["type"] == "start_timer" + assert server.trace == ["history", "delivery", "history", "completion"] + assert server.client.deliver_workflow_cancellation.await_count == 1 + + +async def test_legacy_boolean_cannot_replace_canonical_cooperative_delivery(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer() + worker = await server.worker(monkeypatch) + task = claimed_task() + task["cancel_requested"] = True + commands = await worker._run_workflow_task(task) + assert commands is not None and commands[0]["type"] == "start_timer" + assert server.trace == ["history", "delivery", "history", "completion"] + + +@pytest.mark.parametrize("error", [None, TimeoutError("not accepted"), ServerError(409, {"reason": "lease_expired"})]) +async def test_ack_without_matching_marker_cannot_run_cleanup_or_complete( + monkeypatch: pytest.MonkeyPatch, error: Exception | None, +) -> None: + server = ClaimServer() + server.commit_delivery = False + server.delivery_error = error + worker = await server.worker(monkeypatch) + assert await worker._run_workflow_task(claimed_task()) is None + assert server.trace == ["history", "delivery", "history"] + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + + +async def test_claim_loss_during_refresh_cannot_complete_or_fail_stale_attempt(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer() + server.client.workflow_task_history.side_effect = [ + {"history_events": [request()], "next_history_page_token": None}, + ServerError(409, {"reason": "workflow_task_attempt_mismatch"}), + ] + worker = await server.worker(monkeypatch) + assert await worker._run_workflow_task(claimed_task()) is None + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + + +async def test_prior_commands_commit_before_a_successor_can_deliver(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer() + worker = await server.worker(monkeypatch, workflows=[PriorCommandsWorkflow]) + worker._workflow_memo_updates_supported = True + task = claimed_task() + task.update(workflow_type="worker-cancellation-prior", arguments=serializer.envelope([], codec="avro")) + commands = await worker._run_workflow_task(task) + assert commands is not None and [command["type"] for command in commands] == ["upsert_memo"] + server.client.deliver_workflow_cancellation.assert_not_awaited() + assert server.trace == ["history", "completion"] + + +async def test_cold_replacement_replays_marker_without_redelivery(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer(history=[request(), marker(1, "timer")]) + worker = await server.worker(monkeypatch) + task = claimed_task(observed=False) + task["history_events"] = deepcopy(server.history) + commands = await worker._run_workflow_task(task) + assert commands is not None and commands[0]["delay_seconds"] == 1 + server.client.deliver_workflow_cancellation.assert_not_awaited() + assert server.trace == ["completion"] + + +async def test_local_result_after_request_observation_is_not_serialized_or_committed( + monkeypatch: pytest.MonkeyPatch, +) -> None: + server = ClaimServer(history=[]) + + @activity.defn(name="work") + async def work() -> object: + server.trace.append("local-result") + server.history.append(request()) + return object() # Would fail serialization if the late value escaped the fence. + + server.client.heartbeat_workflow_task.side_effect = [lease_ack(), lease_ack(observed=True)] + worker = await server.worker(monkeypatch, activities=[work]) + commands = await worker._run_workflow_task(claimed_task("local_activity", observed=False)) + assert server.trace == ["local-result", "history", "delivery", "history", "completion"] + assert commands is not None and [command["type"] for command in commands] == ["start_timer"] + assert server.client.deliver_workflow_cancellation.await_args.kwargs["call_kind"] == "local_activity" + server.client.fail_workflow_task.assert_not_awaited() + + +async def test_local_heartbeat_observation_does_not_become_an_activity_failure(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer(history=[]) + + @activity.defn(name="work") + async def work() -> str: + server.history.append(request()) + await activity.context().heartbeat() + raise AssertionError("local execution should have returned to canonical replay") + + server.client.heartbeat_workflow_task.side_effect = [lease_ack(), lease_ack(observed=True)] + worker = await server.worker(monkeypatch, activities=[work]) + commands = await worker._run_workflow_task(claimed_task("local_activity", observed=False)) + assert commands is not None and [command["type"] for command in commands] == ["start_timer"] + server.client.fail_workflow_task.assert_not_awaited() + + +async def test_active_local_call_without_user_heartbeat_is_fenced_before_delivery( + monkeypatch: pytest.MonkeyPatch, +) -> None: + server = ClaimServer(history=[]) + late_result_discarded = asyncio.Event() + + @activity.defn(name="work") + async def work() -> object: + server.history.append(request()) + server.client.heartbeat_workflow_task.return_value = lease_ack(observed=True) + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + late_result_discarded.set() + return object() + + worker = await server.worker(monkeypatch, activities=[work], heartbeat_interval=0.01) + commands = await asyncio.wait_for( + worker._run_workflow_task(claimed_task("local_activity", observed=False)), timeout=2, + ) + await asyncio.wait_for(late_result_discarded.wait(), timeout=1) + assert commands is not None and [command["type"] for command in commands] == ["start_timer"] + assert server.client.deliver_workflow_cancellation.await_args.kwargs["call_kind"] == "local_activity" + server.client.fail_workflow_task.assert_not_awaited() + + +@pytest.mark.parametrize("field,value", [ + ("renewed", False), ("lease_owner", "replacement-worker"), + ("workflow_task_attempt", 5), ("task_id", "replacement-task"), +]) +async def test_active_local_lease_refusal_drops_work_without_delivery( + monkeypatch: pytest.MonkeyPatch, field: str, value: Any, +) -> None: + server = ClaimServer(history=[]) + abandoned = asyncio.Event() + + @activity.defn(name="work") + async def work() -> None: + ack = lease_ack(observed=True) + ack[field] = value + server.client.heartbeat_workflow_task.return_value = ack + try: + await asyncio.Event().wait() + finally: + abandoned.set() + + worker = await server.worker(monkeypatch, activities=[work], heartbeat_interval=0.01) + task = claimed_task("local_activity", observed=False) + assert await asyncio.wait_for(worker._run_workflow_task(task), timeout=2) is None + await asyncio.wait_for(abandoned.wait(), timeout=1) + assert "cancellation_request" not in task + server.client.deliver_workflow_cancellation.assert_not_awaited() + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + + +async def test_shielded_local_cleanup_can_heartbeat_the_delivered_request(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer() + cleanup: list[str] = [] + + @activity.defn(name="cleanup") + async def compensate(request_id: str) -> str: + cleanup.append(request_id) + await activity.context().heartbeat() + return request_id + + server.client.heartbeat_workflow_task.return_value = lease_ack(observed=True) + worker = await server.worker(monkeypatch, workflows=[LocalCleanupWorkflow], activities=[compensate]) + task = claimed_task() + task.update(workflow_type="worker-cancellation-local-cleanup", arguments=serializer.envelope([], codec="avro")) + commands = await worker._run_workflow_task(task) + assert cleanup == ["request-1"] + assert commands is not None + assert [command["type"] for command in commands] == ["record_local_activity", "complete_workflow"] + assert server.client.deliver_workflow_cancellation.await_count == 1 + + +async def test_worker_stop_drains_shielded_local_cleanup_within_its_grace_period( + monkeypatch: pytest.MonkeyPatch, +) -> None: + server = ClaimServer() + entered = asyncio.Event() + release = asyncio.Event() + + @activity.defn(name="cleanup") + async def compensate(request_id: str) -> str: + entered.set() + await release.wait() + await activity.context().heartbeat() + return request_id + + worker = await server.worker( + monkeypatch, workflows=[LocalCleanupWorkflow], activities=[compensate], shutdown_timeout=1, + ) + task = claimed_task() + task.update(workflow_type="worker-cancellation-local-cleanup", arguments=serializer.envelope([], codec="avro")) + execution = worker._track(worker._run_workflow_task(task)) + await asyncio.wait_for(entered.wait(), timeout=1) + stopping = asyncio.create_task(worker.stop()) + await asyncio.sleep(0) + assert worker._stop.is_set() + release.set() + commands = await asyncio.wait_for(execution, timeout=1) + await asyncio.wait_for(stopping, timeout=1) + assert commands is not None + assert [command["type"] for command in commands] == ["record_local_activity", "complete_workflow"] + server.client.deliver_workflow_cancellation.assert_awaited_once() + server.client.fail_workflow_task.assert_not_awaited() + server.client.deregister_worker_registration.assert_awaited_once_with("cooperative-worker") + + +@pytest.mark.parametrize("handler_kind", ["async", "sync"]) +async def test_shutdown_timeout_abandons_cleanup_and_replacement_replays_original_delivery( + monkeypatch: pytest.MonkeyPatch, handler_kind: str, +) -> None: + server = ClaimServer() + entered = asyncio.Event() + discarded = asyncio.Event() + loop = asyncio.get_running_loop() + release_thread = threading.Event() + fenced_heartbeats: list[bool] = [] + + @activity.defn(name="cleanup") + async def compensate(request_id: str) -> object: + assert request_id == "request-1" + entered.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + with pytest.raises(LocalActivityExecutionAborted): + await activity.context().heartbeat() + fenced_heartbeats.append(True) + discarded.set() + return object() # A late, unencodable result must never become an activity failure. + + @activity.defn(name="cleanup") + def synchronous_compensation(request_id: str) -> object: + assert request_id == "request-1" + context = activity.context() + assert context.info.worker_id == "cooperative-worker" + loop.call_soon_threadsafe(entered.set) + try: + assert release_thread.wait(timeout=2) + heartbeat = asyncio.run_coroutine_threadsafe(context.heartbeat(), loop) + try: + heartbeat.result(timeout=1) + except LocalActivityExecutionAborted: + fenced_heartbeats.append(True) + else: + fenced_heartbeats.append(False) + return object() + finally: + loop.call_soon_threadsafe(discarded.set) + + worker = await server.worker( + monkeypatch, workflows=[LocalCleanupWorkflow], + activities=[compensate if handler_kind == "async" else synchronous_compensation], shutdown_timeout=0.01, + ) + task = claimed_task() + task.update(workflow_type="worker-cancellation-local-cleanup", arguments=serializer.envelope([], codec="avro")) + execution = worker._track(worker._run_workflow_task(task)) + await asyncio.wait_for(entered.wait(), timeout=1) + original_history = deepcopy(server.history) + heartbeats_before_stop = server.client.heartbeat_workflow_task.await_count + await asyncio.wait_for(worker.stop(), timeout=1) + release_thread.set() + await asyncio.wait_for(discarded.wait(), timeout=1) + assert fenced_heartbeats == [True] + assert server.client.heartbeat_workflow_task.await_count == heartbeats_before_stop + assert execution.cancelled() or execution.result() is None + assert server.history == original_history + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + server.client.deregister_worker_registration.assert_awaited_once_with("cooperative-worker") + + replacement = ClaimServer(history=original_history, lease_owner="replacement-worker", workflow_task_attempt=5) + cleanup_ids: list[str] = [] + + @activity.defn(name="cleanup") + async def resumed_cleanup(request_id: str) -> str: + cleanup_ids.append(request_id) + await activity.context().heartbeat() + return request_id + + successor = await replacement.worker(monkeypatch, workflows=[LocalCleanupWorkflow], activities=[resumed_cleanup]) + reclaimed = deepcopy(task) + reclaimed["workflow_task_attempt"] = 5 + reclaimed["history_events"] = deepcopy(original_history) + commands = await asyncio.wait_for(successor._run_workflow_task(reclaimed), timeout=1) + assert cleanup_ids == ["request-1"] + assert commands is not None + assert [command["type"] for command in commands] == ["record_local_activity", "complete_workflow"] + assert replacement.history == original_history + replacement.client.deliver_workflow_cancellation.assert_not_awaited() + replacement.client.fail_workflow_task.assert_not_awaited() + assert replacement.client.complete_workflow_task.await_args.kwargs["lease_owner"] == "replacement-worker" + assert replacement.client.complete_workflow_task.await_args.kwargs["workflow_task_attempt"] == 5 + + +async def test_active_synchronous_local_call_observes_request_without_blocking_the_worker( + monkeypatch: pytest.MonkeyPatch, +) -> None: + server = ClaimServer(history=[]) + entered = asyncio.Event() + discarded = asyncio.Event() + release = threading.Event() + loop = asyncio.get_running_loop() + fenced: list[bool] = [] + loop.set_default_executor(ThreadPoolExecutor(max_workers=1)) + + @activity.defn(name="work") + def work() -> object: + context = activity.context() + loop.call_soon_threadsafe(entered.set) + try: + assert release.wait(timeout=2) + heartbeat = asyncio.run_coroutine_threadsafe(context.heartbeat(), loop) + try: + heartbeat.result(timeout=1) + except LocalActivityExecutionAborted: + fenced.append(True) + else: + fenced.append(False) + return object() + finally: + loop.call_soon_threadsafe(discarded.set) + + worker = await server.worker(monkeypatch, activities=[work], heartbeat_interval=0.01) + execution = worker._track(worker._run_workflow_task(claimed_task("local_activity", observed=False))) + await asyncio.wait_for(entered.wait(), timeout=1) + server.history.append(request()) + server.client.heartbeat_workflow_task.return_value = lease_ack(observed=True) + try: + commands = await asyncio.wait_for(execution, timeout=1) + assert commands is not None and [command["type"] for command in commands] == ["start_timer"] + heartbeats_after_delivery = server.client.heartbeat_workflow_task.await_count + finally: + release.set() + await asyncio.wait_for(discarded.wait(), timeout=1) + assert fenced == [True] + assert server.client.heartbeat_workflow_task.await_count == heartbeats_after_delivery + assert server.client.deliver_workflow_cancellation.await_args.kwargs["call_kind"] == "local_activity" + server.client.fail_workflow_task.assert_not_awaited() + await worker.stop() + + +@pytest.mark.parametrize("value", ["", None, {}, "not-a-page"]) +async def test_invalid_or_missing_refresh_token_never_delivers(monkeypatch: pytest.MonkeyPatch, value: Any) -> None: + server = ClaimServer() + worker = await server.worker(monkeypatch) + task = claimed_task() + task["cancellation_request"]["history_refresh_page_token"] = value + assert await worker._run_workflow_task(task) is None + server.client.deliver_workflow_cancellation.assert_not_awaited() + server.client.complete_workflow_task.assert_not_awaited() + + +async def test_canonical_refresh_rejects_repeated_page_token(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer() + server.client.workflow_task_history.side_effect = [{ + "history_events": [request()], "next_history_page_token": "opaque-first-page", + }] + worker = await server.worker(monkeypatch) + assert await worker._run_workflow_task(claimed_task()) is None + assert server.client.workflow_task_history.await_count == 1 + server.client.deliver_workflow_cancellation.assert_not_awaited() + server.client.complete_workflow_task.assert_not_awaited() + + +@pytest.mark.parametrize("version", ["1.19", "2.20", "1.bad"]) +async def test_explicit_cooperative_worker_refuses_incompatible_protocol( + monkeypatch: pytest.MonkeyPatch, version: str, +) -> None: + server = ClaimServer() + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", version) + worker = Worker(server.client, task_queue="queue", capabilities=["cooperative_cancellation"]) + with pytest.raises(RuntimeError, match="protocol 1.20"): + await worker._register() + server.client.register_worker.assert_not_awaited() + + +async def test_current_default_worker_does_not_advertise_cooperative_support(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer() + monkeypatch.delenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", raising=False) + worker = Worker(server.client, task_queue="queue") + await worker._register() + assert worker._cooperative_cancellation_supported is False + assert "cooperative_cancellation" not in server.client.register_worker.await_args.kwargs["capabilities"] + + +async def test_incapable_worker_cannot_execute_a_canonical_cooperative_run(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer(history=[request(), marker(1, "timer")]) + monkeypatch.delenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", raising=False) + worker = Worker(server.client, task_queue="queue", workflows=[CancellationWorkflow]) + await worker._register() + task = claimed_task(observed=False) + task["history_events"] = deepcopy(server.history) + assert await worker._run_workflow_task(task) is None + server.client.deliver_workflow_cancellation.assert_not_awaited() + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() diff --git a/tests/test_cooperative_remote_worker.py b/tests/test_cooperative_remote_worker.py new file mode 100644 index 0000000..ca61e36 --- /dev/null +++ b/tests/test_cooperative_remote_worker.py @@ -0,0 +1,335 @@ +from __future__ import annotations + +import asyncio +import os +import shutil +import time +from copy import deepcopy +from pathlib import Path +from typing import Any +from unittest.mock import AsyncMock, patch + +import httpx +import pytest +import pytest_asyncio + +from durable_workflow import activity, serializer +from durable_workflow.client import Client +from durable_workflow.errors import ActivityCancelled, NonRetryableError, ServerError +from durable_workflow.retry_policy import TransportRetryPolicy +from durable_workflow.worker import Worker +from tests.test_activity_process import wait_for_exit, wait_for_file +from tests.test_worker import compatible_cluster_info + + +def task(*args: Any) -> dict[str, Any]: + return {"task_id": "remote-task", "activity_attempt_id": "attempt", "activity_type": "remote", + "payload_codec": "avro", "arguments": serializer.envelope(list(args), codec="avro")} + + +def status() -> dict[str, Any]: + return {"task_id": "remote-task", "activity_attempt_id": "attempt", "lease_owner": "owner", + "can_continue": True, "cancel_requested": False, "reason": None, "heartbeat_recorded": False, + "lease_expires_at": "2100-01-01T00:00:00Z", "deadlines": None, "worker_session": None} + + +def cancellation_status() -> dict[str, Any]: + return {**status(), "can_continue": False, "cancel_requested": True, "reason": "activity_cancelled", + "cancellation_acknowledgement": {"request_id": "original-request", "root_request_id": "root-request", + "cleanup_deadline_at": "2100-01-01T00:00:00Z", + "cancellation_history_event_id": "cancellation-event", + "callback_state": "unknown"}} + + +def stop_receipt() -> dict[str, Any]: + return {"task_id": "remote-task", "activity_attempt_id": "attempt", "lease_owner": "owner", + "request_id": "original-request", "acknowledged": True, "duplicate": False, "reason": None, + "heartbeat_recorded": False, "history_event_id": "receipt-event"} + + +async def blocked_async(marker: str, heartbeat: bool = False) -> None: + if heartbeat: + await activity.context().heartbeat({"authored": True}) + Path(marker).write_text(str(os.getpid())) + await asyncio.Event().wait() + + +def blocked_sync(marker: str, heartbeat: bool = False) -> None: + if heartbeat: + asyncio.run(activity.context().heartbeat({"authored": True})) + Path(marker).write_text(str(os.getpid())) + time.sleep(60) + + +def released_sync(marker: str) -> bytes: + Path(marker).write_text(str(os.getpid())) + while not Path(marker + ".release").exists(): + time.sleep(0.02) + return b"late result" + + +async def typed_async() -> dict[str, Any]: + await activity.context().heartbeat({"authored": True}) + return {"bytes": b"\x00\xff", "value": 42, "null": None, "task_id": activity.context().info.task_id} + + +def typed_sync() -> dict[str, Any]: + return asyncio.run(typed_async()) + + +async def application_failure(kind: str) -> None: + if kind == "ValueError": + raise ValueError("original") + if kind == "NonRetryableError": + raise NonRetryableError("original") + raise ActivityCancelled() + + +@pytest_asyncio.fixture +async def owner(monkeypatch: pytest.MonkeyPatch): # type: ignore[no-untyped-def] + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + client = AsyncMock(spec=Client) + client.get_cluster_info.return_value = compatible_cluster_info(worker_protocol={ + "version": "1.20", "server_capabilities": {"cooperative_cancellation": True}, + }) + client.register_worker.return_value = {"registered": True} + client.activity_task_status.return_value = status() + client.heartbeat_activity_task.return_value = status() + client.acknowledge_activity_cancellation.return_value = stop_receipt() + worker = Worker(client, task_queue="queue", worker_id="owner", capabilities=["cooperative_cancellation"], + max_concurrent_activity_tasks=1) + await worker._register() + try: + yield worker, client + finally: + await worker.stop() + + +@pytest.mark.parametrize("reason", ["cancel", "replacement", "backend", "shutdown"]) +@pytest.mark.parametrize("synchronous", [False, True]) +@pytest.mark.parametrize("heartbeat", [False, True]) +async def test_authority_loss_stops_and_joins_the_actual_callback_before_reporting( + owner: Any, tmp_path: Path, reason: str, synchronous: bool, heartbeat: bool, +) -> None: + worker, client = owner + marker = tmp_path / "callback" + worker.activities["remote"] = blocked_sync if synchronous else blocked_async + execution = asyncio.create_task(worker._run_activity_task(task(str(marker), heartbeat))) + await wait_for_file(marker) + pid = int(marker.read_text()) + + async def acknowledge(**options: Any) -> dict[str, Any]: + with pytest.raises(ProcessLookupError): + os.kill(pid, 0) + assert not worker._remote_activity_processes + assert options == {"task_id": "remote-task", "activity_attempt_id": "attempt", "lease_owner": "owner", + "request_id": "original-request"} + return stop_receipt() + + client.acknowledge_activity_cancellation.side_effect = acknowledge + if reason == "cancel": + client.activity_task_status.return_value = cancellation_status() + elif reason == "replacement": + client.activity_task_status.return_value = {**status(), "activity_attempt_id": "replacement"} + elif reason == "backend": + client.activity_task_status.side_effect = OSError("backend unavailable") + else: + worker._local_activity_shutdown.set() + assert await asyncio.wait_for(execution, timeout=4.0) == "claim_aborted" + await wait_for_exit(pid) + assert not worker._remote_activity_processes + assert worker._current_task_slots()["activity_available"] == 1 + assert client.heartbeat_activity_task.await_count == int(heartbeat) + assert client.acknowledge_activity_cancellation.await_count == int(reason == "cancel") + client.complete_activity_task.assert_not_awaited() + client.fail_activity_task.assert_not_awaited() + + +async def test_shutdown_timeout_joins_the_remote_callback_before_returning(owner: Any, tmp_path: Path) -> None: + worker, client = owner + worker._shutdown_timeout = 0.01 + marker = tmp_path / "callback" + worker.activities["remote"] = blocked_async + execution = worker._track(worker._run_activity_task(task(str(marker)))) + await wait_for_file(marker) + pid = int(marker.read_text()) + await asyncio.wait_for(worker.stop(), timeout=4.0) + await wait_for_exit(pid) + assert execution.cancelled() or execution.result() == "claim_aborted" + assert not worker._remote_activity_processes + client.complete_activity_task.assert_not_awaited() + client.fail_activity_task.assert_not_awaited() + + +@pytest.mark.parametrize("synchronous", [False, True]) +async def test_authored_heartbeat_stays_on_owner_loop_and_preserves_typed_result(owner: Any, synchronous: bool) -> None: + worker, client = owner + loop = asyncio.get_running_loop() + observed_loops = [] + + async def heartbeat(**kwargs: Any) -> dict[str, Any]: + observed_loops.append(asyncio.get_running_loop()) + assert kwargs["details"] == {"authored": True} + return status() + + client.heartbeat_activity_task.side_effect = heartbeat + worker.activities["remote"] = typed_sync if synchronous else typed_async + assert await worker._run_activity_task(task()) == "completed" + assert observed_loops == [loop] + assert client.complete_activity_task.await_args.kwargs["result"] == { + "bytes": b"\x00\xff", "value": 42, "null": None, "task_id": "remote-task", + } + assert not worker._remote_activity_processes + client.acknowledge_activity_cancellation.assert_not_awaited() + client.fail_activity_task.assert_not_awaited() + + +@pytest.mark.parametrize("field", ["lease", "heartbeat", "start_to_close", "schedule_to_close", "session", "malformed"]) +async def test_invalid_or_elapsed_observed_bounds_prevent_callback_start(owner: Any, field: str) -> None: + worker, client = owner + reply = deepcopy(status()) + expired = "2000-01-01T00:00:00Z" + if field == "lease": + reply["lease_expires_at"] = expired + elif field == "malformed": + reply["deadlines"] = "invalid" + elif field == "session": + reply["worker_session"] = {"status": "active", "lease_owner": "owner", + "lease_expires_at": expired, "ttl_expires_at": "2100-01-01T00:00:00Z"} + else: + reply["deadlines"] = {field: expired} + client.activity_task_status.return_value = reply + callback = AsyncMock() + worker.activities["remote"] = callback + assert await worker._run_activity_task(task()) == "claim_aborted" + callback.assert_not_awaited() + client.acknowledge_activity_cancellation.assert_not_awaited() + client.complete_activity_task.assert_not_awaited() + client.fail_activity_task.assert_not_awaited() + + +@pytest.mark.parametrize("kind", ["ValueError", "NonRetryableError", "ActivityCancelled"]) +async def test_genuine_application_failure_preserves_classification(owner: Any, kind: str) -> None: + worker, client = owner + worker.activities["remote"] = application_failure + outcome = await worker._run_activity_task(task(kind)) + assert outcome in {"failed", "failed_non_retryable", "cancelled"} + report = client.fail_activity_task.await_args.kwargs + assert report["failure_type"] == kind + assert report["non_retryable"] is (kind != "ValueError") + assert report["failure_class"].endswith("." + kind) + assert "application_failure" in report["stack_trace"] + client.complete_activity_task.assert_not_awaited() + + +@pytest.mark.parametrize("refusal", ["refused", "unproved"]) +async def test_stop_receipt_failure_never_becomes_result_or_failure(owner: Any, tmp_path: Path, refusal: str) -> None: + worker, client = owner + marker = tmp_path / "callback" + worker.activities["remote"] = blocked_async + execution = asyncio.create_task(worker._run_activity_task(task(str(marker)))) + await wait_for_file(marker) + if refusal == "refused": + client.acknowledge_activity_cancellation.side_effect = ServerError(409, {"reason": "request_mismatch"}) + else: + client.acknowledge_activity_cancellation.return_value = {**stop_receipt(), "acknowledged": False} + client.activity_task_status.return_value = cancellation_status() + assert await asyncio.wait_for(execution, timeout=4.0) == "claim_aborted" + await wait_for_exit(int(marker.read_text())) + client.acknowledge_activity_cancellation.assert_awaited_once() + client.complete_activity_task.assert_not_awaited() + client.fail_activity_task.assert_not_awaited() + + +async def test_dead_supervisor_retains_capacity_and_never_reports_stop(owner: Any, tmp_path: Path) -> None: + worker, client = owner + marker = tmp_path / "callback" + worker.activities["remote"] = released_sync + execution = asyncio.create_task(worker._run_activity_task(task(str(marker)))) + await wait_for_file(marker) + callback = next(iter(worker._remote_activity_processes)) + pid = int(marker.read_text()) + try: + callback._supervisor.kill() + assert await asyncio.wait_for(execution, timeout=4.0) == "claim_aborted" + os.kill(pid, 0) + assert callback.stopped is False + assert worker._current_task_slots()["activity_available"] == 0 + poller = asyncio.create_task(worker._poll_activity_tasks()) + try: + await asyncio.sleep(0.1) + client.poll_activity_task.assert_not_awaited() + finally: + poller.cancel() + with pytest.raises(asyncio.CancelledError): + await poller + client.acknowledge_activity_cancellation.assert_not_awaited() + client.complete_activity_task.assert_not_awaited() + client.fail_activity_task.assert_not_awaited() + with pytest.raises(RuntimeError, match="unconfirmed remote callback stop.*registration remains active"): + await worker._shutdown() + assert worker._registered is True + client.deregister_worker_registration.assert_not_awaited() + finally: + Path(str(marker) + ".release").touch() + await wait_for_exit(pid) + shutil.rmtree(callback.directory, ignore_errors=True) + # The test performs external recovery after proving the worker cannot + # confirm stop. Release this synthetic state only after the process dies. + worker._remote_activity_processes.discard(callback) + + +async def test_registration_refuses_nonimportable_cooperative_handler_before_claiming(owner: Any) -> None: + _, client = owner + + @activity.defn(name="not-importable") + async def callback() -> None: + pass + + client.register_worker.reset_mock() + worker = Worker(client, task_queue="queue", worker_id="unsupported-owner", activities=[callback], + capabilities=["cooperative_cancellation"]) + with pytest.raises(RuntimeError, match="unsupported-owner.*not-importable.*spawn-compatible"): + await worker._register() + client.register_worker.assert_not_awaited() + client.poll_activity_task.assert_not_awaited() + + +async def test_client_observation_keeps_worker_credentials_namespace_and_fence(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + async with Client("http://runtime.test", control_token="control", worker_token="worker", namespace="ns") as client: + reply = httpx.Response(200, json=status(), request=httpx.Request("POST", "http://runtime.test")) + with patch.object(client._http, "request", new_callable=AsyncMock, return_value=reply) as send: + assert await client.activity_task_status( + task_id="remote-task", activity_attempt_id="attempt", lease_owner="owner", + ) == status() + headers = send.await_args.kwargs["headers"] + assert headers["Authorization"] == "Bearer worker" + assert headers["X-Namespace"] == "ns" + assert headers["X-Durable-Workflow-Protocol-Version"] == "1.20" + assert send.await_args.kwargs["json"] == {"activity_attempt_id": "attempt", "lease_owner": "owner"} + + +async def test_client_observation_requires_explicit_protocol_before_io(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", raising=False) + async with Client("http://runtime.test") as client: + with (patch.object(client._http, "request", new_callable=AsyncMock) as send, + pytest.raises(ValueError, match="1.20")): + await client.activity_task_status(task_id="task", activity_attempt_id="attempt", lease_owner="owner") + send.assert_not_awaited() + + +async def test_client_observation_bounds_the_total_retry_budget(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + retry = TransportRetryPolicy(initial_backoff_seconds=100, max_backoff_seconds=100, jitter=False) + async with Client("http://runtime.test", retry_policy=retry) as client: + refused = httpx.Response(503, json={"reason": "backend_unavailable"}, + request=httpx.Request("POST", "http://runtime.test")) + with patch.object(client._http, "request", new_callable=AsyncMock, return_value=refused) as send: + started = asyncio.get_running_loop().time() + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(client.activity_task_status( + task_id="task", activity_attempt_id="attempt", lease_owner="owner", + ), timeout=7) + assert asyncio.get_running_loop().time() - started < 6 + send.assert_awaited_once() diff --git a/tests/test_prepared_local_activity_client.py b/tests/test_prepared_local_activity_client.py new file mode 100644 index 0000000..b499da2 --- /dev/null +++ b/tests/test_prepared_local_activity_client.py @@ -0,0 +1,163 @@ +from __future__ import annotations + +import asyncio +import json +import time +from typing import Any +from unittest.mock import AsyncMock, patch + +import httpx +import pytest + +from durable_workflow.client import Client, _payload_completion_context +from durable_workflow.errors import ExternalPayloadError, ServerError +from durable_workflow.retry_policy import TransportRetryPolicy + + +@pytest.fixture(autouse=True) +def candidate_protocol(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + + +@pytest.mark.parametrize("operation", ["checkpoint", "checkpoint-group", "prepare", "recover", + "control", "heartbeat", "outcome", "acknowledge-cancellation"]) +async def test_operation_keeps_original_claim_and_encoded_backend_identity(operation: str) -> None: + attempt = "backend/attempt" if operation in { + "control", "heartbeat", "outcome", "acknowledge-cancellation", + } else None + async with Client("http://server", worker_token="worker", namespace="tenant") as client: + with patch.object(client, "_request", new_callable=AsyncMock, return_value={"opaque": "receipt"}) as send: + result = await client.prepared_local_activity_operation(task_id="task/one", lease_owner="original", + workflow_task_attempt=4, operation=operation, body={"checkpoint_id": "group"}, + activity_attempt_id=attempt, timeout_seconds=0.5) + assert result == {"opaque": "receipt"} + suffix = "backend%2Fattempt/" if attempt else "" + send.assert_awaited_once_with("POST", "/worker/workflow-tasks/task%2Fone/local-activities/" + suffix + operation, + worker=True, json={"lease_owner": "original", "workflow_task_attempt": 4, "checkpoint_id": "group"}, + timeout=0.5) + + +async def test_actual_transport_uses_worker_credential_namespace_and_candidate_header() -> None: + reply = httpx.Response(200, json={"active": False}, request=httpx.Request("POST", "http://server")) + async with Client("http://server", control_token="control", worker_token="worker", namespace="tenant") as client: + with patch.object(client._http, "request", new_callable=AsyncMock, return_value=reply) as send: + await client.prepared_local_activity_operation(task_id="task/one", lease_owner="original", + workflow_task_attempt=4, operation="control", activity_attempt_id="backend/attempt") + args, kwargs = send.await_args + assert args == ("POST", "/api/worker/workflow-tasks/task%2Fone/local-activities/backend%2Fattempt/control") + assert kwargs["headers"]["Authorization"] == "Bearer worker" + assert kwargs["headers"]["X-Namespace"] == "tenant" + assert kwargs["headers"]["X-Durable-Workflow-Protocol-Version"] == "1.20" + + +@pytest.mark.parametrize("change", [ + {"task_id": ""}, {"lease_owner": " "}, {"workflow_task_attempt": True}, {"workflow_task_attempt": 0}, + {"operation": "guess"}, {"activity_attempt_id": "unexpected"}, {"body": {"lease_owner": "replacement"}}, + {"body": {"workflow_task_attempt": 5}}, {"body": {"not_finite": float("nan")}}, + {"timeout_seconds": 0}, {"timeout_seconds": float("inf")}, {"timeout_seconds": 5.1}, +]) +async def test_invalid_claim_or_authority_never_starts_transport(change: dict[str, Any]) -> None: + arguments = {"task_id": "task", "lease_owner": "original", "workflow_task_attempt": 4, + "operation": "prepare", **change} + async with Client("http://server") as client: + with patch.object(client, "_request", new_callable=AsyncMock) as send: + with pytest.raises((ValueError, TypeError)): + await client.prepared_local_activity_operation(**arguments) + send.assert_not_awaited() + + +async def test_default_protocol_refuses_before_io(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION") + async with Client("http://server") as client: + with patch.object(client, "_request", new_callable=AsyncMock) as send: + with pytest.raises(ValueError, match="1.20"): + await client.prepared_local_activity_operation(task_id="task", lease_owner="original", + workflow_task_attempt=4, operation="prepare") + send.assert_not_awaited() + + +async def test_one_total_budget_cancels_pending_payload_and_retry_work() -> None: + cancelled = False + + async def pending(*args: Any, **kwargs: Any) -> None: + nonlocal cancelled + try: + await asyncio.sleep(60) + finally: + cancelled = True + + async with Client("http://server") as client: + with patch.object(client, "_request", side_effect=pending): + started = time.monotonic() + with pytest.raises(asyncio.TimeoutError): + await client.prepared_local_activity_operation(task_id="task", lease_owner="original", + workflow_task_attempt=4, operation="prepare", timeout_seconds=0.08) + assert time.monotonic() - started < 0.3 + assert cancelled + + +@pytest.mark.parametrize("receipt", [None, [], "unproved"]) +async def test_non_object_receipt_is_explicitly_refused(receipt: Any) -> None: + async with Client("http://server") as client: + with patch.object(client, "_request", new_callable=AsyncMock, return_value=receipt): + with pytest.raises(ServerError) as failure: + await client.prepared_local_activity_operation(task_id="task", lease_owner="original", + workflow_task_attempt=4, operation="prepare") + assert failure.value.reason() == "invalid_prepared_local_activity_receipt" + + +def test_payload_admission_identity_requires_negotiation_and_original_epoch() -> None: + body = {"lease_owner": "original", "workflow_task_attempt": 4, "checkpoint_id": "atomic-group"} + route = "/worker/workflow-tasks/task%2Fone/local-activities/checkpoint-group" + assert _payload_completion_context(route, body) is None + assert _payload_completion_context(route, body, allow_prepared=True) == { + "schema": "durable-workflow.v2.payload-completion-context.v2", "kind": "workflow", + "task_id": "task/one", "attempt": 4, "lease_owner": "original", + "operation": "local_activity_group_checkpoint", "checkpoint_id": "atomic-group", + } + for invalid in [ + {**body, "lease_owner": ""}, {**body, "workflow_task_attempt": True}, {**body, "checkpoint_id": ""}, + ]: + assert _payload_completion_context(route, invalid, allow_prepared=True) is None + + +@pytest.mark.parametrize("prepared_supported,state", [(True, "draining"), (False, "draining"), (True, "fenced")]) +async def test_group_draining_uploads_bind_each_payload_slot_and_original_claim( + prepared_supported: bool, state: str, +) -> None: + from durable_workflow import serializer + from tests.test_runtime_external_payload_transport import CompletionPayloadServer, runtime_client + + class PreparedPayloadServer(CompletionPayloadServer): + def cluster_info(self) -> dict[str, Any]: + info = super().cluster_info() + context = info["namespace"]["external_payload_storage"]["transport"]["upload"]["completion_context"] + if prepared_supported: + context["prepared_schema"] = "durable-workflow.v2.payload-completion-context.v2" + return info + + server = PreparedPayloadServer(state=state) + async with runtime_client(server, retry_policy=TransportRetryPolicy(max_attempts=1)) as client: + call = client.prepared_local_activity_operation(task_id="task/one", lease_owner="original", + workflow_task_attempt=4, operation="checkpoint-group", body={"checkpoint_id": "group", "commands": [ + {"type": "start_child_workflow", "input": serializer.envelope("child" * 100)}, + {"type": "prepare_local_activity", "arguments": serializer.envelope("local" * 100)}, + ]}) + if not prepared_supported or state == "fenced": + with pytest.raises((ServerError, ExternalPayloadError)): + await call + assert len(server.upload_requests) == 1 + assert "X-Durable-Workflow-Payload-Completion" not in server.upload_requests[0].headers + assert server.requests == [] + return + await call + assert len(server.upload_requests) == 4 + for index, slot in [(0, ["commands", 0, "input"]), (2, ["commands", 1, "arguments"])]: + first, bound = server.upload_requests[index:index + 2] + assert first.content == bound.content + assert bound.headers["authorization"] == first.headers["authorization"] + assert json.loads(bound.headers["X-Durable-Workflow-Payload-Completion"]) == { + "schema": "durable-workflow.v2.payload-completion-context.v2", "kind": "workflow", + "task_id": "task/one", "attempt": 4, "lease_owner": "original", + "operation": "local_activity_group_checkpoint", "checkpoint_id": "group", "slot": slot, + } diff --git a/tests/test_prepared_local_activity_groups.py b/tests/test_prepared_local_activity_groups.py new file mode 100644 index 0000000..3304295 --- /dev/null +++ b/tests/test_prepared_local_activity_groups.py @@ -0,0 +1,390 @@ +from __future__ import annotations + +import asyncio +import json +import os +from copy import deepcopy +from pathlib import Path +from typing import Any +from unittest.mock import AsyncMock + +import pytest + +from durable_workflow import activity, serializer, workflow +from durable_workflow._prepared_local_activity import PreparedCancellationObserved +from durable_workflow.client import Client +from durable_workflow.errors import NonDeterministicReplayError, ServerError, WorkflowCancelled +from durable_workflow.worker import Worker +from durable_workflow.workflow import LocalActivityExecutionAborted, replay +from tests.test_activity_process import blocked_without_python_progress, wait_for_exit, wait_for_file +from tests.test_cancellation_context import delivery, request, snapshot +from tests.test_prepared_local_activity_worker import admission, control, timestamp +from tests.test_replay_regression_corpus import PreparedLocalGroupColdResultsWorkflow +from tests.test_worker import compatible_cluster_info + + +@workflow.defn(name="prepared.two-locals") +class TwoLocals: + def run(self, ctx: workflow.WorkflowContext, marker: str, prefix: bool = False): # type: ignore[no-untyped-def] + if prefix: + yield ctx.upsert_memo({"before": "group"}) + return (yield [ctx.local_activity("prepared.group", [marker + ".0", marker + ".1"]), + ctx.local_activity("prepared.group", [marker + ".1", marker + ".0"])]) + + +@workflow.defn(name="prepared.mixed-group") +class MixedGroup: + def run(self, ctx: workflow.WorkflowContext): # type: ignore[no-untyped-def] + return (yield [ctx.start_child_workflow("python.child", [], task_queue="child"), + [ctx.schedule_activity("remote", []), ctx.local_activity("prepared.group", [])]]) + + +@workflow.defn(name="prepared.group-cleanup") +class GroupCleanup: + def run(self, ctx: workflow.WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled: + with ctx.cancellation_shield(): + yield [ctx.local_activity("prepared.group", []), ctx.local_activity("prepared.group", [])] + return "cleaned" + + +@activity.defn(name="prepared.group") +async def concurrent_callback(marker: str, peer: str) -> dict[str, Any]: + Path(marker).write_text(str(os.getpid())) + for _ in range(200): + if Path(peer).exists(): + return {"bytes": b"\x00\xff", "attempt": activity.context().info.activity_attempt_id} + await asyncio.sleep(0.01) + raise RuntimeError("parallel sibling did not start") + + +def capture(cls: type = TwoLocals, history: list[dict[str, Any]] | None = None, inputs: list[Any] | None = None): # type: ignore[no-untyped-def] + return replay(cls, history or [], ["unused"] if inputs is None else inputs, + prepare_local_activities=True, prepare_local_activity_groups=True, run_id="run-1", + local_activity_executor=lambda _: pytest.fail("callback ran without group admission")) + + +def local_event(kind: str, sequence: int, *, base: int = 1, **extra: Any) -> dict[str, Any]: + metadata = {"parallel_group_id": f"parallel-activities:{base}:2", "parallel_group_kind": "activity", + "parallel_group_base_sequence": base, "parallel_group_size": 2, + "parallel_group_index": sequence - base} + return {"id": kind + str(sequence), "event_type": kind, "payload": { + "sequence": sequence, "activity_type": "prepared.group", "execution_mode": "local", + "activity_execution_id": f"execution-{sequence}", "activity_attempt_id": f"attempt-{sequence}", + **metadata, "parallel_group_path": [metadata], **extra, + }} + + +def test_complete_nested_mixed_batch_is_captured_without_inventing_local_outcomes() -> None: + group = capture(MixedGroup, inputs=[]).prepared_local_activity_group + assert group is not None and not group.committed and group.base_sequence == 1 and group.size == 3 + assert [call.sequence for call in group.calls] == [3] + descriptor = group.calls[0].descriptor("avro") + assert "outcome" not in descriptor and "result" not in descriptor and "queue" not in descriptor + assert len(descriptor["parallel_group_path"]) == 2 + assert descriptor["parallel_group_path"][0]["parallel_group_index"] == 2 + assert descriptor["parallel_group_path"][1]["parallel_group_index"] == 1 + assert group.commands[0].task_queue == "child" # type: ignore[union-attr] + + +def test_group_prefix_is_checkpointed_separately_and_preserves_authored_sequences() -> None: + outcome = capture(inputs=["unused", True]) + assert [type(command).__name__ for command in outcome.commands] == ["UpsertMemo"] + assert outcome.prepared_local_activity_group is not None + assert outcome.prepared_local_activity_group.base_sequence == 2 + assert [call.sequence for call in outcome.prepared_local_activity_group.calls] == [2, 3] + + +@pytest.mark.parametrize("history", [ + [local_event("ActivityScheduled", 1)], + [local_event("ActivityStarted", 1), local_event("ActivityStarted", 2)], + [{**local_event("ActivityScheduled", 1), "payload": { + **local_event("ActivityScheduled", 1)["payload"], "parallel_group_path": []}}, + local_event("ActivityScheduled", 2)], + [local_event("ActivityScheduled", 1, base=2), local_event("ActivityScheduled", 2, base=2)], +]) +def test_partial_or_changed_group_history_cannot_authorize_any_callback(history: list[dict[str, Any]]) -> None: + with pytest.raises(NonDeterministicReplayError): + capture(history=history) + + +def test_partial_cold_replay_skips_a_completed_member_and_recovers_only_the_started_sibling() -> None: + history = [local_event("ActivityScheduled", 1), local_event("ActivityScheduled", 2), + local_event("ActivityCompleted", 1, result=serializer.envelope("first")), + local_event("ActivityStarted", 2)] + group = capture(history=history).prepared_local_activity_group + assert group is not None and group.committed and group.commands == () + assert [call.sequence for call in group.calls] == [2] and group.calls[0].recover is True + history.append(local_event("ActivityRetryScheduled", 2)) + assert capture(history=history).prepared_local_activity_group.calls[0].recover is False + history.append(local_event("ActivityCompleted", 2, result=serializer.envelope("second"))) + outcome = capture(history=history) + assert outcome.prepared_local_activity_group is None and outcome.commands[0].result == ["first", "second"] + + +def test_reverse_completion_order_replays_results_at_the_authored_positions() -> None: + history = [local_event("ActivityScheduled", 1), local_event("ActivityScheduled", 2), + local_event("ActivityCompleted", 2, result=serializer.envelope("second")), + local_event("ActivityCompleted", 1, result=serializer.envelope("first"))] + assert capture(history=history).commands[0].result == ["first", "second"] + + +def test_cleanup_group_preserves_original_root_delivery_and_deadline_for_every_member() -> None: + history = [request(), {**delivery(), "id": "canonical-delivery"}] + group = capture(GroupCleanup, history, []).prepared_local_activity_group + assert group is not None and group.base_sequence == 2 + for call in group.calls: + assert call.cleanup == {"request_id": "request-1", "root_request_id": "root-1", + "delivery_history_event_id": "canonical-delivery", + "cleanup_deadline_at": snapshot()["cleanup_deadline_at"]} + assert group.calls[0].cleanup == group.calls[1].cleanup + + +@pytest.mark.parametrize("selection", [False, True]) +def test_unsupported_group_shapes_are_refused_before_callbacks(selection: bool) -> None: + @workflow.defn(name="prepared.unsupported") + class Unsupported: + def run(self, ctx: workflow.WorkflowContext): # type: ignore[no-untyped-def] + local = ctx.local_activity("prepared.group", []) + yield ctx.select([local]) if selection else [local] * 101 + + with pytest.raises(LocalActivityExecutionAborted): + capture(Unsupported, inputs=[]) + + +class GroupServer: + def __init__(self, markers: tuple[Path, Path]) -> None: + self.client = AsyncMock(spec=Client) + self.history: list[dict[str, Any]] = [] + self.trace: list[tuple[str, Any]] = [] + self.attempts: dict[str, dict[str, Any]] = {} + self.markers = markers + self.stop = False + self.bad_checkpoint = False + self.bad_second_admission = False + self.partial_checkpoint_history = False + self.lost_outcome = False + self.admission_cancelled = False + self.omit_second_outcome = False + self.client.get_cluster_info.return_value = compatible_cluster_info(worker_protocol={ + "version": "1.20", "server_capabilities": { + "prepared_local_activities": True, "prepared_local_activity_groups": True, + "cooperative_cancellation": True, + }, + }) + self.client.register_worker.return_value = {"registered": True} + self.client.complete_workflow_task.return_value = {"outcome": "completed"} + self.client.prepared_local_activity_operation.side_effect = self.operation + self.client.workflow_task_history.side_effect = self.page + + async def page(self, **_: Any) -> dict[str, Any]: + self.trace.append(("history", None)) + return {"history_events": deepcopy(self.history), "next_history_page_token": None} + + async def operation(self, **kwargs: Any) -> dict[str, Any]: + assert kwargs["task_id"] == "task" and kwargs["lease_owner"] == "owner" + assert kwargs["workflow_task_attempt"] == 3 + name = kwargs["operation"] + body = kwargs["body"] + attempt_id = kwargs.get("activity_attempt_id") + self.trace.append((name, body.get("sequence", attempt_id))) + if name == "checkpoint-group": + assert all(not path.exists() for path in self.markers) + assert [command["type"] for command in body["commands"]] == ["prepare_local_activity"] * 2 + self.history = [local_event("ActivityScheduled", 1), local_event("ActivityScheduled", 2)] + if self.partial_checkpoint_history: + self.history.pop() + locals_ = [{"sequence": index, "activity_execution_id": f"execution-{index}"} for index in (1, 2)] + if self.bad_checkpoint: + locals_[1]["activity_execution_id"] = "execution-1" + return {"checkpointed": True, "duplicate": False, "reason": None, + "checkpoint_id": body["checkpoint_id"], "task_id": "task", "workflow_run_id": "run-1", + "workflow_task_attempt": 3, "lease_owner": "owner", "start_sequence": 1, "next_sequence": 3, + "local_activities": locals_, "history_refresh_page_token": "canonical-page"} + if name == "prepare": + assert all(not path.exists() for path in self.markers) + sequence = body["sequence"] + if sequence == 2 and self.admission_cancelled: + self.stop = True + raise ServerError(409, {"reason": "cancellation_requested"}) + receipt = admission( + activity_execution_id=f"execution-{sequence}", activity_attempt_id=f"attempt-{sequence}", + worker_attempt_id=body["worker_attempt_id"], + ) + if sequence == 2 and self.bad_second_admission: + receipt["activity_attempt_id"] = "" + self.attempts[f"attempt-{sequence}"] = receipt + self.history.append(local_event("ActivityStarted", sequence)) + return receipt + assert isinstance(attempt_id, str) and attempt_id in self.attempts + receipt = self.attempts[attempt_id] + sequence = int(attempt_id[-1]) + if name == "control": + assert body == {"renew_lease": True} + if not self.stop: + return control(receipt) + context = {**snapshot(), "requested_at": timestamp(), "cleanup_deadline_at": timestamp(30)} + # One immutable context for the entire group. + if not hasattr(self, "context"): + self.context = context + return control(receipt, active=False, renewed=False, stop_required=True, reason="cancellation_requested", + fenced=True, cancellation_request=self.context, cancellation_history_event_id="cancelled", + history_refresh_page_token="canonical-page") + if name == "acknowledge-cancellation": + assert body == {"request_id": "request-1"} + for marker in self.markers: + # This member must have physically exited before its own ACK. + if marker.exists() and marker == self.markers[sequence - 1]: + with pytest.raises(ProcessLookupError): + os.kill(int(marker.read_text()), 0) + return {"acknowledged": True, "duplicate": False, "reason": None, + "history_event_id": "joined-" + attempt_id} + assert name == "outcome" + assert len(self.attempts) == 2 + if self.lost_outcome: + raise TimeoutError("canonical acknowledgment lost") + if sequence != 2 or not self.omit_second_outcome: + self.history.append(local_event("ActivityCompleted", sequence, result=body["report"]["result"])) + return {**receipt, "recorded": True, "event_id": "ActivityCompleted" + str(sequence), + "event_type": "ActivityCompleted", "workflow_run_id": "run-1", "recorded_at": timestamp(), + "claim_released": False, "created_task_ids": [], "history_refresh_page_token": "canonical-page"} + + async def worker(self, monkeypatch: pytest.MonkeyPatch, *, handler: Any = concurrent_callback) -> Worker: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + worker = Worker(self.client, task_queue="queue", worker_id="owner", workflows=[TwoLocals], + capabilities=["cooperative_cancellation", "prepared_local_activities", + "prepared_local_activity_groups"]) + await worker._register() + worker.activities["prepared.group"] = handler + return worker + + def task(self) -> dict[str, Any]: + return {"task_id": "task", "workflow_type": "prepared.two-locals", "workflow_id": "workflow-1", + "run_id": "run-1", "workflow_task_attempt": 3, "payload_codec": "avro", + "arguments": serializer.envelope([str(self.markers[0])[:-2]]), "history_events": deepcopy(self.history)} + + +def server(tmp_path: Path) -> GroupServer: + return GroupServer((tmp_path / "callback.0", tmp_path / "callback.1")) + + +async def test_atomic_batch_then_complete_admission_precedes_concurrent_callbacks_and_canonical_results( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + fixture = server(tmp_path) + worker = await fixture.worker(monkeypatch) + commands = await worker._run_workflow_task(fixture.task()) + assert commands is not None and [command["type"] for command in commands] == ["complete_workflow"] + result = serializer.decode_envelope(commands[0]["result"]) + assert [item["attempt"] for item in result] == ["attempt-1", "attempt-2"] + assert all(item["bytes"] == b"\x00\xff" for item in result) + assert fixture.trace[:4] == [("checkpoint-group", None), ("history", None), ("prepare", 1), ("prepare", 2)] + for marker in fixture.markers: + await wait_for_exit(int(marker.read_text())) + assert not worker._prepared_local_activity_processes + fixture.client.fail_workflow_task.assert_not_awaited() + + +@pytest.mark.parametrize("failure", ["bad_checkpoint", "partial_checkpoint_history", "bad_second_admission"]) +async def test_incomplete_admission_never_spawns_or_publishes_application_work( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, failure: str, +) -> None: + fixture = server(tmp_path) + setattr(fixture, failure, True) + worker = await fixture.worker(monkeypatch) + assert await worker._run_workflow_task(fixture.task()) is None + assert all(not marker.exists() for marker in fixture.markers) + assert not any(name == "outcome" for name, _ in fixture.trace) + fixture.client.complete_workflow_task.assert_not_awaited() + fixture.client.fail_workflow_task.assert_not_awaited() + + +@pytest.mark.parametrize("failure", ["lost_outcome", "omit_second_outcome"]) +async def test_unknown_or_noncanonical_group_result_joins_siblings_and_never_reinvokes( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, failure: str, +) -> None: + fixture = server(tmp_path) + setattr(fixture, failure, True) + worker = await fixture.worker(monkeypatch) + assert await worker._run_workflow_task(fixture.task()) is None + assert sum(name == "prepare" for name, _ in fixture.trace) == 2 + for marker in fixture.markers: + await wait_for_exit(int(marker.read_text())) + assert not worker._prepared_local_activity_processes + fixture.client.complete_workflow_task.assert_not_awaited() + fixture.client.fail_workflow_task.assert_not_awaited() + + +async def test_cancellation_joins_both_blocked_callbacks_without_application_heartbeats_before_replay( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + fixture = server(tmp_path) + worker = await fixture.worker(monkeypatch, handler=blocked_without_python_progress) + task = fixture.task() + initial = capture(history=fixture.history, inputs=[str(tmp_path / "callback")]) + fixture.history = await worker._execute_prepared_local_activity_group(task, [], initial) + outcome = capture(history=fixture.history, inputs=[str(tmp_path / "callback")]) + # The blocked fixture accepts one marker, with no app heartbeat. + for call in outcome.prepared_local_activity_group.calls: + call.command.arguments = call.command.arguments[:1] + pending = asyncio.create_task(worker._execute_prepared_local_activity_group(task, fixture.history, outcome)) + try: + for marker in fixture.markers: + await wait_for_file(marker) + fixture.stop = True + with pytest.raises(PreparedCancellationObserved): + await asyncio.wait_for(pending, timeout=8) + for marker in fixture.markers: + await wait_for_exit(int(marker.read_text())) + assert sum(name == "acknowledge-cancellation" for name, _ in fixture.trace) == 2 + assert not any(name in {"heartbeat", "outcome"} for name, _ in fixture.trace) + assert not worker._prepared_local_activity_processes + finally: + pending.cancel() + await asyncio.gather(pending, return_exceptions=True) + + +async def test_cancel_during_second_admission_acks_unstarted_first_attempt_without_spawning( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + fixture = server(tmp_path) + worker = await fixture.worker(monkeypatch) + task = fixture.task() + fixture.history = await worker._execute_prepared_local_activity_group(task, [], capture()) + fixture.admission_cancelled = True + fixture.client.heartbeat_workflow_task.return_value = { + "task_id": "task", "lease_owner": "owner", "workflow_task_attempt": 3, "renewed": True, + "cancellation_request": {"request_id": "request-1", "requested_at": timestamp(), + "cleanup_deadline_at": timestamp(30), "history_refresh_page_token": "canonical-page"}, + } + with pytest.raises(LocalActivityExecutionAborted): + await worker._execute_prepared_local_activity_group(task, fixture.history, capture(history=fixture.history)) + assert all(not marker.exists() for marker in fixture.markers) + assert sum(name == "acknowledge-cancellation" for name, _ in fixture.trace) == 1 + + +async def test_group_capability_is_explicit_and_requires_the_installed_atomic_bridge( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + fixture = server(tmp_path) + worker = await fixture.worker(monkeypatch) + manifest = fixture.client.register_worker.await_args.kwargs["capability_manifest"] + assert manifest["prepared_local_activity_groups"]["implementation"] == "durable_atomic_all_admission" + capabilities = fixture.client.get_cluster_info.return_value["worker_protocol"]["server_capabilities"] + del capabilities["prepared_local_activity_groups"] + with pytest.raises(RuntimeError, match="installed atomic bridge"): + await worker._register() + + +def test_immutable_atomic_group_fixture_recovers_only_the_unfinished_original_sibling() -> None: + fixture = json.loads((Path(__file__).parent / "fixtures/replay_regressions/prepared-local-group-cold-results.json") + .read_text()) + outcome = capture(PreparedLocalGroupColdResultsWorkflow, fixture["history"][:-1], []) + group = outcome.prepared_local_activity_group + assert group is not None and group.committed and [call.sequence for call in group.calls] == [2] + assert group.calls[0].recover is True and group.commands == () + completed = capture(PreparedLocalGroupColdResultsWorkflow, fixture["history"], []) + assert completed.prepared_local_activity_group is None + assert completed.commands[0].result == ["fast-value", "fast-value"] diff --git a/tests/test_prepared_local_activity_policies.py b/tests/test_prepared_local_activity_policies.py new file mode 100644 index 0000000..2d9c9ff --- /dev/null +++ b/tests/test_prepared_local_activity_policies.py @@ -0,0 +1,206 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Any + +import pytest + +from durable_workflow import serializer, workflow +from durable_workflow.cancellation import CancellationPolicy, ParentClosePolicy +from durable_workflow.errors import NonDeterministicReplayError +from durable_workflow.worker import Worker +from durable_workflow.workflow import ( + CompleteUpdate, + LocalActivityExecutionAborted, + RecordLocalActivity, + apply_update, + query_state, + replay, + validate_update, +) +from tests.test_cancellation_context import delivery, request +from tests.test_prepared_local_activity_groups import local_event +from tests.test_prepared_local_activity_worker import PreparedServer, SequentialWorkflow +from tests.test_updates import _update_accepted_event + +POLICIES = ("try_cancel", "wait_cancellation_completed") + + +@workflow.defn(name="prepared.policy-group") +class PolicyGroupWorkflow: + def run(self, ctx: Any, policy: str = "wait_cancellation_completed") -> Any: + return (yield [ctx.local_activity("prepared.group", []), + ctx.local_activity("prepared.group", [], cancellation_policy=policy)]) + + +@workflow.defn(name="prepared.policy") +class PolicyWorkflow: + def run(self, ctx: Any, policy: Any = None, group: bool = False) -> Any: + call = ctx.local_activity("prepared.callback", [], cancellation_policy=policy) + self.result = None + self.result = yield [ctx.local_activity("prepared.callback", []), [call]] if group else call + return self.result + + @workflow.query(name="result") + def result_query(self) -> Any: + return self.result + + @workflow.update("inspect") + def inspect(self) -> Any: + return self.result + + @workflow.update_validator("inspect") + def validate_inspect(self) -> Any: + return self.result + + +def history(policy: str | None, *, complete: bool = False) -> list[dict[str, Any]]: + payload: dict[str, Any] = {"sequence": 1, "activity_type": "prepared.callback", "execution_mode": "local"} + if policy is not None: + payload["activity"] = {"cancellation_policy": policy} + events = [{"event_type": kind, "payload": dict(payload)} for kind in ("ActivityScheduled", "ActivityStarted")] + if complete: + events.append({"event_type": "ActivityCompleted", "payload": { + **payload, "result": serializer.envelope(b"completed"), + }}) + return events + + +@pytest.mark.parametrize("policy", [*POLICIES, CancellationPolicy.TRY_CANCEL, + CancellationPolicy.WAIT_CANCELLATION_COMPLETED]) +def test_explicit_policy_is_preserved_before_admission(policy: Any) -> None: + outcome = replay(PolicyWorkflow, [], [policy], prepare_local_activities=True, + local_activity_cancellation_policies=POLICIES, + local_activity_executor=lambda _: pytest.fail("callback ran before admission")) + call = outcome.prepared_local_activity + assert call is not None and not call.recover + assert call.descriptor("avro")["cancellation_policy"] == policy + assert "outcome" not in call.descriptor("avro") + + +@pytest.mark.parametrize("policy", [None, *POLICIES]) +def test_cold_replay_retains_policy_and_skips_completed_callback(policy: str | None) -> None: + options = {"prepare_local_activities": True, "local_activity_cancellation_policies": POLICIES} + pending = replay(PolicyWorkflow, history(policy), [policy], **options) + call = pending.prepared_local_activity + assert call is not None and call.recover + assert call.descriptor("avro").get("cancellation_policy") == policy + completed = replay(PolicyWorkflow, history(policy, complete=True), [policy], **options) + assert completed.prepared_local_activity is None + assert completed.commands[0].result == b"completed" # type: ignore[union-attr] + assert query_state(PolicyWorkflow, history(policy, complete=True), [policy], "result", **options) == b"completed" + + +@pytest.mark.parametrize("recorded,authored", [ + (None, "wait_cancellation_completed"), ("wait_cancellation_completed", None), + ("wait_cancellation_completed", "try_cancel"), ("try_cancel", "wait_cancellation_completed"), +]) +@pytest.mark.parametrize("complete", [False, True]) +def test_changed_policy_fails_pending_and_completed_replay(recorded: str | None, authored: str | None, + complete: bool) -> None: + with pytest.raises(NonDeterministicReplayError, match="local_activity_cancellation_policy_changed"): + replay(PolicyWorkflow, history(recorded, complete=complete), [authored], prepare_local_activities=True, + local_activity_cancellation_policies=POLICIES) + + +@pytest.mark.parametrize("delivered", [False, True]) +def test_changed_policy_is_rejected_before_request_or_committed_delivery(delivered: bool) -> None: + events = [*history("wait_cancellation_completed"), request()] + if delivered: + marker = delivery() + marker["payload"]["call_kind"] = "local_activity" + events.append(marker) + with pytest.raises(NonDeterministicReplayError, match="local_activity_cancellation_policy_changed"): + replay(PolicyWorkflow, events, ["try_cancel"], run_id="run-1", prepare_local_activities=True, + local_activity_cancellation_policies=POLICIES) + + +def test_atomic_group_recovery_preserves_policy_from_canonical_nested_snapshot() -> None: + events = [] + for sequence in (1, 2): + events.append(local_event("ActivityScheduled", sequence, activity={ + "cancellation_policy": "try_cancel" if sequence == 1 else "wait_cancellation_completed", + })) + events.append(local_event("ActivityStarted", sequence)) + options = {"prepare_local_activities": True, "prepare_local_activity_groups": True, + "local_activity_cancellation_policies": POLICIES} + outcome = replay(PolicyGroupWorkflow, events, [], **options) + group = outcome.prepared_local_activity_group + assert group is not None and group.committed and all(call.recover for call in group.calls) + assert group.calls[1].descriptor("avro")["cancellation_policy"] == "wait_cancellation_completed" + with pytest.raises(NonDeterministicReplayError, match="local_activity_cancellation_policy_changed"): + replay(PolicyGroupWorkflow, events, ["try_cancel"], **options) + + +@pytest.mark.parametrize("policy", ["abandon", CancellationPolicy.ABANDON, "unknown", 1, False, + ParentClosePolicy.ABANDON]) +def test_local_abandon_and_malformed_authoring_are_refused(policy: Any) -> None: + with pytest.raises(ValueError, match="local activity"): + RecordLocalActivity("prepared.callback", [], cancellation_policy=policy) + + +@pytest.mark.parametrize("prepared,group", [(False, False), (False, True), (True, False), (True, True)]) +def test_unsupported_policy_refuses_the_entire_call_before_callback_or_checkpoint(prepared: bool, group: bool) -> None: + with pytest.raises(LocalActivityExecutionAborted, match="installed Server discovery"): + replay(PolicyWorkflow, [], ["wait_cancellation_completed", group], prepare_local_activities=prepared, + prepare_local_activity_groups=prepared, + local_activity_executor=lambda _: pytest.fail("unsupported callback ran")) + with pytest.raises(LocalActivityExecutionAborted, match="explicit policies require prepared admission"): + RecordLocalActivity("prepared.callback", [], cancellation_policy="try_cancel").to_server_command("queue") + + +@pytest.mark.parametrize("advertised,accepted", [(None, ()), (True, ()), ("try_cancel", ()), + (["abandon", "unknown"], ()), + (["try_cancel"], ("try_cancel",)), (list(POLICIES), POLICIES)]) +async def test_worker_registers_only_discovered_prepared_policies(monkeypatch: pytest.MonkeyPatch, + advertised: Any, accepted: tuple[str, ...]) -> None: + server = PreparedServer() + capabilities = server.client.get_cluster_info.return_value["worker_protocol"]["server_capabilities"] + capabilities["prepared_local_activity_cancellation_policies"] = advertised + worker = await server.worker(monkeypatch) + assert worker._local_activity_cancellation_policies == accepted + registration = server.client.register_worker.await_args.kwargs + key = "prepared_local_activity_cancellation_policies" + assert (key in registration["capabilities"]) is bool(accepted) + assert (key in registration["capability_manifest"]) is bool(accepted) + if not accepted: + server.client.register_worker.reset_mock() + worker.capabilities = (*worker.capabilities, key) + with pytest.raises(RuntimeError, match="installed prepared policies"): + await worker._register() + server.client.register_worker.assert_not_awaited() + + +def test_explicit_worker_policy_capability_requires_prepared_admission() -> None: + with pytest.raises(ValueError, match="require prepared_local_activities"): + Worker(PreparedServer().client, task_queue="queue", + capabilities=["prepared_local_activity_cancellation_policies"]) + + +async def test_worker_refuses_undiscovered_policy_before_prefix_checkpoint_and_callback( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + server = PreparedServer() + worker = await server.worker(monkeypatch) + marker = tmp_path / "callback" + task = server.task(marker) + task["arguments"] = serializer.envelope([str(marker), True, "wait_cancellation_completed"]) + assert await worker._run_workflow_task(task) is None + assert not marker.exists() and server.trace == [] + server.client.prepared_local_activity_operation.assert_not_awaited() + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + assert worker.workflows["prepared.sequential"] is SequentialWorkflow + + +@pytest.mark.parametrize("complete", [False, True]) +def test_queries_updates_and_validators_replay_explicit_local_policy_without_running_callback(complete: bool) -> None: + policy = "wait_cancellation_completed" + events = history(policy, complete=complete) + options = {"prepare_local_activities": True, "local_activity_cancellation_policies": POLICIES} + expected = b"completed" if complete else None + assert query_state(PolicyWorkflow, events, [policy], "result", **options) == expected + assert validate_update(PolicyWorkflow, events, [policy], "inspect", [], **options) == expected + updated = apply_update(PolicyWorkflow, [*events, _update_accepted_event("inspect-1", "inspect", [])], + [policy], "inspect-1", **options) + assert isinstance(updated, CompleteUpdate) and updated.result == expected diff --git a/tests/test_prepared_local_activity_worker.py b/tests/test_prepared_local_activity_worker.py new file mode 100644 index 0000000..bd3db7f --- /dev/null +++ b/tests/test_prepared_local_activity_worker.py @@ -0,0 +1,526 @@ +from __future__ import annotations + +import asyncio +import json +import os +import shutil +import time +from copy import deepcopy +from datetime import datetime, timedelta, timezone +from pathlib import Path +from typing import Any +from unittest.mock import AsyncMock + +import pytest + +from durable_workflow import activity, serializer, workflow +from durable_workflow._prepared_local_activity import PreparedAttempt, PreparedCancellationObserved +from durable_workflow.client import Client +from durable_workflow.errors import WorkflowCancelled +from durable_workflow.worker import Worker +from durable_workflow.workflow import LocalActivityExecutionAborted, replay +from tests.test_activity_process import blocked_without_python_progress, wait_for_exit, wait_for_file, wait_for_release +from tests.test_cancellation_context import delivery, request, snapshot +from tests.test_replay_regression_corpus import PreparedLocalColdResultsWorkflow +from tests.test_worker import compatible_cluster_info + + +@workflow.defn(name="prepared.sequential") +class SequentialWorkflow: + def run(self, ctx: workflow.WorkflowContext, marker: str, prefix: bool = False, + policy: str | None = None): # type: ignore[no-untyped-def] + if prefix: + yield ctx.upsert_memo({"before": "local"}) + return (yield ctx.local_activity("prepared.callback", [marker], cancellation_policy=policy)) + + +@workflow.defn(name="prepared.cleanup") +class CleanupWorkflow: + def run(self, ctx: workflow.WorkflowContext, shield: bool = True): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled: + if shield: + with ctx.cancellation_shield(): + yield ctx.local_activity("prepared.callback", []) + else: + yield ctx.local_activity("prepared.callback", []) + return "cleaned" + + +@workflow.defn(name="prepared.group") +class GroupWorkflow: + def run(self, ctx: workflow.WorkflowContext): # type: ignore[no-untyped-def] + return (yield [ctx.local_activity("prepared.callback", []), ctx.start_timer(1)]) + + +@activity.defn(name="prepared.callback") +def typed_callback(marker: str) -> dict[str, Any]: + Path(marker).write_text(str(os.getpid())) + return {"value": b"\x00\xff", "attempt": activity.context().info.activity_attempt_id} + + +@activity.defn(name="prepared.progress") +async def progress_callback(marker: str) -> dict[str, Any]: + Path(marker).write_text(str(os.getpid())) + await activity.context().heartbeat({"phase": "processing", "count": 2}) + return {"value": b"\x00\xff", "attempt": activity.context().info.activity_attempt_id} + + +def timestamp(seconds: float = 0) -> str: + return (datetime.now(timezone.utc) + timedelta(seconds=seconds)).isoformat(timespec="microseconds").replace( + "+00:00", "Z", + ) + + +def admission(**changes: Any) -> dict[str, Any]: + return { + "prepared": True, "duplicate": False, "reason": None, + "workflow_task_id": "task", "workflow_task_attempt": 3, "lease_owner": "owner", + "worker_attempt_id": "nonce", "activity_execution_id": "execution", "activity_attempt_id": "attempt", + "attempt_number": 1, "server_time": timestamp(), "lease_expires_at": timestamp(30), + "start_to_close_deadline_at": None, "schedule_to_close_deadline_at": None, "heartbeat_deadline_at": None, + "cancellation_cleanup": None, **changes, + } + + +def admitted(receipt: dict[str, Any], **changes: Any) -> PreparedAttempt: + return PreparedAttempt.admitted( + receipt, task_id="task", run_id="run-1", owner="owner", epoch=3, nonce="nonce", + heartbeat_timeout=None, cleanup=None, request_started=time.monotonic(), **changes, + ) + + +def control(receipt: dict[str, Any], **changes: Any) -> dict[str, Any]: + return { + **receipt, "active": True, "renewed": True, "stop_required": False, + "workflow_lease_expires_at": receipt["lease_expires_at"], + "heartbeat_recorded": False, "heartbeat_history_event_id": None, **changes, + } + + +@pytest.mark.parametrize("changes", [ + {"prepared": False}, {"duplicate": 1}, {"reason": "refused"}, {"workflow_task_id": "other"}, + {"workflow_task_attempt": True}, {"workflow_task_attempt": 4}, {"lease_owner": "other"}, + {"worker_attempt_id": "other"}, {"attempt_number": True}, {"attempt_number": 0}, + {"activity_attempt_id": ""}, {"activity_execution_id": " "}, {"server_time": "2026-02-30T00:00:00Z"}, + {"lease_expires_at": "2026-01-01T00:00:00Z"}, {"heartbeat_deadline_at": "2027-01-01T00:00:00Z"}, + {"cancellation_cleanup": {"request_id": "invented"}}, +]) +def test_admission_refuses_changed_identity_and_invented_authority(changes: dict[str, Any]) -> None: + with pytest.raises(LocalActivityExecutionAborted): + admitted(admission(**changes)) + + +@pytest.mark.parametrize("field", ["reason", "start_to_close_deadline_at", "schedule_to_close_deadline_at", + "heartbeat_deadline_at"]) +def test_admission_requires_explicit_receipt_fields(field: str) -> None: + receipt = admission() + del receipt[field] + with pytest.raises(LocalActivityExecutionAborted): + admitted(receipt) + + +@pytest.mark.parametrize("changes", [ + {"active": 1}, {"renewed": False}, {"stop_required": True}, {"workflow_task_attempt": True}, + {"activity_attempt_id": "other"}, {"lease_owner": "other"}, {"reason": "refused"}, + {"heartbeat_recorded": True, "heartbeat_history_event_id": "invented-progress"}, + {"heartbeat_deadline_at": "2027-01-01T00:00:00Z"}, + {"start_to_close_deadline_at": "2027-01-01T00:00:00Z"}, + {"workflow_lease_expires_at": "2026-01-01T00:00:00Z"}, + {"cancellation_cleanup": {"request_id": "invented"}}, +]) +def test_supervisor_control_cannot_change_authority_or_report_application_progress(changes: dict[str, Any]) -> None: + receipt = admission() + attempt = admitted(receipt) + with pytest.raises(LocalActivityExecutionAborted): + attempt.validate_control(control(receipt, **changes)) + + +def test_application_heartbeat_alone_advances_its_deadline() -> None: + receipt = admission(heartbeat_deadline_at=timestamp(9)) + attempt = PreparedAttempt.admitted( + receipt, task_id="task", run_id="run-1", owner="owner", epoch=3, nonce="nonce", + heartbeat_timeout=10, cleanup=None, request_started=time.monotonic(), + ) + refreshed = control(receipt, renewed=False, heartbeat_recorded=True, heartbeat_history_event_id="heartbeat", + server_time=timestamp(1), heartbeat_deadline_at=timestamp(10)) + attempt.validate_control(refreshed, heartbeat=True) + assert attempt.deadlines["heartbeat_deadline_at"] == refreshed["heartbeat_deadline_at"] + attempt.validate_control(control(refreshed, renewed=True, heartbeat_recorded=False, + heartbeat_history_event_id=None)) + + +def test_transport_elapsed_time_cannot_restart_an_authority_budget() -> None: + receipt = admission(lease_expires_at=timestamp(2)) + with pytest.raises(LocalActivityExecutionAborted, match="budget expired"): + PreparedAttempt.admitted( + receipt, task_id="task", run_id="run-1", owner="owner", epoch=3, nonce="nonce", + heartbeat_timeout=None, cleanup=None, request_started=time.monotonic() - 3, + ) + + +@pytest.mark.parametrize("field", ["request_id", "root_request_id", "delivery_history_event_id", "cleanup_deadline_at"]) +def test_cleanup_admission_and_control_cannot_replace_original_cascade_authority(field: str) -> None: + cleanup = {"request_id": "request-1", "root_request_id": "root-1", "delivery_history_event_id": "delivery-1", + "cleanup_deadline_at": timestamp(20)} + receipt = admission(cancellation_cleanup=cleanup, heartbeat_deadline_at=cleanup["cleanup_deadline_at"], + start_to_close_deadline_at=cleanup["cleanup_deadline_at"], + schedule_to_close_deadline_at=cleanup["cleanup_deadline_at"]) + attempt = PreparedAttempt.admitted( + receipt, task_id="task", run_id="run-1", owner="owner", epoch=3, nonce="nonce", + heartbeat_timeout=None, cleanup=cleanup, request_started=time.monotonic(), + ) + changed = {**cleanup, field: timestamp(30) if field == "cleanup_deadline_at" else "replacement"} + with pytest.raises(LocalActivityExecutionAborted, match="cleanup authority"): + attempt.validate_control(control(receipt, cancellation_cleanup=changed)) + with pytest.raises(LocalActivityExecutionAborted, match="cleanup authority"): + PreparedAttempt.admitted( + {**receipt, "cancellation_cleanup": changed}, task_id="task", run_id="run-1", owner="owner", + epoch=3, nonce="nonce", heartbeat_timeout=None, cleanup=cleanup, request_started=time.monotonic(), + ) + + +def test_replay_captures_a_fresh_local_call_before_invoking_its_executor() -> None: + outcome = replay(SequentialWorkflow, [], ["unused", True], prepare_local_activities=True, + local_activity_executor=lambda _: pytest.fail("callback ran before admission")) + assert [type(command).__name__ for command in outcome.commands] == ["UpsertMemo"] + call = outcome.prepared_local_activity + assert call is not None and call.sequence == 2 and call.recover is False and call.cleanup is None + descriptor = call.descriptor("avro") + assert "outcome" not in descriptor and "result" not in descriptor + assert serializer.decode_envelope(descriptor["arguments"]) == ["unused"] + + +@pytest.mark.parametrize("last_event,recover", [("ActivityStarted", True), ("ActivityRetryScheduled", False)]) +def test_cold_replay_distinguishes_unknown_started_callback_from_durable_retry(last_event: str, recover: bool) -> None: + history = [{"event_type": "ActivityScheduled", "payload": { + "sequence": 1, "activity_type": "prepared.callback", "execution_mode": "local", + }}, {"event_type": last_event, "payload": { + "sequence": 1, "activity_type": "prepared.callback", "execution_mode": "local", + }}] + outcome = replay(SequentialWorkflow, history, ["unused"], prepare_local_activities=True) + assert outcome.commands == [] + assert outcome.prepared_local_activity is not None + assert outcome.prepared_local_activity.sequence == 1 and outcome.prepared_local_activity.recover is recover + + +def test_immutable_cold_results_skip_completed_callback_and_recover_unfinished_original_attempt() -> None: + fixture = json.loads(( + Path(__file__).parent / "fixtures/replay_regressions/prepared-local-cold-results.json" + ).read_text()) + history = fixture["history"] + outcome = replay(PreparedLocalColdResultsWorkflow, history[:-1], [], prepare_local_activities=True, + local_activity_executor=lambda _: pytest.fail("unfinished Started callback was reexecuted")) + call = outcome.prepared_local_activity + assert call is not None and call.sequence == 2 and call.recover is True + assert call.command.activity_type == "prepared.second" and outcome.commands == [] + completed = replay(PreparedLocalColdResultsWorkflow, history, [], prepare_local_activities=True, + local_activity_executor=lambda _: pytest.fail("completed callback was reexecuted")) + assert completed.prepared_local_activity is None + assert completed.commands[0].result == ["fast-value", "fast-value"] # type: ignore[union-attr] + + +def test_cleanup_capture_preserves_the_original_delivery_root_and_deadline() -> None: + marker = {**delivery(), "id": "canonical-delivery"} + original = [request(), marker] + outcome = replay(CleanupWorkflow, original, [], run_id="run-1", prepare_local_activities=True) + call = outcome.prepared_local_activity + assert call is not None and call.sequence == 2 + assert call.cleanup == { + "request_id": "request-1", "root_request_id": "root-1", "delivery_history_event_id": "canonical-delivery", + "cleanup_deadline_at": snapshot()["cleanup_deadline_at"], + } + assert call.descriptor("avro")["cancellation_cleanup"] == { + "request_id": "request-1", "delivery_history_event_id": "canonical-delivery", + } + original[1]["id"] = "" + with pytest.raises(LocalActivityExecutionAborted, match="canonical cancellation delivery"): + replay(CleanupWorkflow, original, [], run_id="run-1", prepare_local_activities=True) + with pytest.raises(LocalActivityExecutionAborted, match="shield"): + replay(CleanupWorkflow, [request(), marker], [False], run_id="run-1", prepare_local_activities=True) + + +def test_sequential_consumer_refuses_a_local_parallel_group_before_any_callback() -> None: + with pytest.raises(LocalActivityExecutionAborted, match="atomic group consumer"): + replay(GroupWorkflow, [], [], prepare_local_activities=True, + local_activity_executor=lambda _: pytest.fail("group callback ran before complete admission")) + + +class PreparedServer: + def __init__(self) -> None: + self.client = AsyncMock(spec=Client) + self.history: list[dict[str, Any]] = [] + self.trace: list[str] = [] + self.receipt: dict[str, Any] = {} + self.stop = False + self.bad_admission = False + self.omit_outcome_history = False + self.fail_outcome = False + self.marker: Path | None = None + self.client.get_cluster_info.return_value = compatible_cluster_info(worker_protocol={ + "version": "1.20", "server_capabilities": { + "prepared_local_activities": True, "cooperative_cancellation": True, + "workflow_memo_updates": {"supported": True}, + "supported_workflow_task_commands": ["upsert_memo"], + }, + }) + self.client.register_worker.return_value = {"registered": True} + self.client.prepared_local_activity_operation.side_effect = self.operation + self.client.workflow_task_history.side_effect = self.page + self.client.complete_workflow_task.return_value = {"outcome": "completed"} + + async def page(self, **kwargs: Any) -> dict[str, Any]: + assert kwargs == {"task_id": "task", "lease_owner": "owner", "workflow_task_attempt": 3, + "next_history_page_token": "canonical-page"} + self.trace.append("history") + return {"history_events": deepcopy(self.history), "next_history_page_token": None} + + async def operation(self, **kwargs: Any) -> dict[str, Any]: + assert kwargs["task_id"] == "task" and kwargs["lease_owner"] == "owner" + assert kwargs["workflow_task_attempt"] == 3 + name = kwargs["operation"] + self.trace.append(name) + body = kwargs["body"] + if name == "prepare": + assert "outcome" not in body["descriptor"] + self.receipt = admission(worker_attempt_id=body["worker_attempt_id"]) + if self.bad_admission: + self.receipt["activity_attempt_id"] = "" + payload = {"sequence": body["sequence"], "activity_type": "prepared.callback", "execution_mode": "local", + "activity_execution_id": "execution", "activity_attempt_id": "attempt"} + if "cancellation_policy" in body["descriptor"]: + payload["activity"] = {"cancellation_policy": body["descriptor"]["cancellation_policy"]} + self.history.extend([{"id": "scheduled", "event_type": "ActivityScheduled", "payload": payload}, + {"id": "started", "event_type": "ActivityStarted", "payload": payload}]) + return self.receipt + assert kwargs["activity_attempt_id"] == "attempt" + if name == "control": + assert body == {"renew_lease": True} + if not self.stop: + return control(self.receipt) + context = snapshot() + context.update(requested_at=timestamp(), cleanup_deadline_at=timestamp(30)) + return control(self.receipt, active=False, stop_required=True, renewed=False, + reason="cancellation_requested", fenced=True, cancellation_request=context, + cancellation_history_event_id="cancelled", history_refresh_page_token="canonical-page") + if name == "acknowledge-cancellation": + assert body == {"request_id": "request-1"} + assert self.marker is not None + pid = int(self.marker.read_text()) + with pytest.raises(ProcessLookupError): + os.kill(pid, 0) + return {"acknowledged": True, "duplicate": False, "reason": None, "history_event_id": "joined"} + if name == "heartbeat": + assert body == {"progress": {"details": {"phase": "processing", "count": 2}}} + payload = {**self.history[-1]["payload"], "progress": body["progress"]} + self.history.append({"id": "progress", "event_type": "ActivityHeartbeatRecorded", "payload": payload}) + return control(self.receipt, renewed=False, heartbeat_recorded=True, + heartbeat_history_event_id="progress") + assert name == "outcome" + if self.fail_outcome: + raise TimeoutError("outcome acknowledgment lost") + if not self.omit_outcome_history: + payload = {**self.history[-1]["payload"], "result": body["report"]["result"]} + self.history.append({"id": "completed", "event_type": "ActivityCompleted", "payload": payload}) + return {**self.receipt, "recorded": True, "event_id": "completed", "event_type": "ActivityCompleted", + "workflow_run_id": "run-1", "recorded_at": timestamp(), "claim_released": False, + "created_task_ids": [], "history_refresh_page_token": "canonical-page"} + + async def worker(self, monkeypatch: pytest.MonkeyPatch, *, handler: Any = typed_callback) -> Worker: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + worker = Worker(self.client, task_queue="queue", worker_id="owner", workflows=[SequentialWorkflow], + capabilities=["cooperative_cancellation", "prepared_local_activities"]) + await worker._register() + worker.activities["prepared.callback"] = handler + return worker + + def task(self, marker: Path) -> dict[str, Any]: + self.marker = marker + return {"task_id": "task", "workflow_type": "prepared.sequential", "workflow_id": "workflow-1", + "run_id": "run-1", "workflow_task_attempt": 3, "payload_codec": "avro", + "arguments": serializer.envelope([str(marker)]), "history_events": deepcopy(self.history)} + + +@pytest.mark.parametrize("policy", [None, "try_cancel", "wait_cancellation_completed"]) +async def test_worker_runs_only_an_admitted_process_then_replays_canonical_result( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, policy: str | None, +) -> None: + server = PreparedServer() + server.client.get_cluster_info.return_value["worker_protocol"]["server_capabilities"].update( + prepared_local_activity_cancellation_policies=["try_cancel", "wait_cancellation_completed"], + ) + worker = await server.worker(monkeypatch) + task = server.task(tmp_path / "callback") + task["arguments"] = serializer.envelope([str(tmp_path / "callback"), False, policy]) + commands = await worker._run_workflow_task(task) + assert commands is not None and [command["type"] for command in commands] == ["complete_workflow"] + assert serializer.decode_envelope(commands[0]["result"]) == {"value": b"\x00\xff", "attempt": "attempt"} + assert server.trace == ["prepare", "control", "control", "outcome", "history"] + assert server.history[0]["payload"].get("activity", {}).get("cancellation_policy") == policy + assert not worker._prepared_local_activity_processes + server.client.fail_workflow_task.assert_not_awaited() + await wait_for_exit(int((tmp_path / "callback").read_text())) + + +async def test_bad_admission_never_spawns_and_never_completes_or_fails_a_claim( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + server = PreparedServer() + server.bad_admission = True + worker = await server.worker(monkeypatch) + assert await worker._run_workflow_task(server.task(tmp_path / "callback")) is None + assert not (tmp_path / "callback").exists() + assert server.trace == ["prepare"] + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + + +async def test_supervised_application_heartbeat_preserves_progress_in_canonical_history( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + server = PreparedServer() + worker = await server.worker(monkeypatch, handler=progress_callback) + marker = tmp_path / "callback" + commands = await worker._run_workflow_task(server.task(marker)) + assert commands is not None and commands[0]["type"] == "complete_workflow" + progress = [event for event in server.history if event["event_type"] == "ActivityHeartbeatRecorded"] + assert len(progress) == 1 and progress[0]["payload"]["progress"]["details"] == {"phase": "processing", "count": 2} + assert server.trace.count("heartbeat") == 1 and server.trace.count("outcome") == 1 + await wait_for_exit(int(marker.read_text())) + server.client.fail_workflow_task.assert_not_awaited() + + +@pytest.mark.parametrize("failure", ["fail_outcome", "omit_outcome_history"]) +async def test_lost_or_noncanonical_outcome_abandons_without_reexecuting_callback( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, failure: str, +) -> None: + server = PreparedServer() + setattr(server, failure, True) + worker = await server.worker(monkeypatch) + assert await worker._run_workflow_task(server.task(tmp_path / "callback")) is None + assert server.trace.count("prepare") == 1 and server.trace.count("outcome") == 1 + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + assert not worker._prepared_local_activity_processes + + +@pytest.mark.parametrize("policy", [None, "try_cancel", "wait_cancellation_completed"]) +async def test_control_stops_and_joins_a_gil_blocked_local_callback_without_application_heartbeats( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, policy: str | None, +) -> None: + server = PreparedServer() + server.client.get_cluster_info.return_value["worker_protocol"]["server_capabilities"].update( + prepared_local_activity_cancellation_policies=["try_cancel", "wait_cancellation_completed"], + ) + worker = await server.worker(monkeypatch, handler=blocked_without_python_progress) + marker = tmp_path / "callback" + task = server.task(marker) + outcome = replay(SequentialWorkflow, [], [str(marker), False, policy], prepare_local_activities=True, + local_activity_cancellation_policies=worker._local_activity_cancellation_policies) + pending = asyncio.create_task(worker._execute_prepared_local_activity(task, [], outcome)) + try: + await wait_for_file(marker) + pid = int(marker.read_text()) + started = time.monotonic() + server.stop = True + with pytest.raises(PreparedCancellationObserved): + await asyncio.wait_for(pending, timeout=7.0) + assert time.monotonic() - started < 4.0 + await wait_for_exit(pid) + assert server.trace.count("acknowledge-cancellation") == 1 + assert "heartbeat" not in server.trace and "outcome" not in server.trace + assert task["_prepared_cancellation_context"]["request_id"] == "request-1" + assert not worker._prepared_local_activity_processes + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + finally: + if not pending.done(): + pending.cancel() + await asyncio.gather(pending, return_exceptions=True) + + +async def test_prepared_capability_requires_actual_bridge_and_never_advertises_group_support( + monkeypatch: pytest.MonkeyPatch, +) -> None: + server = PreparedServer() + worker = await server.worker(monkeypatch) + manifest = server.client.register_worker.await_args.kwargs["capability_manifest"] + assert manifest["prepared_local_activities"]["implementation"] == "durable_sequential_admission" + assert "prepared_local_activity_groups" not in manifest + info = server.client.get_cluster_info.return_value + del info["worker_protocol"]["server_capabilities"]["prepared_local_activities"] + with pytest.raises(RuntimeError, match="installed admission bridge"): + await worker._register() + with pytest.raises(ValueError, match="require prepared_local_activities"): + Worker(server.client, task_queue="queue", capabilities=["prepared_local_activity_groups"]) + + +async def test_descendant_stop_keeps_local_observation_time_and_canonical_root_time_distinct( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + server = PreparedServer() + worker = await server.worker(monkeypatch) + task = server.task(tmp_path / "unused") + root = snapshot() + observed = {"request_id": root["request_id"], "requested_at": "2026-10-01T00:00:05Z", + "cleanup_deadline_at": root["cleanup_deadline_at"], "history_refresh_page_token": "canonical-page"} + server.client.heartbeat_workflow_task.return_value = { + "task_id": "task", "lease_owner": "owner", "workflow_task_attempt": 3, "renewed": True, + "cancellation_request": observed, + } + + async def stopped(*_: Any) -> list[dict[str, Any]]: + task["_prepared_cancellation_context"] = root + server.history = [request()] + raise PreparedCancellationObserved("original callback joined") + + async def delivered(**kwargs: Any) -> dict[str, Any]: + assert kwargs["call_kind"] == "local_activity" + server.history.append({**delivery(), "payload": {**delivery()["payload"], "call_kind": "local_activity"}}) + return {"delivered": True} + + worker._execute_prepared_local_activity = stopped # type: ignore[method-assign] + server.client.deliver_workflow_cancellation.side_effect = delivered + outcome, _ = await worker._replay_workflow_claim( + SequentialWorkflow, task, [], ["unused"], payload_codec="avro", + execute_local=lambda _: pytest.fail("callback reexecuted after joined stop"), + ) + assert outcome.prepared_local_activity is None and len(outcome.commands) == 1 + assert task["cancellation_request"] == observed + assert task["cancellation_request"]["requested_at"] != root["requested_at"] + assert server.client.deliver_workflow_cancellation.await_count == 1 + + +async def test_lost_prepared_supervisor_retains_workflow_capacity_and_never_reports_stop( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + server = PreparedServer() + worker = await server.worker(monkeypatch, handler=wait_for_release) + worker.max_concurrent_workflow_tasks = 1 + worker._wf_semaphore = asyncio.Semaphore(1) + marker = tmp_path / "callback" + await worker._reserve_workflow_capacity() + pending = worker._admit_workflow_work(server.task(marker), "workflow") + await wait_for_file(marker) + callback = next(iter(worker._prepared_local_activity_processes)) + pid = int(marker.read_text()) + try: + callback._supervisor.kill() + assert await asyncio.wait_for(pending, timeout=4.0) is None + os.kill(pid, 0) + assert callback.stopped is False + assert worker._current_task_slots()["workflow_available"] == 0 and worker._workflow_reserved == 1 + assert "acknowledge-cancellation" not in server.trace and "outcome" not in server.trace + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + with pytest.raises(RuntimeError, match="unconfirmed prepared local callback stop.*registration remains active"): + await worker._shutdown() + server.client.deregister_worker_registration.assert_not_awaited() + finally: + Path(str(marker) + ".release").touch() + await wait_for_exit(pid) + shutil.rmtree(callback.directory, ignore_errors=True) + worker._prepared_local_activity_processes.discard(callback) + worker._abandoned_prepared_local_claims.discard("task") + worker._release_workflow_capacity() diff --git a/tests/test_replay_regression_corpus.py b/tests/test_replay_regression_corpus.py index 52d1af0..6e39eef 100644 --- a/tests/test_replay_regression_corpus.py +++ b/tests/test_replay_regression_corpus.py @@ -10,8 +10,13 @@ from durable_workflow import Replayer, Worker, serializer, workflow from durable_workflow.client import Client, WorkflowStreamAppendItem -from durable_workflow.errors import NonDeterministicReplayError, WorkflowPayloadDecodeError -from durable_workflow.workflow import WorkflowContext, commands_to_server_commands, query_state +from durable_workflow.errors import NonDeterministicReplayError, WorkflowCancelled, WorkflowPayloadDecodeError +from durable_workflow.workflow import ( + LocalActivityExecutionAborted, + WorkflowContext, + commands_to_server_commands, + query_state, +) from tests.test_golden_history_replay import ( GoldenSagaCompensationWorkflow, GoldenSignalWaitWorkflow, @@ -212,14 +217,50 @@ def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] return (yield ctx.local_activity("golden.local", [])) +@workflow.defn(name="tests.replay.prepared-local-cold-results") +class PreparedLocalColdResultsWorkflow: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + first = yield ctx.local_activity("prepared.first", []) + second = yield ctx.local_activity("prepared.second", []) + return [first, second] + + +@workflow.defn(name="tests.replay.prepared-local-group-cold-results") +class PreparedLocalGroupColdResultsWorkflow: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + return (yield [ctx.local_activity("prepared.first", []), ctx.local_activity("prepared.second", [])]) + + +@workflow.defn(name="tests.replay.cooperative-reopened-condition-cleanup") +class CooperativeReopenedConditionCleanupWorkflow: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.wait_condition(lambda: False, key="forward-wait") + except WorkflowCancelled as exc: + with ctx.cancellation_shield(): + yield ctx.start_timer(1) + return exc.request_id + return "not cancelled" + + +@workflow.defn(name="tests.replay.child-policy-author") +class ChildPolicyAuthorWorkflow: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + return (yield ctx.start_child_workflow("child", [])) + + WORKFLOWS = [ + ChildPolicyAuthorWorkflow, ColdReplacementSatisfiedConditionWorkflow, + CooperativeReopenedConditionCleanupWorkflow, GoldenSagaCompensationWorkflow, GoldenSignalWaitWorkflow, GoldenSingleActivityWorkflow, GoldenTimeoutWaitWorkflow, GoldenVersionMarkerWorkflow, LocalActivityColdResultWorkflow, + PreparedLocalColdResultsWorkflow, + PreparedLocalGroupColdResultsWorkflow, MessageStreamConsumerWorkflow, NestedParallelPathWorkflow, ParallelMetadataProducerWorkflow, @@ -380,12 +421,16 @@ def test_checked_in_replay_regression_corpus_uses_official_replayer( expected_error = expected.get("error") expected_replay_error = fixture.get("expected_replay_error") if isinstance(expected_replay_error, dict): - assert expected_replay_error.get("type") == "NonDeterministicReplayError" message = expected_replay_error.get("message_contains") - workflow_sequence = expected_replay_error.get("workflow_sequence") assert isinstance(message, str) and message - assert isinstance(workflow_sequence, int) assert expected.get("command_sequence") == [] + if expected_replay_error.get("type") == "LocalActivityExecutionAborted": + with pytest.raises(LocalActivityExecutionAborted, match=message): + _execute_fixture(fixture) + return + assert expected_replay_error.get("type") == "NonDeterministicReplayError" + workflow_sequence = expected_replay_error.get("workflow_sequence") + assert isinstance(workflow_sequence, int) with pytest.raises(NonDeterministicReplayError, match=message) as captured: _execute_fixture(fixture) assert captured.value.workflow_sequence == workflow_sequence diff --git a/tests/test_scoped_run_cancellation_context.py b/tests/test_scoped_run_cancellation_context.py new file mode 100644 index 0000000..76b2e82 --- /dev/null +++ b/tests/test_scoped_run_cancellation_context.py @@ -0,0 +1,201 @@ +from __future__ import annotations + +import json +from copy import deepcopy +from dataclasses import FrozenInstanceError +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +import pytest + +from durable_workflow import CancellationContext, ScopedCancellationContext, serializer +from durable_workflow._cooperative_cancellation import read_cancellation_history +from durable_workflow.errors import NonDeterministicReplayError, WorkflowCancelled +from durable_workflow.workflow import CompleteWorkflow, ScheduleActivity, WorkflowContext, replay +from tests.test_cooperative_cancellation import completed_activity + + +def fixtures() -> dict[str, Any]: + return json.loads((Path(__file__).parent / "fixtures/scoped-run-cancellation-context.json").read_text()) + + +def history() -> list[dict[str, Any]]: + snapshot = fixtures()["child"] + return [ + {"event_type": "CooperativeCancellationRequested", "timestamp": "2026-10-04T00:00:05Z", "payload": { + "workflow_run_id": "child-run", "workflow_instance_id": "child-instance", + "workflow_command_id": "child-request", "reason": "maintenance", + "cleanup_deadline_at": snapshot["cleanup_deadline_at"], "cancellation": deepcopy(snapshot), + }}, + {"event_type": "CooperativeCancellationDelivered", "timestamp": "2026-10-04T00:00:08Z", "payload": { + "workflow_run_id": "child-run", "workflow_command_id": "child-request", + "sequence": 1, "call_kind": "timer", "cancellation": deepcopy(snapshot), + }}, + ] + + +@pytest.mark.parametrize("name", ["child", "grandchild"]) +def test_native_context_preserves_the_complete_tree_and_original_metadata(name: str) -> None: + snapshot = fixtures()[name] + context = CancellationContext.from_dict(snapshot) + assert context.to_dict() == snapshot + assert CancellationContext.from_dict(context.to_dict()) == context + assert context.root_request_id == "root-request" + assert context.reason == "maintenance" + assert context.source == "api" + assert context.requester == {"type": "operator", "id": "operator-1"} + assert context.requested_at == datetime(2026, 10, 4, 0, 0, 0, 123456, timezone.utc) + assert context.scope_origin is not None + assert context.scope_origin.root_deadline == datetime(2026, 10, 4, 0, 0, 30, 123456, timezone.utc) + if name == "child": + assert context.parent_request_id == "inner-request" + assert [entry.scope_id for entry in context.scope_origin.lineage] == ["outer", "inner"] + assert context.scope_origin.deadline == datetime(2026, 10, 4, 0, 0, 20, 123456, timezone.utc) + assert context.deadline == datetime(2026, 10, 4, 0, 0, 15, 123456, timezone.utc) + else: + assert context.parent_request_id == "child-scope-request" + assert [entry.scope_id for entry in context.scope_origin.lineage] == ["outer", "inner", "root", "child-scope"] + assert [entry.request_id for entry in context.lineage] == [ + "root-request", "child-scope-request", "grandchild-request", + ] + assert context.deadline == datetime(2026, 10, 4, 0, 0, 12, 123456, timezone.utc) + + +def test_avro_timezone_and_object_order_do_not_change_original_scope_metadata() -> None: + original = fixtures()["grandchild"] + snapshot = deepcopy(original) + snapshot["requested_at"] = "2026-10-03T20:00:00.123456-04:00" + snapshot["cleanup_deadline_at"] = snapshot["scope_authority_deadline_at"] = "2026-10-03T20:00:12.123456-04:00" + snapshot["requester"] = dict(reversed(list(snapshot["requester"].items()))) + snapshot["scope_origin"]["lineage"] = [ + dict(reversed(list(entry.items()))) for entry in snapshot["scope_origin"]["lineage"] + ] + decoded = serializer.decode(serializer.encode(snapshot), codec="avro") + assert CancellationContext.from_dict(decoded).to_dict() == original + + +def test_origin_and_nested_scope_addresses_are_immutable() -> None: + snapshot = fixtures()["grandchild"] + context = CancellationContext.from_dict(snapshot) + snapshot["scope_origin"]["lineage"][3]["scope_id"] = "changed" + origin = context.scope_origin + assert isinstance(origin, ScopedCancellationContext) + assert origin.scope_id == "child-scope" and origin.root_scope_id == "outer" + assert origin.request_id == "child-scope-request" and origin.parent_request_id == "child-request" + assert origin.workflow_instance_id == "child-instance" and origin.workflow_run_id == "child-run" + assert origin.requested_at == context.requested_at + with pytest.raises(FrozenInstanceError): + origin.lineage[3].scope_id = "changed" # type: ignore[misc] + with pytest.raises(TypeError): + origin.root_context.requester["id"] = "changed" # type: ignore[index] + detached = origin.to_dict() + detached["lineage"][3]["scope_id"] = "changed" + assert origin.to_dict() == fixtures()["grandchild"]["scope_origin"] + + +def test_cold_replay_keeps_original_origin_and_the_same_narrowed_cleanup_clock() -> None: + contexts: list[CancellationContext] = [] + observations: list[list[float]] = [] + + class Cleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + assert ctx.cancellation_context is None + try: + yield ctx.start_timer(10) + except WorkflowCancelled as error: + assert error.context is not None and error.context is ctx.cancellation_context + assert error.context.to_dict() == fixtures()["child"] + contexts.append(error.context) + seen = [error.context.remaining()] + observations.append(seen) + with ctx.cancellation_shield(): + yield ctx.schedule_activity("cleanup", []) + seen.append(error.context.remaining()) + return {"remaining": seen, "cancellation": error.context.to_dict()} + return None + + events = history() + pending = replay(Cleanup, events, [], run_id="child-run") + assert len(pending.commands) == 1 and isinstance(pending.commands[0], ScheduleActivity) + assert observations == [[7.123456]] + events.extend([ + {**completed_activity(2, "cleanup", "cleaned"), "timestamp": "2026-10-04T00:00:12Z"}, + {"event_type": "SignalReceived", "timestamp": "2026-10-04T00:00:29Z", "payload": { + "signal_name": "future", "arguments": serializer.envelope([]), + }}, + ]) + for _replacement in range(2): + assert replay(Cleanup, events, [], run_id="child-run").commands == [CompleteWorkflow({ + "remaining": [7.123456, 3.123456], "cancellation": fixtures()["child"], + })] + for context in contexts: + with pytest.raises(RuntimeError, match="active workflow replay"): + context.remaining() + + +def test_canonical_delivery_cannot_replace_a_committed_scope_origin() -> None: + events = history() + events[1]["payload"]["cancellation"]["scope_origin"]["lineage"][1]["scope_id"] = "different" + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history(events, run_id="child-run") + + +def invalid_contexts() -> list[Any]: + original = fixtures()["grandchild"] + rows = [] + for key, value in { + "schema": "durable-workflow.cancellation-context/v1", "root_request_id": "other", + "root_workflow_instance_id": "other", "root_workflow_run_id": "other", + "parent_request_id": "root-request", "reason": "other", "source": "other", + "requested_at": "2026-10-04T00:00:01.123456Z", "scope_origin": [], + "cleanup_deadline_at": "2026-10-04T00:00:15.123456Z", + "scope_authority_deadline_at": "2026-10-04T00:00:15.123456Z", + }.items(): + rows.append(pytest.param({**original, key: value}, id=key)) + for key in ["scope_origin", "scope_authority_deadline_at", "parent_request_id"]: + snapshot = deepcopy(original) + del snapshot[key] + rows.append(pytest.param(snapshot, id="missing "+key)) + for name in ["requester", "discarded hop", "widened global budget", "reused request", "reused run", + "discarded root", "widened scope budget", "run reassigned", "run reentry", "repeated scope request", + "repeated address", "recursive root", "unsupported scope field", "invalid schema type"]: + snapshot = deepcopy(original) + if name == "requester": + snapshot["requester"]["id"] = "other" + elif name == "discarded hop": + snapshot["lineage"][1]["request_id"] = "child-request" + elif name == "widened global budget": + snapshot["cleanup_deadline_at"] = snapshot["scope_authority_deadline_at"] = "2026-10-04T00:00:30.123456Z" + elif name == "reused request": + snapshot["request_id"] = snapshot["lineage"][2]["request_id"] = "inner-request" + elif name == "reused run": + snapshot["lineage"][2]["workflow_run_id"] = "child-run" + elif name == "discarded root": + snapshot["scope_origin"]["lineage"].pop(0) + elif name == "widened scope budget": + snapshot["scope_origin"]["lineage"][3]["cleanup_deadline_at"] = "2026-10-04T00:00:16.123456Z" + elif name == "run reassigned": + snapshot["scope_origin"]["lineage"][3]["workflow_instance_id"] = "other" + elif name == "run reentry": + snapshot["scope_origin"]["lineage"][3]["workflow_run_id"] = "root-run" + snapshot["scope_origin"]["lineage"][3]["workflow_instance_id"] = "root-instance" + elif name == "repeated scope request": + snapshot["scope_origin"]["lineage"][3]["request_id"] = "inner-request" + elif name == "repeated address": + snapshot["scope_origin"]["lineage"][3]["scope_id"] = "root" + elif name == "recursive root": + snapshot["scope_origin"]["root_context"]["schema"] = "durable-workflow.cancellation-context/v2" + snapshot["scope_origin"]["root_context"]["scope_origin"] = original["scope_origin"] + elif name == "unsupported scope field": + snapshot["scope_origin"]["lineage"][3]["authority"] = "unrecorded" + else: + snapshot["schema"] = [] + rows.append(pytest.param(snapshot, id=name)) + return rows + + +@pytest.mark.parametrize("snapshot", invalid_contexts()) +def test_invalid_origin_identity_or_budget_is_refused(snapshot: dict[str, Any]) -> None: + with pytest.raises(ValueError): + CancellationContext.from_dict(snapshot)