diff --git a/.cargo/config.toml b/.cargo/config.toml index 7f5daf4df..33f8e44c8 100644 --- a/.cargo/config.toml +++ b/.cargo/config.toml @@ -1,6 +1,8 @@ [env] # This temporarily overrides the version of the CLI used for integration tests, locally and in CI -# CLI_VERSION_OVERRIDE = "v1.6.3-serverless" +# TEMP: workflow task completion pagination requires server >=v1.32.0-162.0, for which there's no +# published stable CLI release yet. +CLI_VERSION_OVERRIDE = "v1.8.3-server-1.32.0-162.0" [alias] # Not sure why --all-features doesn't work @@ -13,6 +15,8 @@ integ-test = [ "--features", "ephemeral-server", "--features", + "temporalio-sdk/experimental", + "--features", "temporalio-sdk-core/otel", "--package", "temporalio-sdk-core", diff --git a/.github/workflows/changelog.yml b/.github/workflows/changelog.yml new file mode 100644 index 000000000..e5554a225 --- /dev/null +++ b/.github/workflows/changelog.yml @@ -0,0 +1,14 @@ +name: Changelog + +on: + pull_request: + types: [opened, synchronize, reopened, labeled, unlabeled] + +permissions: + contents: read + pull-requests: read + +jobs: + changelog-check: + name: Changelog checkpoint + uses: temporalio/.github/.github/workflows/changelog.yml@d8129c755fbd1d3f18c5a7a5420d5754acc63e3f diff --git a/.github/workflows/create-release.yml b/.github/workflows/create-release.yml new file mode 100644 index 000000000..031f619d2 --- /dev/null +++ b/.github/workflows/create-release.yml @@ -0,0 +1,127 @@ +name: Create Release + +on: + workflow_dispatch: + +concurrency: + group: publish-crates + cancel-in-progress: false + +jobs: + publish: + name: Publish crates + runs-on: ubuntu-latest + timeout-minutes: 30 + environment: release + permissions: + contents: write + id-token: write + + steps: + - name: Require a release branch + shell: bash + run: | + if [[ "$GITHUB_REF" != "refs/heads/main" && ! "$GITHUB_REF_NAME" =~ ^releases/[0-9]+\.[0-9]+\.x$ ]]; then + echo "Publishing must be run from main or a releases/..x branch; got $GITHUB_REF." + exit 1 + fi + + - name: Check out repository + uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + with: + fetch-depth: 0 + + # Necessary as `cargo publish` builds crates as a check before publishing + - name: Set up repository + uses: ./.github/actions/setup + with: + mise-install-args: protoc + + - name: Read release metadata + id: release + shell: bash + run: | + metadata="$(cargo metadata --format-version 1 --no-deps)" + version="$(jq -er '[.packages[] | select(.name == "temporalio-sdk") | .version] | if length == 1 then .[0] else error("expected one temporalio-sdk package") end' <<<"$metadata")" + core_version="$(jq -er '[.packages[] | select(.name == "temporalio-sdk-core") | .version] | if length == 1 then .[0] else error("expected one temporalio-sdk-core package") end' <<<"$metadata")" + + if [[ "$GITHUB_REF_NAME" =~ ^releases/([0-9]+)\.([0-9]+)\.x$ ]]; then + release_line="${BASH_REMATCH[1]}.${BASH_REMATCH[2]}" + if [[ "$version" != "$release_line".* ]]; then + echo "SDK version $version does not belong to release branch $GITHUB_REF_NAME." + exit 1 + fi + fi + + { + echo "version=$version" + echo "tag=v$version" + echo "core_tag=core-v$core_version" + } >>"$GITHUB_OUTPUT" + + - name: Generate release notes + run: | + previous_tag="$(git describe --tags --match 'v[0-9]*' --abbrev=0 HEAD)" + cargo run -p changelog-release-notes -- \ + --from "$previous_tag" \ + --to HEAD \ + --changelog rust \ + >"$RUNNER_TEMP/sdk-release-notes.md" + + if ! previous_core_tag="$(git describe --tags --match 'core-v[0-9]*' --abbrev=0 HEAD 2>/dev/null)"; then + previous_core_tag="$previous_tag" + fi + cargo run -p changelog-release-notes -- \ + --from "$previous_core_tag" \ + --to HEAD \ + --changelog core \ + >"$RUNNER_TEMP/core-release-notes.md" + + - name: Plan crate publication + id: publish-plan + run: cargo run --quiet -p changelog-release-notes --bin plan-release-publish >>"$GITHUB_OUTPUT" + + - name: Authenticate to crates.io + if: steps.publish-plan.outputs.packages != '[]' + id: auth + uses: rust-lang/crates-io-auth-action@c6f97d42243bad5fab37ca0427f495c86d5b1a18 # v1 + + - name: Publish crates + if: steps.publish-plan.outputs.packages != '[]' + env: + CARGO_REGISTRY_TOKEN: ${{ steps.auth.outputs.token }} + PUBLISH_PACKAGES: ${{ steps.publish-plan.outputs.packages }} + shell: bash + run: | + jq -e 'type == "array" and length > 0 and all(.[]; type == "string" and length > 0)' \ + <<<"$PUBLISH_PACKAGES" >/dev/null + mapfile -t packages < <(jq -r '.[]' <<<"$PUBLISH_PACKAGES") + publish_args=() + for package in "${packages[@]}"; do + publish_args+=(--package "$package") + done + cargo publish "${publish_args[@]}" + + - name: Create draft SDK GitHub release + env: + GH_TOKEN: ${{ github.token }} + RELEASE_TAG: ${{ steps.release.outputs.tag }} + RELEASE_TITLE: temporalio-sdk ${{ steps.release.outputs.tag }} + run: | + gh release create "$RELEASE_TAG" \ + --target "$GITHUB_SHA" \ + --draft \ + --title "$RELEASE_TITLE" \ + --notes-file "$RUNNER_TEMP/sdk-release-notes.md" + + - name: Create draft Core GitHub release + env: + GH_TOKEN: ${{ github.token }} + RELEASE_TAG: ${{ steps.release.outputs.core_tag }} + RELEASE_TITLE: temporalio-sdk-core ${{ steps.release.outputs.core_tag }} + run: | + gh release create "$RELEASE_TAG" \ + --target "$GITHUB_SHA" \ + --draft \ + --title "$RELEASE_TITLE" \ + --notes-file "$RUNNER_TEMP/core-release-notes.md" diff --git a/.github/workflows/per-pr.yml b/.github/workflows/per-pr.yml index d04428830..c088bb13a 100644 --- a/.github/workflows/per-pr.yml +++ b/.github/workflows/per-pr.yml @@ -31,6 +31,12 @@ jobs: - run: cargo lint - run: cargo test-lint - run: cargo check + - run: cargo check --features experimental + - name: Check vendored proto compilation + run: cargo check --features temporalio-client/vendored-protox + env: + PROTOC: /does/not/exist + CARGO_TARGET_DIR: /tmp/vendored-protox-target - run: git diff --exit-code test: @@ -67,12 +73,12 @@ jobs: with: script: | core.exportVariable('RUSTFLAGS', '-Csymbol-mangling-version=v0'); - - run: cargo test -- --include-ignored --nocapture + - run: cargo test --features experimental -- --include-ignored --nocapture - name: Find test executable for cgroup tests id: find-cgroup-test if: runner.os == 'Linux' && runner.arch == 'X64' run: | - test_executable=$(cargo build --tests --message-format json | jq -r 'select(.profile?.test == true and .target?.name == "temporalio_sdk_core" and .executable) | .executable') + test_executable=$(cargo build --tests --features experimental --message-format json | jq -r 'select(.profile?.test == true and .target?.name == "temporalio_sdk_core" and .executable) | .executable') cp $test_executable ./core-tests - name: Upload cgroup test executable uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7 @@ -90,17 +96,29 @@ jobs: path: machine_coverage/ msrv: - name: "Verify MSRV" - timeout-minutes: ${{ github.ref == 'refs/heads/main' && 15 || 10 }} + name: "Verify MSRV (${{ matrix.group }})" + timeout-minutes: ${{ github.ref == 'refs/heads/main' && 20 || 15 }} runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + include: + - group: rust-1.88 + crates: client common common-wasm macros protos sdk-core workflow + - group: rust-1.92 + crates: sdk steps: - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 - uses: ./.github/actions/setup with: mise-install-args: protoc cargo:cargo-msrv - rust-cache-key: msrv - - run: cargo msrv verify - working-directory: ./crates/sdk-core + rust-cache-key: msrv-${{ matrix.group }} + - name: Verify crate MSRVs + run: | + read -ra crates <<< "${{ matrix.crates }}" + for crate in "${crates[@]}"; do + cargo msrv verify --all-features --path "crates/$crate" + done cgroup-tests: name: cgroup tests @@ -171,13 +189,10 @@ jobs: - run: cargo integ-test -t wasm_workflow_tests -- --test-threads 1 cloud-tests: - if: github.event.pull_request.head.repo.full_name == '' || github.event.pull_request.head.repo.full_name == 'temporalio/sdk-rust' + if: github.actor != 'dependabot[bot]' && (github.event.pull_request.head.repo.full_name == '' || github.event.pull_request.head.repo.full_name == 'temporalio/sdk-rust') name: Cloud tests env: - TEMPORAL_CLOUD_ADDRESS: https://${{ vars.TEMPORAL_CLIENT_NAMESPACE }}.tmprl.cloud:7233 - TEMPORAL_NAMESPACE: ${{ vars.TEMPORAL_CLIENT_NAMESPACE }} - TEMPORAL_CLIENT_CERT: ${{ secrets.TEMPORAL_CLIENT_CERT }} - TEMPORAL_CLIENT_KEY: ${{ secrets.TEMPORAL_CLIENT_KEY }} + TEMPORAL_CLIENT_CLOUD_API_VERSION: v0.19.1 timeout-minutes: ${{ github.ref == 'refs/heads/main' && 25 || 20 }} runs-on: ubuntu-latest steps: @@ -185,7 +200,62 @@ jobs: - uses: ./.github/actions/setup with: mise-install-args: protoc - - run: cargo test --features=test-utilities --test cloud_tests + - env: + TEMPORAL_CLOUD_ADDRESS: https://${{ vars.TEMPORAL_CLIENT_NAMESPACE }}.tmprl.cloud:7233 + TEMPORAL_NAMESPACE: ${{ vars.TEMPORAL_CLIENT_NAMESPACE }} + TEMPORAL_CLIENT_CERT: ${{ secrets.TEMPORAL_CLIENT_CERT }} + TEMPORAL_CLIENT_KEY: ${{ secrets.TEMPORAL_CLIENT_KEY }} + run: cargo test --features=test-utilities,temporalio-sdk/experimental --test cloud_tests + - name: Generate Cloud test certificates + run: | + umask 077 + cert_dir="$RUNNER_TEMP/cloud-test-certs" + mkdir "$cert_dir" + openssl req -x509 -newkey rsa:2048 -nodes -days 1 \ + -keyout "$cert_dir/ca.key" -out "$cert_dir/ca.pem" \ + -subj '/CN=Temporal Rust SDK Cloud CI CA' + openssl req -newkey rsa:2048 -nodes \ + -keyout "$cert_dir/client.key" -out "$cert_dir/client.csr" \ + -subj '/CN=Temporal Rust SDK Cloud CI' + openssl x509 -req -days 1 -in "$cert_dir/client.csr" \ + -CA "$cert_dir/ca.pem" -CAkey "$cert_dir/ca.key" -CAcreateserial \ + -out "$cert_dir/client.pem" -extfile <(printf 'extendedKeyUsage=clientAuth') + { + echo "TEMPORAL_CLOUD_CLIENT_CA_PATH=$cert_dir/ca.pem" + echo "TEMPORAL_TLS_CLIENT_CERT_PATH=$cert_dir/client.pem" + echo "TEMPORAL_TLS_CLIENT_KEY_PATH=$cert_dir/client.key" + } >> "$GITHUB_ENV" + - name: Create Cloud namespace + id: create-cloud-namespace + env: + TEMPORAL_CLIENT_CLOUD_API_KEY: ${{ secrets.TEMPORAL_CLIENT_CLOUD_API_KEY }} + run: cargo integ-test cloud-namespace create + - name: Run Cloud-eligible integration tests + timeout-minutes: 20 + env: + TEMPORAL_ADDRESS: ${{ steps.create-cloud-namespace.outputs.namespace }}.tmprl.cloud:7233 + TEMPORAL_NAMESPACE: ${{ steps.create-cloud-namespace.outputs.namespace }} + run: | + set -o pipefail + cargo integ-test -s envconfig --cloud 2>&1 | tee cloud-integration.log + - name: Delete Cloud namespace + id: delete-cloud-namespace + if: ${{ always() && steps.create-cloud-namespace.outputs.namespace != '' }} + continue-on-error: true + env: + TEMPORAL_CLIENT_CLOUD_API_KEY: ${{ secrets.TEMPORAL_CLIENT_CLOUD_API_KEY }} + run: cargo integ-test cloud-namespace delete "${{ steps.create-cloud-namespace.outputs.namespace }}" + - name: Report Cloud namespace cleanup failure + if: ${{ always() && steps.delete-cloud-namespace.outcome == 'failure' }} + run: echo "::warning title=Cloud namespace cleanup failed::Failed to delete Cloud namespace ${{ steps.create-cloud-namespace.outputs.namespace }}" + - name: Upload Cloud test output + if: always() + uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7 + with: + name: cloud-integration-tests-output + path: cloud-integration.log + if-no-files-found: ignore + retention-days: 14 docker-integ-tests: name: Docker integ tests @@ -203,7 +273,7 @@ jobs: with: mise-install-args: protoc - name: Start container for otel-collector and prometheus - uses: hoverkraft-tech/compose-action@d2bee4f07e8ca410d6b196d00f90c12e7d48c33a # v2.6.0 + uses: hoverkraft-tech/compose-action@ee6af68587292d72db67743171809c19787df4c9 # v3.1.0 with: compose-file: ./etc/docker/docker-compose-ci.yaml - run: cargo integ-test docker_ @@ -223,11 +293,23 @@ jobs: )" test -x "$integ_test_binary" test -x "$integ_runner_binary" + # The CLI is cached under a version-derived name, so the container only finds the one + # the previous step downloaded if it asks for the same version. Cargo's `[env]` does not + # reach a directly invoked binary, hence passing the pin, if any, explicitly. + cli_version_override="$( + sed -n 's/^CLI_VERSION_OVERRIDE = "\(.*\)"$/\1/p' .cargo/config.toml + )" + cli_version_env=() + if [ -n "$cli_version_override" ]; then + cli_version_env=(--env "CLI_VERSION_OVERRIDE=$cli_version_override") + fi docker run --rm \ --env TMPDIR="$RUNNER_TEMP" \ --env TEMPORAL_TEST_EXPECT_DOCKER=true \ + "${cli_version_env[@]}" \ --volume "$RUNNER_TEMP:$RUNNER_TEMP" \ --volume "$GITHUB_WORKSPACE:$GITHUB_WORKSPACE" \ + --volume /etc/ssl/certs/ca-certificates.crt:/etc/ssl/certs/ca-certificates.crt:ro \ --workdir "$GITHUB_WORKSPACE" \ ubuntu:24.04 \ "$integ_runner_binary" --test-executable "$integ_test_binary" \ diff --git a/.github/workflows/sdk-sentinel-pr-responder.yml b/.github/workflows/sdk-sentinel-pr-responder.yml new file mode 100644 index 000000000..5389ff08a --- /dev/null +++ b/.github/workflows/sdk-sentinel-pr-responder.yml @@ -0,0 +1,57 @@ +name: SDK Sentinel PR responder relay + +on: + issue_comment: + types: [created] + +permissions: {} + +jobs: + relay: + if: >- + github.event.issue.pull_request && + github.event.comment.user.login != 'sdk-sentinel-bot' && + contains(github.event.comment.body, '@sdk-sentinel-bot') && + contains(fromJSON('["OWNER","MEMBER","COLLABORATOR"]'), github.event.comment.author_association) + runs-on: ubuntu-latest + timeout-minutes: 5 + steps: + - name: Mint a Sentinel-only dispatch token from the Relay App + id: dispatch-token + uses: actions/create-github-app-token@fee1f7d63c2ff003460e3d139729b119787bc349 # v2 + with: + app-id: ${{ vars.SDK_SENTINEL_RELAY_APP_ID }} + private-key: ${{ secrets.SDK_SENTINEL_RELAY_PRIVATE_KEY }} + owner: temporalio + repositories: sdk-sentinel + permission-actions: write + permission-metadata: read + - name: Verify the Sentinel Relay installation + env: + ACTUAL_APP_SLUG: ${{ steps.dispatch-token.outputs.app-slug }} + ACTUAL_INSTALLATION_ID: ${{ steps.dispatch-token.outputs.installation-id }} + EXPECTED_INSTALLATION_ID: ${{ vars.SDK_SENTINEL_RELAY_INSTALLATION_ID }} + run: | + test "$ACTUAL_APP_SLUG" = sdk-sentinel-relay + test "$ACTUAL_INSTALLATION_ID" = "$EXPECTED_INSTALLATION_ID" + - name: Dispatch the trusted central responder + env: + COMMENT_ID: ${{ github.event.comment.id }} + DISPATCH_TOKEN: ${{ steps.dispatch-token.outputs.token }} + PR_NUMBER: ${{ github.event.issue.number }} + TARGET_ID: rust + run: | + payload="$( + jq -cn \ + --arg target "$TARGET_ID" \ + --arg pr_number "$PR_NUMBER" \ + --arg comment_id "$COMMENT_ID" \ + '{ref:"main",inputs:{target:$target,pr_number:$pr_number,comment_id:$comment_id}}' + )" + curl --fail --silent --show-error \ + --request POST \ + --header "Accept: application/vnd.github+json" \ + --header "Authorization: Bearer $DISPATCH_TOKEN" \ + --header "X-GitHub-Api-Version: 2022-11-28" \ + --data "$payload" \ + https://api.github.com/repos/temporalio/sdk-sentinel/actions/workflows/sdk-pr-responder.yml/dispatches diff --git a/CHANGELOG.md b/CHANGELOG.md index 1264c8240..e559ee2b6 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -33,9 +33,213 @@ relevant information. ## Unreleased +### Fixed +* `temporalio-common`'s build script now generates its payload-visitor implementations in a + stable order. The generated code was emitted in `HashSet`/`HashMap` iteration order, so its + content changed on every build and a compilation cache such as sccache missed + `temporalio-common` and every crate downstream of it on every build. +* Autoscaled task pollers now preserve polling concurrency after transient cancellations and + timeouts while still applying retry backoff. +* Sticky workflow backlog no longer prevents normal pollers from using capacity after sticky + pollers reach their polling limit. +* Workflow task failures are now reported to the server only on a task's first attempt, no + matter why the task failed. Previously a completion rejected for exceeding the worker's payload + size error limit, or a failure to fetch workflow history, was re-reported on every retry. +* The `temporal_workflow_task_execution_failed` metric now counts every failed workflow task + attempt, including ones whose failure was not sent to the server, and tags size-related failures + with their specific `failure_reason` (`GrpcMessageTooLarge`, `PayloadsTooLarge`, + `RequestTooLarge`) on every path. + +## [1.0.0] - 2026-09-04 + +### Changed +* Published Rust SDK crates now declare their minimum supported Rust version: Rust 1.92 for + `temporalio-sdk`, and Rust 1.88 for the other crates. + ### Added +* `temporalio-client` now provides a `vendored-protox` feature for compiling protobuf definitions + with `protox`, allowing client and Rust SDK builds without an installed `protoc`. +* `CancelExternalWorkflowError` and `workflow_interceptors::CancelExternalWorkflowResult` + for use in interceptors. +* `WorkflowContextKey` and context-value scopes provide replay-safe, workflow-run-owned context + storage for application code and workflow interceptors. Values survive async suspension while + remaining isolated between concurrent branches and signal/update handlers. Read-only workflow + views can observe values established by synchronous inbound interceptors, and outbound + interceptors can use values to propagate metadata to activities, child workflows, signals, + Nexus operations, and continue-as-new runs. +### Breaking Changes +* `ActivityEnvironmentBuilder` no longer accepts a `tokio_util::sync::CancellationToken`. + Use `ActivityEnvironment::cancel` to cancel activities running in the test environment. +* `ChildWorkflowStartError::StartFailed` now reports an SDK-owned, non-exhaustive + `StartChildWorkflowExecutionFailedCause` instead of a generated protobuf enum. Add a wildcard + branch when matching the cause. +* Client options now use client-owned, non-exhaustive `WorkflowIdReusePolicy`, + `WorkflowIdConflictPolicy`, `QueryRejectCondition`, `ArchivalState`, and + `HistoryEventFilterType` enums instead of generated protobuf enums. Add wildcard branches when + matching these types. +* `ActivityError`, `PayloadConversionError`, `ActivityExecutionError`, + `ChildWorkflowStartError`, `ChildWorkflowExecutionError`, and `WorkflowSignalError` are now + non-exhaustive. Add wildcard branches when matching these enums. + for use in interceptors, with `WorkflowCancelFailureError` exposing the decoded cancellation + failure. +* `WorkflowSignalError::NotFound` and `CancelExternalWorkflowError::NotFound` let workflows + distinguish a missing signal or cancellation target from other delivery failures. + +### Breaking Changes +* `temporalio-sdk` now exposes SDK-owned worker tuner and slot supplier types instead of + re-exporting the corresponding `temporalio-sdk-core` traits. +* `temporalio-sdk` now owns its runtime, polling, workflow-error, worker-validation, and local test + server configuration types instead of re-exporting their `temporalio-sdk-core` equivalents. + `TokioRuntimeBuilder` is non-exhaustive and provides a builder for construction. Unrelated Core + worker configuration and replay helpers are no longer re-exported by the SDK. + `PollerBehavior::Autoscaling` now holds builder-created `AutoscalingOptions`. +* `Runtime::from_current_tokio` replaces `Runtime::new_assume_tokio` as the preferred constructor. + The old name remains as a deprecated alias, but now returns `RuntimeError` instead of + `anyhow::Error`. + +### Fixed +* `Runtime::from_current_tokio` now returns `RuntimeError::NoCurrentTokioRuntime` when called + without an active Tokio runtime instead of panicking. +* `Memo::keys` and `SearchAttributes::keys` now iterate in lexicographic order so workflow + decisions based on collection traversal remain deterministic during replay. + +## [0.8.0] - 2026-09-02 + +### Breaking Changes +* `SdkWakeGuard` is no longer `Send` or `Sync`, preventing the thread-local guard from being moved + to or referenced from another thread. +* `WorkflowContextView` and the containing `PatchActivationInput` are no longer `Send` or `Sync` + because workflow context views now share single-threaded workflow randomness with replay-sensitive + SDK integrations. + +### Added +* `DefaultFailureConverter::new(true)` moves failure messages and stack traces into encoded + attributes so payload codecs can encrypt them. +* `LocalActivityOptions::include_arguments_in_marker` allows Rust workflows to opt in to + recording local activity arguments in Workflow history. +* `WorkflowContext::random_stream`, `SyncWorkflowContext::random_stream`, and + `WorkflowInterceptorContext::random_stream` provide deterministic, workflow-run-scoped + pseudo-random streams isolated by a stable caller-supplied name. Repeated lookup continues a + named stream without consuming the workflow's default randomness or any other named stream. +* `WorkflowHandle::get_update_handle` creates a typed handle for an existing Workflow Update from + its update ID, allowing callers to wait for the result independently of the original handle. +* `WorkflowContext::all_handlers_finished` and `SyncWorkflowContext::all_handlers_finished` let + Rust workflows wait for active signal and update handler chains before completing or continuing + as new. +* `WorkflowStartOptions::memo` attaches a non-indexed memo when starting a workflow, using the + same `MemoValues` type already used by continue-as-new and `WorkflowContext::upsert_memo`. + Values are serialized with the client's payload converter and codec, matching how `describe` + and `list` read them back. +* `MemoValue` and `MemoValues` are now exported from `temporalio_common` as well as + `temporalio_workflow`, so the same types can be used from clients and workflows. +* The `temporal_activity_execution_failed` and `temporal_local_activity_execution_failed` worker + metrics now carry a `failure_reason` attribute. Each is now split into one time series per + reason, which may affect existing dashboards. +* Workflow task completions larger than the gRPC request size limit are now paginated automatically when the namespace supports it. Paginated workflow task completions require Temporal Server 1.32.0 or later. +* Update-with-Start support: `Client::start_update_with_start_workflow` and + `Client::execute_update_with_start_workflow` start a workflow and send it an update in one atomic + operation. `WorkflowUpdateWithStartOptions` requires an ID conflict policy (use `UseExisting` to + attach an update to an already-running workflow), provides distinct start and update headers, + and controls the atomic RPC. The operation can be intercepted via + `ClientInterceptor::update_with_start_workflow`. + +### Breaking Changes :boom: +* `DefaultFailureConverter` is no longer a unit struct. Use `DefaultFailureConverter::default()` instead. +* `WorkflowHandle::fetch_history` now returns a lazy `WorkflowHistory` stream instead + of eagerly fetching every history page. Use `WorkflowHistory::into_events` if eager fetching + is desired. +* `WorkflowHistory::to_json` is now async, `WorkflowHistoryError` reports fetch and JSON conversion failures; and the eager `events`, + `Clone`, and `From for History` APIs have been removed. Replay results expose + their eagerly fetched events through `ReplayHistory`. +* Rust SDK APIs previously marked experimental now require the `experimental` Cargo feature. This + includes: + * Nexus operation caller and workflow interceptor APIs, including `NexusOperationOptions`, + `NexusOperationCancellationType`, `StartedNexusOperation`, the workflow-context start methods, + and the `WorkflowInterceptor::start_nexus_operation` hook and input/result types. + * Worker deployment versioning APIs: `ContinueAsNewVersioningBehavior`, + `ContinueAsNewOptions::initial_versioning_behavior`, + `WorkflowContext::target_worker_deployment_version_changed`, and + `SyncWorkflowContext::target_worker_deployment_version_changed`. + * Client, worker, and workflow replayer plugin APIs. + * Client payload warning thresholds (`PayloadLimitsOptions` and + `ConnectionOptions::payload_limits`) and `WorkerOptions::disable_payload_error_limit`. + * Patch activation callback types and the corresponding worker option. + * Worker lifecycle interception APIs (`WorkerInterceptor`, its input types and registration + methods, and `ReturnWorkflowExitValueInterceptor`). + * Event Group marker fields on activity, local activity, child workflow, timer, and external + signal options. +* The following types are now non-exhaustive: `Priority`, `WorkerDeploymentVersion`, + `WorkerCallbacks`, `WorkflowExecutionInfo`, `ActivityCloseTimeouts`, + `ActivityExecutionDecodeHint`, child-workflow and signal decode hints, + `SerializationContext`, `SerializationContextData`, `PayloadConverter`, `IncomingError`, + `ScheduleSpec`, and `ScheduleOverlapPolicy`. Construct structs using their respective builders + or constructors (`WorkerCallbacks::new`, `ActivityExecutionDecodeHint::new`, or + `SerializationContext::new`); use `Default` for `PayloadConverter`; and add wildcard branches + when matching enums. +* `SerializationContextData::{Workflow, Activity, Nexus}` now contain corresponding context + structs. `SerializationContextData` is no longer `Copy`. +* Renamed `ActivityCloseTimeouts::Both` to `ActivityCloseTimeouts::ScheduleAndStartToClose`. +* Removed the unused `ActExitValue` type. Use `ActivityError::WillCompleteAsync` to mark an + activity for asynchronous completion. +* Removed the test-only `FailOnNondeterminismInterceptor` from the public Rust SDK API. +* Environment configuration values (`DataSource`, `ClientConfig`, and related profile, TLS, and + codec types) are now non-exhaustive. Use their `bon` builders to construct configuration structs, + and add a wildcard branch when matching `DataSource`. +* Values stored in a `MemoValue` must now be `Send + Sync`. It previously held its value in an + `Rc` and now uses an `Arc`, so that memos can be built outside a workflow and handed to the + client. Only affects memo values that are themselves non-`Send`/non-`Sync`, such as those + holding an `Rc` or `RefCell`. + +* `Client::signal_with_start_workflow` starts a workflow and sends a typed signal atomically. + +### Breaking Changes :boom: + +* Signal-with-start is now invoked with `Client::signal_with_start_workflow`; remove uses of + `WorkflowStartOptions::start_signal` and `WorkflowStartSignal`. + +### Fixed +* `Worker` shutdown no longer loses an activity result it was still reporting to the server. If + shutdown raced such a completion — most likely while the activity's final heartbeat RPC was + still in flight — the worker could strand the completion forever: debug builds panicked with + `Waiting for all slot permits to release took too long!`, and release builds logged that error + and dropped the result, leaving the server to time the activity out before retrying it. + Shutdown now drains in-flight completions first. +* Standalone activity result and describe APIs now apply the configured payload codec when + decoding failures, so encoded failure attributes and details are restored correctly. +* The default payload converter now serializes Serde `null` values such as `Option::None` as + `binary/null` and accepts both `binary/null` and legacy `json/plain` null payloads. +* The Prometheus exporter now respects `PrometheusExporterOptions::counters_total_suffix`, + appending `_total` to counter metric names when enabled. +* Workflow start requests now include the client's identity. +* An activity failure caused by oversized final heartbeat details is now counted in the + `temporal_activity_execution_failed` metric as `failure_reason="PayloadsTooLarge"`. Previously it + was counted under the reason for the failure the activity itself reported, and was not counted at + all when that failure was benign, even though the worker reported a payload-limit failure to the + server. +* Workers configured with a small `max_cached_workflows` no longer briefly stop accepting new + workflows. Sticky workflow-task pollers could consume every workflow-cache permit and starve the + non-sticky poller, so the worker would stop picking up new workflows until a poll timed out (up to + ~60s). The poll balancer now reserves a non-sticky slot against the cache size rather than the + slot-supplier size. +* The default payload converter now encodes `Vec` and `Option>` as `binary/plain` when + present. + +## [0.7.0] - 2026-08-17 + +### Added +* Support for running Standalone Activities in Rust SDK Worker. +* Client methods for starting and managing execution of Standalone Activities. +* `WorkflowTermination::cancelled_with_details` for recording structured details when a Workflow + Execution completes as cancelled. * `LoggerFormat` for selecting compact, pretty, or JSON Core console log output. Configured log filters continue to apply to JSON output. +* The Rust SDK now has an optional `testing` feature with a typed activity test environment and + local or external workflow test environments. Local workflow environments manage a Temporal CLI + dev server and expose shutdown through their local-server type state. +* Worker heartbeats now report the SDK runtime, hosting environments, operating system, and + architecture once per worker, retrying until the first successful delivery. Runtime options can + disable this reporting, and language SDK bridges can supply their own runtime details. The Rust + SDK exposes separate runtime options that omit bridge-only runtime overrides. * `RpcOptions::builder()` for constructing per-call RPC options. * `DnsLoadBalancingOptions::builder()` for configuring DNS re-resolution intervals. * Experimental plugin APIs for packaging reusable client and worker configuration, including data @@ -66,6 +270,14 @@ relevant information. `WorkerInterceptor::with_workflow_replay_worker`. ### Breaking Changes :boom: +* `WorkflowTermination::Cancelled` now has an optional `details` field. Use + `WorkflowTermination::cancelled()` to construct a cancellation without details. +* Changes to `ActivityInfo`: instead of `workflow_namespace`, `workflow_execution` and `run_id`, + there is now `namespace`, `workflow_id`, `workflow_run_id` and `activity_run_id`. + Also, `workflow_type` is now `Option`. +* `ActivityIdentifier::ById` was split into 2 variants, `ByIdWorkflow` and `ByIdStandalone`. + `ActivityIdentifier::by_id` method was renamed to `by_id_workflow`, and `by_id_standalone` + was added. * `anyhow::Error` no longer converts directly into `WorkflowTermination`. Wrap an error in `ApplicationFailure` to explicitly fail the Workflow Execution. * `OutgoingWorkflowError` now has a dedicated `PayloadConversion` variant. Converting activity, @@ -73,6 +285,12 @@ relevant information. `OutgoingError`, `OutgoingActivityError`, and `OutgoingWorkflowError` are now non-exhaustive; downstream matches must include a wildcard arm. * Removed `InterceptorWithNext`. Register worker interceptors as an ordered vector instead. +* Ephemeral server APIs now return `EphemeralServerError` instead of `anyhow::Error`, and dev-server + log format and level use the non-exhaustive `DevServerLogFormat` and `DevServerLogLevel` enums. +* Ephemeral server APIs now return the operation-oriented `EphemeralServerError` instead of + `anyhow::Error`. +* Activity macro support now exposes instance requirements through `ExecutableActivity`; the + redundant `HasOnlyStaticMethods` marker trait has been removed. * `Worker::run` now returns `WorkerRunError` instead of `anyhow::Error`. Non-validation failures are reported as `WorkerRunError::Fatal` with a message and source. * `Logger::Console` now requires a `format: Option` field. Use `None` to preserve the @@ -126,6 +344,9 @@ relevant information. * Workflow workers now preserve the outstanding workflow-task token if an internal admission invariant is violated, buffering the replacement task instead of overwriting the task in flight in release builds. +* Rust SDK workers now warn when autoscaling task polling encounters errors continuously for one + minute. Repeated warnings use exponential backoff up to 15-minute intervals and stop after + polling recovers. * Unhandled workflow payload conversion errors now fail the Workflow Task so it can retry instead of failing the Workflow Execution. Workflows may still explicitly handle these errors. * Workers no longer send worker heartbeats or appear in centralized heartbeat reports before diff --git a/Cargo.toml b/Cargo.toml index c9165dbb0..cf6a8884d 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -9,6 +9,7 @@ members = [ "crates/sdk", "crates/workflow", "crates/sdk-core-c-bridge", + "crates/changelog-release-notes", ] resolver = "2" diff --git a/DOC_HIDDEN_IDENTIFIERS.md b/DOC_HIDDEN_IDENTIFIERS.md new file mode 100644 index 000000000..3fd785324 --- /dev/null +++ b/DOC_HIDDEN_IDENTIFIERS.md @@ -0,0 +1,47 @@ +# `doc(hidden)` identifiers + +This inventory was generated from Rust source with: + +```console +rg -n --glob '*.rs' '^\s*#\s*!?\s*\[\s*doc\s*\(\s*hidden\s*\)\s*\]' crates +``` + +The anchored expression avoids matching `#[doc(hidden)]` mentioned in comments. Identifiers in a +grouped `pub use` are listed separately. The generated function whose source is emitted by +`crates/common/build.rs` is included as well. Locations point to the identifier's declaration or +re-export rather than to the preceding attribute. + +| Identifier | Kind | Location | +| --- | --- | --- | +| `temporalio_client::retry::jittered` | function | [`crates/client/src/retry.rs:160`](crates/client/src/retry.rs#L160) | +| `temporalio_client::grpc::PayloadLimitsClient` | struct | [`crates/client/src/grpc.rs:415`](crates/client/src/grpc.rs#L415) | +| `temporalio_client::request_extensions::PayloadErrorLimits` | struct | [`crates/client/src/request_extensions.rs:31`](crates/client/src/request_extensions.rs#L31) | +| `temporalio_client::worker::NamespaceDescriptionSource` | struct | [`crates/client/src/worker.rs:324`](crates/client/src/worker.rs#L324) | +| `temporalio_client::worker::ClientWorkerSet::namespace_description_source` | method | [`crates/client/src/worker.rs:420`](crates/client/src/worker.rs#L420) | +| `temporalio_client::jittered` | re-exported function | [`crates/client/src/lib.rs:47`](crates/client/src/lib.rs#L47) | +| `temporalio_client::MESSAGE_TOO_LARGE_KEY` | static | [`crates/client/src/lib.rs:190`](crates/client/src/lib.rs#L190) | +| `temporalio_client::ERROR_RETURNED_DUE_TO_SHORT_CIRCUIT` | static | [`crates/client/src/lib.rs:193`](crates/client/src/lib.rs#L193) | +| `temporalio_common::fsm_trait` | module | [`crates/common/src/lib.rs:13`](crates/common/src/lib.rs#L13) | +| `temporalio_common::payload_limits` | module | [`crates/common/src/lib.rs:15`](crates/common/src/lib.rs#L15) | +| `temporalio_common::payload_limits::LimitClass` | enum | [`crates/common/src/payload_limits.rs:20`](crates/common/src/payload_limits.rs#L20) | +| `temporalio_common::payload_limits::PayloadLimits` | struct | [`crates/common/src/payload_limits.rs:151`](crates/common/src/payload_limits.rs#L151) | +| `temporalio_common::payload_limits::LimitSeverity` | enum | [`crates/common/src/payload_limits.rs:175`](crates/common/src/payload_limits.rs#L175) | +| `temporalio_common::payload_limits::PayloadLimitViolation` | struct | [`crates/common/src/payload_limits.rs:185`](crates/common/src/payload_limits.rs#L185) | +| `temporalio_common::payload_visitor::validate_known_payload_limits` | generated function | [`crates/common/build.rs:993`](crates/common/build.rs#L993) | +| `temporalio_sdk::WorkerOptions::to_core_options` | method | [`crates/sdk/src/lib.rs:662`](crates/sdk/src/lib.rs#L662) | +| `temporalio_workflow::component::bindings` | module | [`crates/workflow/src/component.rs:31`](crates/workflow/src/component.rs#L31) | +| `temporalio_workflow::__private` | module | [`crates/workflow/src/lib.rs:14`](crates/workflow/src/lib.rs#L14) | +| `temporalio_workflow::InternalPatchActivationCallback` | re-exported type alias | [`crates/workflow/src/lib.rs:88`](crates/workflow/src/lib.rs#L88) | +| `temporalio_workflow::PatchActivationCaller` | re-exported struct | [`crates/workflow/src/lib.rs:88`](crates/workflow/src/lib.rs#L88) | +| `temporalio_workflow::__temporal_select` | macro | [`crates/workflow/src/lib.rs:94`](crates/workflow/src/lib.rs#L94) | +| `temporalio_workflow::__temporal_join` | macro | [`crates/workflow/src/lib.rs:102`](crates/workflow/src/lib.rs#L102) | +| `temporalio_workflow::__temporalio_export_workflow_component` | macro | [`crates/workflow/src/lib.rs:110`](crates/workflow/src/lib.rs#L110) | +| `temporalio_workflow::PatchActivationCaller` | struct | [`crates/workflow/src/workflow_context.rs:221`](crates/workflow/src/workflow_context.rs#L221) | +| `temporalio_workflow::BaseWorkflowContext::from_raw` | method | [`crates/workflow/src/workflow_context.rs:659`](crates/workflow/src/workflow_context.rs#L659) | + +The former public `temporalio_workflow::component` and `temporalio_workflow::runtime` modules no +longer appear in this inventory. They are private implementation modules; exports needed by macro +expansions and the SDK are now routed through the explicitly internal +`temporalio_workflow::__private` module. The generated component bindings retain their own +annotation because the exported macro must be able to name them from the workflow author's crate, +even though they are implementation details declared within the private component module. diff --git a/README.md b/README.md index e1ab5ae63..08c56fb29 100644 --- a/README.md +++ b/README.md @@ -5,14 +5,14 @@ [![crates.io](https://img.shields.io/crates/v/temporalio-sdk.svg)](https://crates.io/crates/temporalio-sdk) [![docs.rs](https://docs.rs/temporalio-sdk/badge.svg)](https://docs.rs/temporalio-sdk) -Currently in Public Preview, see more in the [SDK README.md](crates/sdk/README.md) +See the [SDK README](crates/sdk/README.md) for usage and examples. # Temporal Rust Client [![crates.io](https://img.shields.io/crates/v/temporalio-sdk.svg)](https://crates.io/crates/temporalio-client) [![docs.rs](https://docs.rs/temporalio-sdk/badge.svg)](https://docs.rs/temporalio-client) -Currently in Public Preview, see more in the [client README.md](crates/client/README.md) +See the [client README](crates/client/README.md) for usage and examples. # Temporal Core SDK @@ -25,8 +25,8 @@ Core SDK that can be used as a base for other Temporal SDKs. It is currently use # Documentation -Rust & Core SDK documentation can be generated with `cargo doc`, output will be placed in the -`target/doc` directory. +Rust & Core SDK documentation can be generated with `cargo doc --workspace --all-features`, output +will be placed in the `target/doc` directory. [Architecture](ARCHITECTURE.md) doc provides some high-level information about how Core SDK works and how language layers interact with it. @@ -47,7 +47,7 @@ This repo is composed of multiple crates: - temporalio-sdk-core `./crates/core` - The Core implementation. - temporalio-sdk-core-c-bridge `./crates/core-c-bridge` - Provides C bindings for Core. - temporalio-macros `./crates/macros` - Implements procedural macros used by core and the SDK. -- temporalio-sdk `./crates/sdk` - A Public Preview Rust SDK built on top of Core. Used for testing. +- temporalio-sdk `./crates/sdk` - A Rust SDK built on top of Core. Visualized (dev dependencies are in blue): @@ -59,12 +59,55 @@ All the following commands are enforced for each pull request: You can build and test the project using cargo: `cargo build` -`cargo test` +`cargo test --features experimental` Run integ tests with `cargo integ-test`. By default it will start an ephemeral server. You can also use an already-running server by passing `-s external`. -Run load tests with `cargo test --test heavy_tests`. +To run target-compatible integration tests against a client configured with +[envconfig](./crates/client/README.md), pass `-s envconfig` and select a test suitable for that +target: + +```bash +TEMPORAL_ADDRESS=namespace.account.tmprl.cloud:7233 \ +TEMPORAL_NAMESPACE=namespace.account \ +TEMPORAL_API_KEY=... \ +cargo integ-test -s envconfig -- \ + integ_tests::workflow_tests::timers::timer_workflow_workflow_driver --exact +``` + +`TEMPORAL_CONFIG_FILE` and `TEMPORAL_PROFILE` can select a TOML profile instead. The harness does +not start, configure, or clean up the target server or namespace in this mode. + +Pass `--cloud` to skip integration tests that do not exercise a server, require a local server or +unavailable Cloud provisioning, or are not yet compatible with Temporal Cloud. Cloud filtering is +independent of server selection. Run the eligible tests through envconfig, or list the excluded +cases without connecting to a server: + +```bash +cargo integ-test -s envconfig --cloud +cargo integ-test -s external --cloud -- --ignored --list +``` + +Tests are Cloud-eligible by default. An incompatible test must select a +`CloudTestExclusionReason` at the test function: + +```rust +#[temporalio_macros::cloud_test_exclusion( + crate::CloudTestExclusionReason::DoesNotUseServer +)] +#[tokio::test] +async fn example() { + // ... +} +``` + +The supported categories are `DoesNotUseServer`, `RequiresLocalServer`, `RequiresOssOnlyApis`, +`RequiresCloudProvisioning`, and `NeedsCloudAdaptation`. An optional second string should only be +used when it adds information beyond the category. In Cloud mode exclusions become native ignored +tests, so `--ignored --list` lists the skipped cases. + +Run load tests with `cargo test --features experimental --test heavy_tests`. NOTE: Integration tests should pass locally, if running on MacOS and you see integration tests consistently failing with an error that mentions `Too many open files`, this is likely due to `ulimit -n` being too low. You can raise diff --git a/crates/changelog-release-notes/Cargo.toml b/crates/changelog-release-notes/Cargo.toml new file mode 100644 index 000000000..1cfbc28b6 --- /dev/null +++ b/crates/changelog-release-notes/Cargo.toml @@ -0,0 +1,18 @@ +[package] +name = "changelog-release-notes" +version = "0.1.0" +edition = "2024" +license-file = { workspace = true } +publish = false +rust-version = "1.88.0" +default-run = "changelog-release-notes" + +[dependencies] +chrono = { version = "=0.4.45", default-features = false, features = ["clock"] } +reqwest = { version = "0.13", default-features = false, features = ["blocking", "rustls"] } +semver = "=1.0.28" +serde = { version = "1.0", features = ["derive"] } +serde_json = { workspace = true } + +[lints] +workspace = true diff --git a/crates/changelog-release-notes/src/bin/plan-release-publish.rs b/crates/changelog-release-notes/src/bin/plan-release-publish.rs new file mode 100644 index 000000000..d22fd78bc --- /dev/null +++ b/crates/changelog-release-notes/src/bin/plan-release-publish.rs @@ -0,0 +1,166 @@ +use std::{process::Command, time::Duration}; + +use reqwest::{StatusCode, blocking::Client}; +use serde::Deserialize; + +const CRATES_IO_SPARSE_INDEX: &str = "https://index.crates.io"; +const USER_AGENT: &str = + "temporalio/sdk-rust release planner (https://github.com/temporalio/sdk-rust)"; + +#[derive(Debug, Deserialize, PartialEq, Eq)] +struct Metadata { + packages: Vec, +} + +#[derive(Debug, Deserialize, PartialEq, Eq)] +struct Package { + name: String, + version: String, + publish: Option>, +} + +#[derive(Debug, Deserialize, PartialEq, Eq)] +struct IndexEntry { + vers: String, +} + +fn workspace_metadata() -> Result { + let output = Command::new("cargo") + .args(["metadata", "--format-version", "1", "--no-deps"]) + .output() + .map_err(|err| format!("failed to run cargo metadata: {err}"))?; + if !output.status.success() { + return Err(format!( + "cargo metadata failed: {}", + String::from_utf8_lossy(&output.stderr).trim() + )); + } + serde_json::from_slice(&output.stdout) + .map_err(|err| format!("failed to parse cargo metadata: {err}")) +} + +fn crates_io_packages(metadata: Metadata) -> impl Iterator { + metadata.packages.into_iter().filter(|package| { + package + .publish + .as_ref() + .is_none_or(|registries| registries.iter().any(|registry| registry == "crates-io")) + }) +} + +fn sparse_index_path(name: &str) -> String { + let name = name.to_ascii_lowercase(); + // See https://doc.rust-lang.org/cargo/reference/registry-index.html#index-files + match name.len() { + 1 => format!("1/{name}"), + 2 => format!("2/{name}"), + 3 => format!("3/{}/{name}", &name[..1]), + _ => format!("{}/{}/{name}", &name[..2], &name[2..4]), + } +} + +fn version_in_index(index: &str, version: &str) -> Result { + for line in index.lines() { + let entry: IndexEntry = serde_json::from_str(line) + .map_err(|err| format!("failed to parse sparse index entry: {err}"))?; + if entry.vers == version { + return Ok(true); + } + } + Ok(false) +} + +fn version_is_published(client: &Client, package: &Package) -> Result { + let url = format!( + "{CRATES_IO_SPARSE_INDEX}/{}", + sparse_index_path(&package.name) + ); + let response = client.get(&url).send().map_err(|err| { + format!( + "failed to check {}@{}: {err}", + package.name, package.version + ) + })?; + match response.status() { + StatusCode::OK => { + let index = response.text().map_err(|err| { + format!( + "failed to check {}@{}: failed to read sparse index entry: {err}", + package.name, package.version + ) + })?; + version_in_index(&index, &package.version).map_err(|err| { + format!( + "failed to check {}@{}: {err}", + package.name, package.version + ) + }) + } + StatusCode::NOT_FOUND => Ok(false), + status => Err(format!( + "failed to check {}@{}: sparse index returned unexpected HTTP status {status}", + package.name, package.version + )), + } +} + +fn main() -> Result<(), String> { + let client = Client::builder() + .user_agent(USER_AGENT) + .timeout(Duration::from_secs(30)) + .build() + .map_err(|err| format!("failed to create crates.io client: {err}"))?; + + let mut unpublished = Vec::new(); + for package in crates_io_packages(workspace_metadata()?) { + if version_is_published(&client, &package)? { + eprintln!( + "{}@{} is already published; skipping.", + package.name, package.version + ); + } else { + unpublished.push(package.name); + } + } + println!( + "packages={}", + serde_json::to_string(&unpublished) + .map_err(|err| format!("failed to serialize publish plan: {err}"))? + ); + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn selects_packages_publishable_to_crates_io() { + let packages = crates_io_packages(workspace_metadata().expect("workspace metadata")) + .map(|package| package.name) + .collect::>(); + + assert!(packages.iter().any(|package| package == "temporalio-sdk")); + assert!( + !packages + .iter() + .any(|package| package == "temporalio-sdk-core-c-bridge") + ); + } + + #[test] + fn builds_sparse_index_paths() { + assert_eq!(sparse_index_path("temporalio-sdk"), "te/mp/temporalio-sdk"); + } + + #[test] + fn finds_versions_in_sparse_index_entries() { + let index = r#"{"vers":"0.7.0","yanked":true} +{"vers":"0.8.0","yanked":false}"#; + + assert_eq!(version_in_index(index, "0.7.0"), Ok(true)); + assert_eq!(version_in_index(index, "0.8.0"), Ok(true)); + assert_eq!(version_in_index(index, "0.9.0"), Ok(false)); + assert!(version_in_index("not json", "0.7.0").is_err()); + } +} diff --git a/crates/changelog-release-notes/src/bin/prepare-release.rs b/crates/changelog-release-notes/src/bin/prepare-release.rs new file mode 100644 index 000000000..f886c38be --- /dev/null +++ b/crates/changelog-release-notes/src/bin/prepare-release.rs @@ -0,0 +1,214 @@ +use std::{ + env, fs, + path::{Path, PathBuf}, + process::Command, +}; + +use chrono::{NaiveDate, Utc}; +use semver::Version; + +const CORE_CRATE: &str = "temporalio-sdk-core"; +const BRIDGE_CRATE: &str = "temporalio-sdk-core-c-bridge"; +const PROTOS_CRATE: &str = "temporalio-protos"; +const RELEASE_TOOL_CRATE: &str = "changelog-release-notes"; + +fn workspace_root() -> PathBuf { + Path::new(env!("CARGO_MANIFEST_DIR")).join("../..") +} + +fn cargo_set_version(root: &Path, args: &[String]) -> Result<(), String> { + let status = Command::new("cargo") + .args(args) + .current_dir(root) + .status() + .map_err(|err| format!("failed to run `cargo {}`: {err}", args.join(" ")))?; + if !status.success() { + return Err(format!("`cargo {}` failed with {status}", args.join(" "))); + } + Ok(()) +} + +fn turn_over_changelog( + changelog: &str, + version: &Version, + date: NaiveDate, +) -> Result { + let release_prefix = format!("## [{version}]"); + if changelog + .lines() + .any(|line| line.starts_with(&release_prefix)) + { + return Err(format!("CHANGELOG.md already contains {release_prefix}")); + } + + let header = "## Unreleased"; + let header_start = changelog + .match_indices(header) + .find_map(|(index, _)| { + let at_line_start = index == 0 || changelog.as_bytes().get(index - 1) == Some(&b'\n'); + let after = index + header.len(); + let at_line_end = matches!(changelog.as_bytes().get(after), None | Some(b'\n')); + (at_line_start && at_line_end).then_some(index) + }) + .ok_or("CHANGELOG.md is missing an `## Unreleased` section")?; + let body_start = header_start + header.len(); + let following = &changelog[body_start..]; + let next_heading = following + .find("\n## ") + .ok_or("CHANGELOG.md is missing a released-version section")?; + let body = following[..next_heading].trim(); + let previous_releases = &following[next_heading + 1..]; + + let mut output = String::from(&changelog[..header_start]); + output.push_str(header); + output.push_str("\n\n"); + output.push_str(&format!("## [{version}] - {date}")); + if !body.is_empty() { + output.push_str("\n\n"); + output.push_str(body); + } + output.push_str("\n\n"); + output.push_str(previous_releases.trim_start_matches('\n')); + Ok(output) +} + +fn parse_args(args: impl IntoIterator) -> Result<(Version, Version), String> { + let mut args = args.into_iter(); + let expected = "expected "; + let sdk_version = args.next().ok_or(expected)?; + let core_protos_version = args.next().ok_or(expected)?; + if args.next().is_some() { + return Err(expected.into()); + } + let sdk_version = + Version::parse(&sdk_version).map_err(|err| format!("invalid SDK version: {err}"))?; + let core_protos_version = Version::parse(&core_protos_version) + .map_err(|err| format!("invalid Core/Protos version: {err}"))?; + if core_protos_version.major != 0 { + return Err(format!( + "Core/Protos version must remain 0.x, found {core_protos_version}" + )); + } + Ok((sdk_version, core_protos_version)) +} + +fn main() -> Result<(), String> { + let (target_sdk, target_core_protos) = parse_args(env::args().skip(1))?; + let root = workspace_root(); + let release_date = Utc::now().date_naive(); + + let changelog_path = root.join("CHANGELOG.md"); + let changelog = fs::read_to_string(&changelog_path) + .map_err(|err| format!("failed to read {}: {err}", changelog_path.display()))?; + let changelog = turn_over_changelog(&changelog, &target_sdk, release_date)?; + let core_changelog_path = root.join("crates/sdk-core/CHANGELOG.md"); + let core_changelog = fs::read_to_string(&core_changelog_path) + .map_err(|err| format!("failed to read {}: {err}", core_changelog_path.display()))?; + let core_changelog = turn_over_changelog(&core_changelog, &target_core_protos, release_date)?; + + let sdk_update = vec![ + "set-version".into(), + "--workspace".into(), + target_sdk.to_string(), + "--exclude".into(), + CORE_CRATE.into(), + "--exclude".into(), + BRIDGE_CRATE.into(), + "--exclude".into(), + PROTOS_CRATE.into(), + "--exclude".into(), + RELEASE_TOOL_CRATE.into(), + ]; + let core_update = vec![ + "set-version".into(), + "--package".into(), + CORE_CRATE.into(), + target_core_protos.to_string(), + ]; + let protos_update = vec![ + "set-version".into(), + "--package".into(), + PROTOS_CRATE.into(), + target_core_protos.to_string(), + ]; + let mut sdk_dry_run = sdk_update.clone(); + sdk_dry_run.push("--dry-run".into()); + let mut core_dry_run = core_update.clone(); + core_dry_run.push("--dry-run".into()); + let mut protos_dry_run = protos_update.clone(); + protos_dry_run.push("--dry-run".into()); + cargo_set_version(&root, &sdk_dry_run)?; + cargo_set_version(&root, &core_dry_run)?; + cargo_set_version(&root, &protos_dry_run)?; + + cargo_set_version(&root, &sdk_update)?; + cargo_set_version(&root, &core_update)?; + cargo_set_version(&root, &protos_update)?; + fs::write(&changelog_path, changelog) + .map_err(|err| format!("failed to write {}: {err}", changelog_path.display()))?; + fs::write(&core_changelog_path, core_changelog) + .map_err(|err| format!("failed to write {}: {err}", core_changelog_path.display()))?; + + println!("Prepared SDK {target_sdk} with Core and Protos {target_core_protos}"); + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn version(value: &str) -> Version { + Version::parse(value).unwrap() + } + + #[test] + fn rejects_core_protos_one_x() { + assert!( + parse_args(["1.0.0".into(), "1.0.0".into()]) + .unwrap_err() + .contains("Core/Protos version must remain 0.x") + ); + } + + #[test] + fn turns_over_a_populated_changelog() { + let input = "# Changelog\n\n## Unreleased\n\n### Added\n* Feature.\n\n## [0.7.0] - 2026-01-01\n\nOld.\n"; + assert_eq!( + turn_over_changelog( + input, + &version("1.0.0"), + NaiveDate::from_ymd_opt(2026, 8, 27).unwrap() + ) + .unwrap(), + "# Changelog\n\n## Unreleased\n\n## [1.0.0] - 2026-08-27\n\n### Added\n* Feature.\n\n## [0.7.0] - 2026-01-01\n\nOld.\n" + ); + } + + #[test] + fn turns_over_an_empty_changelog() { + let input = "# Changelog\n\n## Unreleased\n\n## [0.7.0] - 2026-01-01\n"; + assert_eq!( + turn_over_changelog( + input, + &version("1.0.0-rc.1"), + NaiveDate::from_ymd_opt(2026, 8, 27).unwrap() + ) + .unwrap(), + "# Changelog\n\n## Unreleased\n\n## [1.0.0-rc.1] - 2026-08-27\n\n## [0.7.0] - 2026-01-01\n" + ); + } + + #[test] + fn rejects_a_duplicate_changelog_release() { + let input = "# Changelog\n\n## Unreleased\n\n## [1.0.0] - 2026-01-01\n"; + assert!( + turn_over_changelog( + input, + &version("1.0.0"), + NaiveDate::from_ymd_opt(2026, 8, 27).unwrap() + ) + .unwrap_err() + .contains("already contains") + ); + } +} diff --git a/crates/sdk-core/src/changelog_release_notes.rs b/crates/changelog-release-notes/src/main.rs similarity index 91% rename from crates/sdk-core/src/changelog_release_notes.rs rename to crates/changelog-release-notes/src/main.rs index b8af885ac..884d0f1f1 100644 --- a/crates/sdk-core/src/changelog_release_notes.rs +++ b/crates/changelog-release-notes/src/main.rs @@ -141,8 +141,8 @@ fn update_entries(previous: &Entries, current: BTreeMap> updated } -fn changelog_notes(from: &str, to: &str) -> Result, String> { - let mut entries: Entries = changelog_entries(&git(&["show", &format!("{from}:CHANGELOG.md")])?) +fn changelog_notes(from: &str, to: &str, path: &str) -> Result, String> { + let mut entries: Entries = changelog_entries(&git(&["show", &format!("{from}:{path}")])?) .into_iter() .map(|(header, entries)| { ( @@ -163,12 +163,12 @@ fn changelog_notes(from: &str, to: &str) -> Result, String> { "--reverse", &format!("{from}..{to}"), "--", - "CHANGELOG.md", + path, ])?; for commit in commits.lines().filter(|commit| !commit.is_empty()) { entries = update_entries( &entries, - changelog_entries(&git(&["show", &format!("{commit}:CHANGELOG.md")])?), + changelog_entries(&git(&["show", &format!("{commit}:{path}")])?), ); } let mut categorized: Entries = Entries::new(); @@ -225,7 +225,7 @@ fn link_prs(subject: &str) -> String { output } -fn release_notes(from: &str, to: &str) -> Result, String> { +fn release_notes(from: &str, to: &str, changelog: &str) -> Result, String> { let log = git(&[ "log", "--format=%H%x00%h%x00%s", @@ -235,7 +235,7 @@ fn release_notes(from: &str, to: &str) -> Result, String> { if log.is_empty() { return Ok(Vec::new()); } - let mut output = changelog_notes(from, to)?; + let mut output = changelog_notes(from, to, changelog)?; if !output.is_empty() { output.insert(0, String::new()); output.insert(0, "#### Changelog".into()); @@ -253,6 +253,14 @@ fn release_notes(from: &str, to: &str) -> Result, String> { Ok(output) } +fn changelog_path(changelog: &str) -> Result<&'static str, String> { + match changelog { + "rust" => Ok("CHANGELOG.md"), + "core" => Ok("crates/sdk-core/CHANGELOG.md"), + _ => Err("expected --changelog ".into()), + } +} + fn main() -> Result<(), String> { let mut args = env::args().skip(1); let from = args @@ -265,7 +273,15 @@ fn main() -> Result<(), String> { .filter(|arg| arg == "--to") .and_then(|_| args.next()) .ok_or("expected --to ")?; - println!("{}", release_notes(&from, &to)?.join("\n")); + let changelog = match args.next().as_deref() { + None => "core".to_owned(), + Some("--changelog") => args.next().ok_or("expected --changelog ")?, + Some(_) => return Err("expected --changelog ".into()), + }; + println!( + "{}", + release_notes(&from, &to, changelog_path(&changelog)?)?.join("\n") + ); Ok(()) } @@ -305,6 +321,15 @@ mod tests { assert_eq!(clean_subject(":boom: Change"), "Change"); } + #[test] + fn selects_the_requested_changelog() { + assert_eq!(changelog_path("rust").unwrap(), "CHANGELOG.md"); + assert_eq!( + changelog_path("core").unwrap(), + "crates/sdk-core/CHANGELOG.md" + ); + } + #[test] fn keeps_final_wording_for_introduced_entry() { let previous = Entries::from([( diff --git a/crates/client/Cargo.toml b/crates/client/Cargo.toml index 7659bf876..5ac51ff51 100644 --- a/crates/client/Cargo.toml +++ b/crates/client/Cargo.toml @@ -1,7 +1,8 @@ [package] name = "temporalio-client" -version = "0.6.0" +version = "1.0.0" edition = "2024" +rust-version = "1.88.0" authors = ["Temporal Technologies Inc. "] license-file = { workspace = true } description = "Clients for interacting with Temporal" @@ -11,20 +12,25 @@ keywords = ["temporal", "workflow"] categories = ["development-tools"] readme = "README.md" +[package.metadata.docs.rs] +features = ["experimental"] + [features] default = ["envconfig", "tls-ring"] +experimental = [] tls-ring = ["tonic/tls-ring"] tls-aws-lc = ["tonic/tls-aws-lc"] telemetry = ["dep:opentelemetry"] core-based-sdk = [] envconfig = ["temporalio-common/envconfig"] dynamic-tls = ["dep:rustls-native-certs"] +vendored-protox = ["temporalio-common/vendored-protox"] [dependencies] anyhow = "1.0" async-trait = "0.1" backon = { version = "1.6", default-features = false } -base64 = "0.22" +base64 = "0.23" bon = { version = "3", default-features = false, features = ["alloc"] } derive_more = { workspace = true } dyn-clone = "1.0" @@ -45,7 +51,11 @@ tokio = { version = "1.47", default-features = false, features = [ "sync", "time", ] } -tonic = { workspace = true, default-features = false, features = ["tls-native-roots", "channel", "gzip"] } +tonic = { workspace = true, default-features = false, features = [ + "tls-native-roots", + "channel", + "gzip", +] } tokio-rustls = { version = "0.26", default-features = false } rustls-native-certs = { version = "0.8", optional = true } tower = { version = "0.5", features = ["util"] } @@ -57,17 +67,21 @@ serde_json = { workspace = true } [dependencies.temporalio-common] path = "../common" -version = "0.6" +version = "~1.0.0" default-features = false features = ["serde_serialize"] [dev-dependencies] assert_matches = "1" +hyper = { version = "1.7.0", features = ["http1", "server"] } +hyper-util = { version = "0.1.16", features = ["http1", "server", "tokio"] } mockall = "0.14" prost = "0.14" prost-types = { workspace = true } rstest = "0.26" tempfile = "3" +temporalio-macros = { path = "../macros", version = "~1.0.0" } +temporalio-workflow = { path = "../workflow", version = "~1.0.0" } tokio = { version = "1.47", default-features = false, features = [ "io-util", "macros", @@ -76,6 +90,9 @@ tokio = { version = "1.47", default-features = false, features = [ "sync", "time", ] } +tokio-stream = { version = "0.1", default-features = false, features = ["net"] } +tonic = { workspace = true, default-features = false, features = ["router", "server"] } +trybuild = { version = "1.0", features = ["diff"] } [lints] workspace = true diff --git a/crates/client/README.md b/crates/client/README.md index 749455f5a..495b6fdf3 100644 --- a/crates/client/README.md +++ b/crates/client/README.md @@ -7,9 +7,6 @@ This crate provides a Rust client for interacting with the Temporal service. It standalone to start and manage workflows, or together with the [`temporalio-sdk`](https://crates.io/crates/temporalio-sdk) crate to run workers. -⚠️ **This crate is in Public Preview and under active development.** The API can and -will continue to evolve. - ## Quick Start ### Connecting and Starting Workflows with Environment Configuration @@ -162,6 +159,16 @@ while let Some(result) = stream.next().await { } ``` +## Experimental APIs + +APIs that are still under development require the `experimental` Cargo feature and may change or +be removed before stabilization. + +## Building Without `protoc` + +Enable the `vendored-protox` Cargo feature to compile protobuf definitions with `protox` +instead of a system `protoc` binary. + ## Raw gRPC Access For operations not covered by the high-level API, access the underlying gRPC service clients diff --git a/crates/client/src/activity.rs b/crates/client/src/activity.rs new file mode 100644 index 000000000..eb250c8f0 --- /dev/null +++ b/crates/client/src/activity.rs @@ -0,0 +1,142 @@ +mod activity_execution_info; +mod activity_handle; + +use crate::errors::ClientError; +pub use activity_execution_info::{ + ActivityExecutionDescription, ActivityExecutionInfo, ActivityExecutionInfoLike, + ActivityExecutionStatus, PendingActivityState, +}; +pub use activity_handle::ActivityHandle; +use futures_util::{Stream, StreamExt}; +use std::{ + collections::VecDeque, + pin::Pin, + task::{Context, Poll}, +}; +use temporalio_common::{ + protos::temporal::api::{ + activity::v1::ActivityExecutionListInfo, + workflowservice::v1::{ + CountActivityExecutionsResponse, count_activity_executions_response, + }, + }, + search_attributes::{SearchAttributeError, SearchAttributeValue}, +}; + +/// A stream of activity executions from a list query. +/// Internally paginates through results from the server. +pub struct ListActivitiesStream { + inner: Pin, ClientError>> + Send>>, + buffer: VecDeque, +} + +impl ListActivitiesStream { + pub(crate) fn new( + stream: impl Stream, ClientError>> + Send + 'static, + ) -> Self { + Self { + inner: Box::pin(stream), + buffer: VecDeque::new(), + } + } +} + +impl Stream for ListActivitiesStream { + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + loop { + if let Some(info) = self.buffer.pop_front() { + return Poll::Ready(Some(Ok(info.into()))); + } + match self.inner.poll_next_unpin(cx) { + Poll::Ready(Some(Ok(items))) => { + self.buffer = items.into(); + } + Poll::Ready(Some(Err(e))) => { + return Poll::Ready(Some(Err(e))); + } + Poll::Ready(None) => { + return Poll::Ready(None); + } + Poll::Pending => { + return Poll::Pending; + } + } + } + } +} + +/// Result of an activity count operation. +/// +/// If the query includes a group-by clause, `groups` will contain the aggregated +/// counts and `count` will be the sum of all group counts. +#[derive(Debug, Clone)] +pub struct ActivityExecutionCount { + count: usize, + groups: Vec, +} + +impl ActivityExecutionCount { + pub(crate) fn from_response(resp: CountActivityExecutionsResponse) -> Self { + Self { + count: resp.count as usize, + groups: resp + .groups + .into_iter() + .map(ActivityExecutionCountAggregationGroup::from_proto) + .collect(), + } + } + + /// The approximate number of activities matching the query. + /// If grouping was applied, this is the sum of all group counts. + pub fn count(&self) -> usize { + self.count + } + + /// The groups if the query had a group-by clause, or empty if not. + pub fn groups(&self) -> &[ActivityExecutionCountAggregationGroup] { + &self.groups + } +} + +/// Aggregation group from an activity count query with a group-by clause. +#[derive(Debug, Clone)] +pub struct ActivityExecutionCountAggregationGroup { + raw: count_activity_executions_response::AggregationGroup, +} + +impl ActivityExecutionCountAggregationGroup { + fn from_proto(proto: count_activity_executions_response::AggregationGroup) -> Self { + Self { raw: proto } + } + + /// Retrieve a typed group value at `index`. + /// + /// Returns `None` if the index is out of bounds or deserialization fails. + /// Use [`Self::try_get`] for explicit error handling. + pub fn get(&self, index: usize) -> Option { + self.try_get(index).ok().flatten() + } + + /// Retrieve a typed group value at `index`, preserving deserialization + /// errors. + /// + /// Returns `Ok(None)` if the index is out of bounds and `Err` if the + /// payload cannot be deserialized. + pub fn try_get( + &self, + index: usize, + ) -> Result, SearchAttributeError> { + match self.raw.group_values.get(index) { + Some(payload) => T::from_search_attribute_payload(payload).map(Some), + None => Ok(None), + } + } + + /// The approximate number of workflows matching for this group. + pub fn count(&self) -> usize { + self.raw.count as usize + } +} diff --git a/crates/client/src/activity/activity_execution_info.rs b/crates/client/src/activity/activity_execution_info.rs new file mode 100644 index 000000000..21f4c2220 --- /dev/null +++ b/crates/client/src/activity/activity_execution_info.rs @@ -0,0 +1,579 @@ +use crate::Priority; +use std::{ + error::Error, + marker::PhantomData, + time::{Duration, SystemTime}, +}; +use temporalio_common::{ + ActivityDefinition, RetryPolicy, UntypedActivity, WorkerDeploymentVersion, + data_converters::{ + DataConverter, NoopDecodeHint, PayloadConversionError, SerializationContextData, + TemporalDeserializable, + }, + error::IncomingError, + payload_visitor::decode_payloads, + protos::{ + proto_ts_to_system_time, + temporal::api::{ + activity::v1::{ + ActivityExecutionInfo as RawInfo, ActivityExecutionListInfo as RawListInfo, + activity_execution_outcome::Value as ActivityExecutionOutcomeValue, + }, + common::v1::{Payload, Payloads}, + enums::v1::{ + ActivityExecutionStatus as ProtoActivityExecutionStatus, + PendingActivityState as ProtoPendingActivityState, + }, + failure::v1::Failure, + workflowservice::v1::DescribeActivityExecutionResponse, + }, + utilities::TryIntoOrNone, + }, + search_attributes::SearchAttributes, +}; + +/// Common methods of [`ActivityExecutionInfo`] and [`ActivityExecutionDescription`]. +pub trait ActivityExecutionInfoLike { + /// ID of the activity. + fn activity_id(&self) -> &str; + /// Run ID of a particular execution of the activity. + fn activity_run_id(&self) -> &str; + /// Type of the activity. + fn activity_type(&self) -> &str; + /// Time the activity was originally scheduled. + fn schedule_time(&self) -> Option; + /// Time when the activity transitioned to a closed state. + fn close_time(&self) -> Option; + /// A general status for this activity, indicates whether it is currently running or in one of + /// the terminal statuses. + fn status(&self) -> ActivityExecutionStatus; + /// The task queue this activity was scheduled on. + fn task_queue(&self) -> &str; + /// The difference between close time and scheduled time. This field is only populated if + /// the activity is closed. + fn execution_duration(&self) -> Option; +} + +/// Contains basic information about an activity. +/// Obtained from [`Client::list_activities`](crate::Client::list_activities). +pub struct ActivityExecutionInfo { + raw: RawListInfo, +} + +impl From for ActivityExecutionInfo { + fn from(raw: RawListInfo) -> Self { + Self { raw } + } +} + +impl ActivityExecutionInfoLike for ActivityExecutionInfo { + fn activity_id(&self) -> &str { + &self.raw.activity_id + } + + fn activity_run_id(&self) -> &str { + &self.raw.run_id + } + + fn activity_type(&self) -> &str { + self.raw + .activity_type + .as_ref() + .map(|t| t.name.as_str()) + .unwrap_or("") + } + + fn schedule_time(&self) -> Option { + self.raw + .schedule_time + .as_ref() + .and_then(proto_ts_to_system_time) + } + + fn close_time(&self) -> Option { + self.raw + .close_time + .as_ref() + .and_then(proto_ts_to_system_time) + } + + fn status(&self) -> ActivityExecutionStatus { + ProtoActivityExecutionStatus::try_from(self.raw.status) + .map(Into::into) + .unwrap_or(ActivityExecutionStatus::Unknown) + } + + fn task_queue(&self) -> &str { + &self.raw.task_queue + } + + fn execution_duration(&self) -> Option { + self.raw.execution_duration.try_into_or_none() + } +} + +impl ActivityExecutionInfo { + /// Raw Protobuf object from server response. + pub fn raw_info(&self) -> &RawListInfo { + &self.raw + } +} + +/// Contains the current state of the activity execution. +/// Obtained from [`ActivityHandle::describe`](crate::ActivityHandle::describe). +/// Methods that deserialize payloads (e.g. [`heartbeat_details`](Self::heartbeat_details)) use +/// [`DataConverter`] of the client associated with the activity handle. +pub struct ActivityExecutionDescription +where + ActivityT: ActivityDefinition, +{ + raw_info: RawInfo, + raw_input: Option, + raw_outcome: Option, + data_converter: DataConverter, + serialization_context: SerializationContextData, + _phantom: PhantomData, +} + +impl ActivityExecutionInfoLike for ActivityExecutionDescription +where + ActivityT: ActivityDefinition, +{ + fn activity_id(&self) -> &str { + &self.raw_info.activity_id + } + + fn activity_run_id(&self) -> &str { + &self.raw_info.run_id + } + + fn activity_type(&self) -> &str { + self.raw_info + .activity_type + .as_ref() + .map(|t| t.name.as_str()) + .unwrap_or("") + } + + fn schedule_time(&self) -> Option { + self.raw_info + .schedule_time + .as_ref() + .and_then(proto_ts_to_system_time) + } + + fn close_time(&self) -> Option { + self.raw_info + .close_time + .as_ref() + .and_then(proto_ts_to_system_time) + } + + fn status(&self) -> ActivityExecutionStatus { + ProtoActivityExecutionStatus::try_from(self.raw_info.status) + .map(Into::into) + .unwrap_or(ActivityExecutionStatus::Unknown) + } + + fn task_queue(&self) -> &str { + &self.raw_info.task_queue + } + + fn execution_duration(&self) -> Option { + self.raw_info.execution_duration.try_into_or_none() + } +} + +impl ActivityExecutionDescription +where + ActivityT: ActivityDefinition, +{ + pub(crate) async fn new( + data_converter: DataConverter, + serialization_context: SerializationContextData, + response: DescribeActivityExecutionResponse, + ) -> Result> { + let Some(mut raw_info) = response.info else { + return Err("info missing in describe response".into()); + }; + if let Some(failure) = raw_info.last_failure.as_mut() { + decode_payloads(failure, data_converter.codec(), &serialization_context).await?; + } + let mut raw_outcome = response.outcome.and_then(|o| o.value); + if let Some(ActivityExecutionOutcomeValue::Failure(failure)) = raw_outcome.as_mut() { + decode_payloads(failure, data_converter.codec(), &serialization_context).await?; + } + Ok(Self { + raw_info, + raw_input: response.input, + raw_outcome, + data_converter, + serialization_context, + _phantom: PhantomData, + }) + } + + /// Convert to an untyped description object. + pub fn untyped(self) -> ActivityExecutionDescription { + ActivityExecutionDescription { + raw_info: self.raw_info, + raw_input: self.raw_input, + raw_outcome: self.raw_outcome, + data_converter: self.data_converter, + serialization_context: self.serialization_context, + _phantom: PhantomData, + } + } + + /// Raw Protobuf object from server response. + pub fn raw_info(&self) -> &RawInfo { + &self.raw_info + } + + /// True if activity input is present. + /// See [`ActivityDescribeOptions::include_input`](crate::ActivityDescribeOptions::include_input). + /// Use [`input`](Self::input) or [`raw_input`](Self::raw_input) to retrieve it. + pub fn has_input(&self) -> bool { + self.raw_input.is_some() + } + + /// Raw payload of activity input, if it was requested. + pub fn raw_input(&self) -> Option<&Payloads> { + self.raw_input.as_ref() + } + + /// Deserialize activity input. Returns `Ok(None)` if not present. + /// See [`ActivityDescribeOptions::include_input`](crate::ActivityDescribeOptions::include_input). + pub async fn input(&self) -> Result, PayloadConversionError> { + let Some(input) = &self.raw_input else { + return Ok(None); + }; + Ok(Some(self.convert_payloads(input).await?)) + } + + /// True if activity outcome is present. + /// See [`ActivityDescribeOptions::include_outcome`](crate::ActivityDescribeOptions::include_outcome). + /// Use [`outcome`](Self::outcome) or [`raw_outcome`](Self::outcome) to retrieve it. + pub fn has_outcome(&self) -> bool { + self.raw_outcome.is_some() + } + + /// Raw payload of activity output, if it was requested and available. + pub fn raw_outcome(&self) -> Option<&ActivityExecutionOutcomeValue> { + self.raw_outcome.as_ref() + } + + /// Deserialize activity outcome. Returns `Ok(None)` if not present. + /// See [`ActivityDescribeOptions::include_outcome`](crate::ActivityDescribeOptions::include_outcome). + pub async fn outcome( + &self, + ) -> Result>, PayloadConversionError> { + match &self.raw_outcome { + None => Ok(None), + Some(ActivityExecutionOutcomeValue::Result(payloads)) => { + Ok(Some(Ok(self.convert_payloads(payloads).await?))) + } + Some(ActivityExecutionOutcomeValue::Failure(failure)) => { + Ok(Some(Err(self.convert_failure(failure)?))) + } + } + } + + /// More detailed breakdown of [`ActivityExecutionStatus::Running`]. + pub fn run_state(&self) -> PendingActivityState { + ProtoPendingActivityState::try_from(self.raw_info.run_state) + .map(Into::into) + .unwrap_or(PendingActivityState::Unknown) + } + + /// Indicates how long the caller is willing to wait for an activity completion. Limits how long + /// retries will be attempted. + pub fn schedule_to_close_timeout(&self) -> Option { + self.raw_info.schedule_to_close_timeout.try_into_or_none() + } + + /// Limits time an activity task can stay in a task queue before a worker picks it up. This + /// timeout is always non-retryable. + pub fn schedule_to_start_timeout(&self) -> Option { + self.raw_info.schedule_to_start_timeout.try_into_or_none() + } + + /// Maximum time a single activity attempt is allowed to execute after being picked up by + /// a worker. This timeout is always retryable. + pub fn start_to_close_timeout(&self) -> Option { + self.raw_info.start_to_close_timeout.try_into_or_none() + } + + /// Maximum permitted time between successful worker heartbeats. + pub fn heartbeat_timeout(&self) -> Option { + self.raw_info.heartbeat_timeout.try_into_or_none() + } + + /// The retry policy for the activity. + pub fn retry_policy(&self) -> Option { + self.raw_info.retry_policy.clone().map(Into::into) + } + + /// True if heartbeat details are present. + /// See [`ActivityDescribeOptions::include_heartbeat_details`](crate::ActivityDescribeOptions::include_heartbeat_details). + /// Use [`heartbeat_details`](Self::heartbeat_details) or + /// [`raw_info()`](Self::raw_info)`.`[`heartbeat_details`](RawInfo::heartbeat_details) + /// to retrieve them. + pub fn has_heartbeat_details(&self) -> bool { + self.raw_info.heartbeat_details.is_some() + } + + /// Deserialize heartbeat details. Returns `Ok(None)` if not present. + /// See [`ActivityDescribeOptions::include_heartbeat_details`](crate::ActivityDescribeOptions::include_heartbeat_details). + pub async fn heartbeat_details( + &self, + ) -> Result, PayloadConversionError> { + let Some(details) = &self.raw_info.heartbeat_details else { + return Ok(None); + }; + Ok(Some(self.convert_payloads(details).await?)) + } + + /// Time the last heartbeat was recorded. + pub fn last_heartbeat_time(&self) -> Option { + self.raw_info + .last_heartbeat_time + .as_ref() + .and_then(proto_ts_to_system_time) + } + + /// Time the last attempt was started. + pub fn last_started_time(&self) -> Option { + self.raw_info + .last_started_time + .as_ref() + .and_then(proto_ts_to_system_time) + } + + /// The attempt this activity is currently on. Incremented each time a new attempt is scheduled. + pub fn attempt(&self) -> u32 { + self.raw_info.attempt.try_into().unwrap_or_default() + } + + /// How long this activity has been running for, including all attempts and backoff between + /// attempts. + pub fn execution_duration(&self) -> Option { + self.raw_info.execution_duration.try_into_or_none() + } + + /// Scheduled time + schedule to close timeout. + pub fn expiration_time(&self) -> Option { + self.raw_info + .expiration_time + .as_ref() + .and_then(proto_ts_to_system_time) + } + + /// True if last failure is present. + /// See [`ActivityDescribeOptions::include_last_failure`](crate::ActivityDescribeOptions::include_last_failure). + /// Use [`last_failure()`](Self::last_failure) or + /// [`raw_info()`](Self::raw_info)`.`[`last_failure`](RawInfo::last_failure) + /// to retrieve it. + pub fn has_last_failure(&self) -> bool { + self.raw_info.last_failure.is_some() + } + + /// Deserialize last failure. Returns `Ok(None)` if not present. + /// See [`ActivityDescribeOptions::include_last_failure`](crate::ActivityDescribeOptions::include_last_failure). + pub fn last_failure(&self) -> Result, PayloadConversionError> { + let Some(failure) = &self.raw_info.last_failure else { + return Ok(None); + }; + Ok(Some(self.convert_failure(failure)?)) + } + + /// Identity of the last worker that attempted this activity. + pub fn last_worker_identity(&self) -> Option<&str> { + self.raw_info + .last_worker_identity + .is_empty() + .then_some(self.raw_info.last_worker_identity.as_str()) + } + + /// Time from the last attempt failure to the next activity retry. + pub fn current_retry_interval(&self) -> Option { + self.raw_info.current_retry_interval.try_into_or_none() + } + + /// The time when the last activity attempt completed. + pub fn last_attempt_complete_time(&self) -> Option { + self.raw_info + .last_attempt_complete_time + .as_ref() + .and_then(proto_ts_to_system_time) + } + + /// The time when the next activity attempt will be scheduled. + pub fn next_attempt_schedule_time(&self) -> Option { + self.raw_info + .next_attempt_schedule_time + .as_ref() + .and_then(proto_ts_to_system_time) + } + + /// The Worker Deployment Version this activity was dispatched to most recently. + pub fn last_deployment_version(&self) -> Option { + self.raw_info + .last_deployment_version + .clone() + .map(Into::into) + } + + /// Priority metadata. + pub fn priority(&self) -> Priority { + self.raw_info.priority.clone().unwrap_or_default().into() + } + + /// Search attributes of the activity. + pub fn search_attributes(&self) -> Option { + self.raw_info + .search_attributes + .as_ref() + .map(SearchAttributes::from_proto) + } + + /// Deserialize static summary that was set when activity was scheduled. + /// Returns `Ok(None)` if not present. + pub async fn static_summary(&self) -> Result, PayloadConversionError> { + let Some(summary) = self + .raw_info + .user_metadata + .as_ref() + .and_then(|m| m.summary.clone()) + else { + return Ok(None); + }; + Ok(Some(self.convert_payload(summary).await?)) + } + + /// Deserialize static details that were set when activity was scheduled. + /// Returns `Ok(None)` if not present. + pub async fn static_details(&self) -> Result, PayloadConversionError> { + let Some(details) = self + .raw_info + .user_metadata + .as_ref() + .and_then(|m| m.details.clone()) + else { + return Ok(None); + }; + Ok(Some(self.convert_payload(details).await?)) + } + + /// Reason for activity cancellation if activity was canceled and reason was provided. + pub fn canceled_reason(&self) -> Option<&str> { + let reason = self.raw_info.canceled_reason.as_str(); + (!reason.is_empty()).then_some(reason) + } + + /// Time to wait before dispatching the first activity task. + /// This delay is not applied to retry attempts. + pub fn start_delay(&self) -> Option { + self.raw_info.start_delay.try_into_or_none() + } + + async fn convert_payload( + &self, + payload: Payload, + ) -> Result { + self.data_converter + .from_payload(&self.serialization_context, payload) + .await + } + + async fn convert_payloads( + &self, + payloads: &Payloads, + ) -> Result { + self.data_converter + .from_payloads(&self.serialization_context, payloads.payloads.clone()) + .await + } + + fn convert_failure(&self, failure: &Failure) -> Result { + self.data_converter + .to_error(&self.serialization_context, failure.clone(), NoopDecodeHint) + } +} + +/// Execution status of an activity. See [`ActivityExecutionInfoLike::status`]. +#[non_exhaustive] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)] +pub enum ActivityExecutionStatus { + #[default] + /// This variant indicates the server did not specify a value. + Unspecified, + /// The activity has not reached a terminal status. + /// See [`ActivityExecutionDescription::run_state`] for the run state. + Running, + /// The activity completed successfully. + Completed, + /// The activity failed with an error. + Failed, + /// The activity was canceled. Note that cancellation is cooperative and a cancel request does + /// not always result in canceled status. + Canceled, + /// The activity was terminated. + Terminated, + /// The activity timed out. + TimedOut, + /// The activity is paused. + Paused, + /// This variant indicates the server used a value not known by this version of the SDK. + Unknown, +} + +impl From for ActivityExecutionStatus { + fn from(value: ProtoActivityExecutionStatus) -> Self { + match value { + ProtoActivityExecutionStatus::Unspecified => Self::Unspecified, + ProtoActivityExecutionStatus::Running => Self::Running, + ProtoActivityExecutionStatus::Completed => Self::Completed, + ProtoActivityExecutionStatus::Failed => Self::Failed, + ProtoActivityExecutionStatus::Canceled => Self::Canceled, + ProtoActivityExecutionStatus::Terminated => Self::Terminated, + ProtoActivityExecutionStatus::TimedOut => Self::TimedOut, + ProtoActivityExecutionStatus::Paused => Self::Paused, + } + } +} + +/// Detailed state of an activity with [`ActivityExecutionStatus::Running`]. +/// See [`ActivityExecutionDescription::run_state`]. +#[non_exhaustive] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)] +pub enum PendingActivityState { + #[default] + /// This variant indicates the server did not specify a state. + Unspecified, + /// Activity is scheduled for execution but not yet running on a worker. + Scheduled, + /// Activity is running on a worker. + Started, + /// Activity has been requested to cancel. + CancelRequested, + /// Activity is paused on the server, and is not running on a worker. + Paused, + /// Activity is currently running on a worker, but paused on the server. + PauseRequested, + /// This variant indicates the server used a value not known by this version of the SDK. + Unknown, +} + +impl From for PendingActivityState { + fn from(value: ProtoPendingActivityState) -> Self { + match value { + ProtoPendingActivityState::Unspecified => Self::Unspecified, + ProtoPendingActivityState::Scheduled => Self::Scheduled, + ProtoPendingActivityState::Started => Self::Started, + ProtoPendingActivityState::CancelRequested => Self::CancelRequested, + ProtoPendingActivityState::Paused => Self::Paused, + ProtoPendingActivityState::PauseRequested => Self::PauseRequested, + } + } +} diff --git a/crates/client/src/activity/activity_handle.rs b/crates/client/src/activity/activity_handle.rs new file mode 100644 index 000000000..e8d51af3a --- /dev/null +++ b/crates/client/src/activity/activity_handle.rs @@ -0,0 +1,350 @@ +use crate::{ + ActivityCancelOptions, ActivityDescribeOptions, ActivityExecutionDescription, + ActivityTerminateOptions, NamespacedClient, + errors::{ActivityInteractionError, ActivityResultError}, + grpc::WorkflowService, +}; +use std::marker::PhantomData; +use temporalio_common::{ + ActivityDefinition, + data_converters::{ + ActivitySerializationContext, DecodablePayloads, NoopDecodeHint, SerializationContextData, + }, + payload_visitor::decode_payloads, + protos::temporal::api::{ + activity::v1::{ActivityExecutionOutcome, activity_execution_outcome}, + failure::v1::failure::FailureInfo, + workflowservice::v1::{ + DescribeActivityExecutionRequest, PollActivityExecutionRequest, + RequestCancelActivityExecutionRequest, TerminateActivityExecutionRequest, + }, + }, +}; +use tonic::IntoRequest; +use uuid::Uuid; + +/// Handle associated with a standalone activity execution that can be used to wait for the result +/// or to manage execution of the activity. Obtained from +/// [`Client::start_activity`](crate::Client::start_activity) or +/// [`Client::get_activity_handle`](crate::Client::get_activity_handle). +/// +/// If [`run_id`](Self::run_id) is set, the handle always targets that specific execution. +/// If [`run_id`](Self::run_id) is `None`, each method call targets the latest run of the specified +/// [`activity_id`](Self::activity_id) at the time the method is called - this means consecutive +/// method calls may target different executions if an activity was started again with the same ID. +pub struct ActivityHandle +where + ActivityT: ActivityDefinition, +{ + client: ClientT, + activity_id: String, + run_id: Option, + _phantom: PhantomData, +} + +impl ActivityHandle +where + ActivityT: ActivityDefinition, +{ + pub(crate) fn new(client: ClientT, activity_id: String, run_id: Option) -> Self { + Self { + client, + activity_id, + run_id, + _phantom: PhantomData, + } + } + + /// Activity ID this handle is associated with. + pub fn activity_id(&self) -> &str { + &self.activity_id + } + + /// Run ID of the activity execution this handle is associated with. If `None`, each method call + /// targets the latest run of the specified [`activity_id`](Self::activity_id) at the time the + /// method is called - this means consecutive method calls may target different executions if + /// an activity was started again with the same ID. + pub fn run_id(&self) -> Option<&str> { + self.run_id.as_deref() + } +} + +impl ActivityHandle +where + ClientT: WorkflowService + NamespacedClient + Clone, + ActivityT: ActivityDefinition, +{ + /// Wait for the activity to complete and fetch its result. If the activity was not successful + /// (e.g. failed, canceled, timed out), this method returns [`ActivityResultError::ActivityFailed`]. + pub async fn result(&self) -> Result { + let mut client = self.client.clone(); + loop { + let resp = client + .poll_activity_execution( + PollActivityExecutionRequest { + namespace: client.namespace(), + activity_id: self.activity_id.clone(), + run_id: self.run_id.clone().unwrap_or_default(), + } + .into_request(), + ) + .await? + .into_inner(); + + // If resp.outcome.value is None, poll again + let Some(ActivityExecutionOutcome { + value: Some(outcome), + .. + }) = resp.outcome + else { + continue; + }; + + let dc = client.data_converter(); + let ctx = SerializationContextData::Activity(ActivitySerializationContext::new()); + + return match outcome { + activity_execution_outcome::Value::Result(payloads) => { + Ok(dc.from_payloads(&ctx, payloads.payloads).await?) + } + activity_execution_outcome::Value::Failure(mut failure) => { + decode_payloads(&mut failure, dc.codec(), &ctx).await?; + Err(match failure.failure_info { + Some(FailureInfo::CanceledFailureInfo(info)) => { + let payloads = info.details.unwrap_or_default().payloads; + let details = DecodablePayloads::new( + payloads, + dc.payload_converter().clone(), + ctx, + ); + ActivityResultError::Cancelled { details } + } + Some(FailureInfo::TerminatedFailureInfo(_)) => { + ActivityResultError::Terminated + } + _ => ActivityResultError::ActivityFailed(dc.to_error( + &ctx, + failure, + NoopDecodeHint, + )?), + }) + } + }; + } + } + + /// Describes the current state of the activity execution. + pub async fn describe( + &self, + options: ActivityDescribeOptions, + ) -> Result, ActivityInteractionError> { + let mut client = self.client.clone(); + let resp = client + .describe_activity_execution( + DescribeActivityExecutionRequest { + namespace: client.namespace(), + activity_id: self.activity_id.clone(), + run_id: self.run_id.clone().unwrap_or_default(), + include_input: options.include_input, + include_outcome: options.include_outcome, + include_heartbeat_details: options.include_heartbeat_details, + include_last_failure: options.include_last_failure, + ..Default::default() + } + .into_request(), + ) + .await? + .into_inner(); + + Ok(ActivityExecutionDescription::new( + client.data_converter().clone(), + SerializationContextData::Activity(ActivitySerializationContext::new()), + resp, + ) + .await?) + } + + /// Requests cancellation of the activity. Does not wait for the cancellation to complete. + pub async fn cancel( + &self, + options: ActivityCancelOptions, + ) -> Result<(), ActivityInteractionError> { + let mut client = self.client.clone(); + client + .request_cancel_activity_execution( + RequestCancelActivityExecutionRequest { + namespace: client.namespace(), + activity_id: self.activity_id.clone(), + run_id: self.run_id.clone().unwrap_or_default(), + identity: client.identity(), + request_id: Uuid::new_v4().to_string(), + reason: options.reason, + } + .into_request(), + ) + .await?; + + Ok(()) + } + + /// Terminates activity execution. + pub async fn terminate( + &self, + options: ActivityTerminateOptions, + ) -> Result<(), ActivityInteractionError> { + let mut client = self.client.clone(); + client + .terminate_activity_execution( + TerminateActivityExecutionRequest { + namespace: client.namespace(), + activity_id: self.activity_id.clone(), + run_id: self.run_id.clone().unwrap_or_default(), + identity: client.identity(), + request_id: Uuid::new_v4().to_string(), + reason: options.reason, + } + .into_request(), + ) + .await?; + + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::test_helpers::XorCodec; + use futures_util::future::BoxFuture; + use temporalio_common::{ + UntypedActivity, + data_converters::{DataConverter, DefaultFailureConverter, PayloadConverter}, + error::{ApplicationFailure, OutgoingActivityError, OutgoingError}, + payload_visitor::encode_payloads, + protos::temporal::api::{ + activity::v1::ActivityExecutionInfo, + failure::v1::Failure, + workflowservice::v1::{ + DescribeActivityExecutionResponse, PollActivityExecutionResponse, + }, + }, + }; + use tonic::{Request, Response, Status}; + + #[derive(Clone)] + struct MockActivityClient { + data_converter: DataConverter, + failure: Failure, + } + + impl NamespacedClient for MockActivityClient { + fn namespace(&self) -> String { + "test-namespace".to_owned() + } + + fn identity(&self) -> String { + "test-identity".to_owned() + } + + fn data_converter(&self) -> &DataConverter { + &self.data_converter + } + } + + impl WorkflowService for MockActivityClient { + fn poll_activity_execution( + &mut self, + _request: Request, + ) -> BoxFuture<'_, Result, Status>> { + let failure = self.failure.clone(); + Box::pin(async move { + Ok(Response::new(PollActivityExecutionResponse { + outcome: Some(ActivityExecutionOutcome { + value: Some(activity_execution_outcome::Value::Failure(failure)), + ..Default::default() + }), + ..Default::default() + })) + }) + } + + fn describe_activity_execution( + &mut self, + _request: Request, + ) -> BoxFuture<'_, Result, Status>> { + let failure = self.failure.clone(); + Box::pin(async move { + Ok(Response::new(DescribeActivityExecutionResponse { + info: Some(ActivityExecutionInfo { + last_failure: Some(failure.clone()), + ..Default::default() + }), + outcome: Some(ActivityExecutionOutcome { + value: Some(activity_execution_outcome::Value::Failure(failure)), + ..Default::default() + }), + ..Default::default() + })) + }) + } + } + + async fn activity_client_with_encoded_failure() -> MockActivityClient { + let data_converter = DataConverter::new( + PayloadConverter::default(), + DefaultFailureConverter::new(true), + XorCodec, + ); + let context = SerializationContextData::Activity(ActivitySerializationContext::new()); + let mut failure = data_converter.to_failure( + &context, + OutgoingError::Activity(OutgoingActivityError::Application(Box::new( + ApplicationFailure::new(anyhow::anyhow!("private message")), + ))), + ); + encode_payloads(&mut failure, data_converter.codec(), &context) + .await + .unwrap(); + MockActivityClient { + data_converter, + failure, + } + } + + #[tokio::test] + async fn result_decodes_failure_attributes_with_codec() { + let handle = ActivityHandle::<_, UntypedActivity>::new( + activity_client_with_encoded_failure().await, + "activity-id".to_owned(), + None, + ); + + let ActivityResultError::ActivityFailed(error) = handle.result().await.unwrap_err() else { + panic!("expected failed activity"); + }; + assert_eq!(error.failure().message, "private message"); + } + + #[tokio::test] + async fn describe_decodes_failure_attributes_with_codec() { + let handle = ActivityHandle::<_, UntypedActivity>::new( + activity_client_with_encoded_failure().await, + "activity-id".to_owned(), + None, + ); + + let description = handle + .describe( + ActivityDescribeOptions::builder() + .include_outcome(true) + .include_last_failure(true) + .build(), + ) + .await + .unwrap(); + let outcome = description.outcome().await.unwrap().unwrap().unwrap_err(); + let last_failure = description.last_failure().unwrap().unwrap(); + assert_eq!(outcome.failure().message, "private message"); + assert_eq!(last_failure.failure().message, "private message"); + } +} diff --git a/crates/client/src/async_activity_handle.rs b/crates/client/src/async_activity_handle.rs index d0b5e2c5b..66b100b59 100644 --- a/crates/client/src/async_activity_handle.rs +++ b/crates/client/src/async_activity_handle.rs @@ -7,7 +7,10 @@ use crate::{ }; use futures_util::future::BoxFuture; use temporalio_common::{ - data_converters::{SerializationContext, SerializationContextData, TemporalSerializable}, + data_converters::{ + ActivitySerializationContext, SerializationContext, SerializationContextData, + TemporalSerializable, + }, error::{ApplicationFailure, OutgoingActivityError, OutgoingError}, payload_visitor::encode_payloads, protos::{ @@ -35,16 +38,17 @@ async fn encode_optional_value( }; let unencoded_payloads = { let payload_converter = data_converter.payload_converter(); - let context = SerializationContext { - data: &SerializationContextData::Activity, - converter: payload_converter, - }; + let context_data = SerializationContextData::Activity(ActivitySerializationContext::new()); + let context = SerializationContext::new(&context_data, payload_converter); value.serialize_payloads(&context)? }; drop(value); let payloads = data_converter .codec() - .encode(&SerializationContextData::Activity, unencoded_payloads) + .encode( + &SerializationContextData::Activity(ActivitySerializationContext::new()), + unencoded_payloads, + ) .await?; Ok(Some(Payloads { payloads })) } @@ -54,8 +58,8 @@ async fn encode_optional_value( pub enum ActivityIdentifier { /// Identify activity by its task token TaskToken(TaskToken), - /// Identify activity by workflow and activity IDs. - ById { + /// Identify workflow activity by workflow and activity IDs. + ByIdWorkflow { /// ID of the workflow that scheduled this activity. workflow_id: String, /// Run ID of the workflow (optional - if not provided, targets the latest run). @@ -63,6 +67,13 @@ pub enum ActivityIdentifier { /// ID of the activity to complete. activity_id: String, }, + /// Identify standalone activity by activity ID. + ByIdStandalone { + /// ID of the activity to complete. + activity_id: String, + /// Run ID of the activity (optional - if not provided, targets the latest run). + run_id: String, + }, } impl ActivityIdentifier { @@ -73,17 +84,42 @@ impl ActivityIdentifier { /// Create an identifier from workflow and activity IDs. Use an empty run id to target the /// latest workflow execution. - pub fn by_id( + pub fn by_id_workflow( workflow_id: impl Into, run_id: impl Into, activity_id: impl Into, ) -> Self { - Self::ById { + Self::ByIdWorkflow { workflow_id: workflow_id.into(), run_id: run_id.into(), activity_id: activity_id.into(), } } + + /// Create an identifier from standalone activity ID. Use an empty run id to target the + /// latest activity execution. + pub fn by_id_standalone(activity_id: impl Into, run_id: impl Into) -> Self { + Self::ByIdStandalone { + activity_id: activity_id.into(), + run_id: run_id.into(), + } + } + + /// Returns tuple of (workflow_id, run_id, activity_id). + fn into_parts(self) -> Option<(String, String, String)> { + match self { + Self::TaskToken(_) => None, + Self::ByIdWorkflow { + workflow_id, + run_id, + activity_id, + } => Some((workflow_id, run_id, activity_id)), + Self::ByIdStandalone { + activity_id, + run_id, + } => Some((String::new(), run_id, activity_id)), + } + } } /// Handle for completing activities asynchronously (outside the worker). @@ -131,47 +167,41 @@ impl AsyncActivityHandle { Box::pin(async move { let (identifier, result, rpc_options) = input.into_parts(); let result = encode_optional_value(result, client.data_converter()).await?; - match identifier { - ActivityIdentifier::TaskToken(token) => { - let mut request = RespondActivityTaskCompletedRequest { - task_token: token.0, - result, - identity: client.identity(), - namespace: client.namespace(), - ..Default::default() - } - .into_request(); - rpc_options.apply_to(&mut request); - WorkflowService::respond_activity_task_completed( - &mut client, - request, - ) - .await - .map_err(AsyncActivityError::from_status)?; + if let ActivityIdentifier::TaskToken(token) = identifier { + let mut request = RespondActivityTaskCompletedRequest { + task_token: token.into_inner(), + result, + identity: client.identity(), + namespace: client.namespace(), + ..Default::default() } - ActivityIdentifier::ById { + .into_request(); + rpc_options.apply_to(&mut request); + WorkflowService::respond_activity_task_completed( + &mut client, + request, + ) + .await + .map_err(AsyncActivityError::from_status)?; + } else { + let (workflow_id, run_id, activity_id) = identifier.into_parts().unwrap(); + let mut request = RespondActivityTaskCompletedByIdRequest { + namespace: client.namespace(), workflow_id, run_id, activity_id, - } => { - let mut request = RespondActivityTaskCompletedByIdRequest { - namespace: client.namespace(), - workflow_id, - run_id, - activity_id, - result, - identity: client.identity(), - resource_id: Default::default(), - } - .into_request(); - rpc_options.apply_to(&mut request); - WorkflowService::respond_activity_task_completed_by_id( - &mut client, - request, - ) - .await - .map_err(AsyncActivityError::from_status)?; + result, + identity: client.identity(), + resource_id: Default::default(), } + .into_request(); + rpc_options.apply_to(&mut request); + WorkflowService::respond_activity_task_completed_by_id( + &mut client, + request, + ) + .await + .map_err(AsyncActivityError::from_status)?; } Ok(()) }) @@ -211,7 +241,7 @@ impl AsyncActivityHandle { input.into_parts(); let data_converter = client.data_converter().clone(); let mut failure = data_converter.to_failure( - &SerializationContextData::Activity, + &SerializationContextData::Activity(ActivitySerializationContext::new()), OutgoingError::Activity(OutgoingActivityError::Application(Box::new( application_failure, ))), @@ -219,54 +249,49 @@ impl AsyncActivityHandle { encode_payloads( &mut failure, data_converter.codec(), - &SerializationContextData::Activity, + &SerializationContextData::Activity(ActivitySerializationContext::new()), ) .await?; let last_heartbeat_details = encode_optional_value(details, &data_converter).await?; - match identifier { - ActivityIdentifier::TaskToken(token) => { - let mut request = RespondActivityTaskFailedRequest { - task_token: token.0, - failure: Some(failure), - identity: client.identity(), - namespace: client.namespace(), - last_heartbeat_details, - ..Default::default() - } - .into_request(); - rpc_options.apply_to(&mut request); - WorkflowService::respond_activity_task_failed( - &mut client, - request, - ) - .await - .map_err(AsyncActivityError::from_status)?; + if let ActivityIdentifier::TaskToken(token) = identifier { + let mut request = RespondActivityTaskFailedRequest { + task_token: token.into_inner(), + failure: Some(failure), + identity: client.identity(), + namespace: client.namespace(), + last_heartbeat_details, + ..Default::default() } - ActivityIdentifier::ById { + .into_request(); + rpc_options.apply_to(&mut request); + WorkflowService::respond_activity_task_failed( + &mut client, + request, + ) + .await + .map_err(AsyncActivityError::from_status)?; + } else { + let (workflow_id, run_id, activity_id) = identifier.into_parts().unwrap(); + let mut request = RespondActivityTaskFailedByIdRequest { + namespace: client.namespace(), workflow_id, run_id, activity_id, - } => { - let mut request = RespondActivityTaskFailedByIdRequest { - namespace: client.namespace(), - workflow_id, - run_id, - activity_id, - failure: Some(failure), - identity: client.identity(), - last_heartbeat_details, - resource_id: Default::default(), - } - .into_request(); - rpc_options.apply_to(&mut request); - WorkflowService::respond_activity_task_failed_by_id( - &mut client, - request, - ) - .await - .map_err(AsyncActivityError::from_status)?; + failure: Some(failure), + identity: client.identity(), + last_heartbeat_details, + resource_id: Default::default(), + ..Default::default() } + .into_request(); + rpc_options.apply_to(&mut request); + WorkflowService::respond_activity_task_failed_by_id( + &mut client, + request, + ) + .await + .map_err(AsyncActivityError::from_status)?; } Ok(()) }) @@ -301,47 +326,41 @@ impl AsyncActivityHandle { Box::pin(async move { let (identifier, details, rpc_options) = input.into_parts(); let details = encode_optional_value(details, client.data_converter()).await?; - match identifier { - ActivityIdentifier::TaskToken(token) => { - let mut request = RespondActivityTaskCanceledRequest { - task_token: token.0, - details, - identity: client.identity(), - namespace: client.namespace(), - ..Default::default() - } - .into_request(); - rpc_options.apply_to(&mut request); - WorkflowService::respond_activity_task_canceled( - &mut client, - request, - ) - .await - .map_err(AsyncActivityError::from_status)?; + if let ActivityIdentifier::TaskToken(token) = identifier { + let mut request = RespondActivityTaskCanceledRequest { + task_token: token.into_inner(), + details, + identity: client.identity(), + namespace: client.namespace(), + ..Default::default() } - ActivityIdentifier::ById { + .into_request(); + rpc_options.apply_to(&mut request); + WorkflowService::respond_activity_task_canceled( + &mut client, + request, + ) + .await + .map_err(AsyncActivityError::from_status)?; + } else { + let (workflow_id, run_id, activity_id) = identifier.into_parts().unwrap(); + let mut request = RespondActivityTaskCanceledByIdRequest { + namespace: client.namespace(), workflow_id, run_id, activity_id, - } => { - let mut request = RespondActivityTaskCanceledByIdRequest { - namespace: client.namespace(), - workflow_id, - run_id, - activity_id, - details, - identity: client.identity(), - ..Default::default() - } - .into_request(); - rpc_options.apply_to(&mut request); - WorkflowService::respond_activity_task_canceled_by_id( - &mut client, - request, - ) - .await - .map_err(AsyncActivityError::from_status)?; + details, + identity: client.identity(), + ..Default::default() } + .into_request(); + rpc_options.apply_to(&mut request); + WorkflowService::respond_activity_task_canceled_by_id( + &mut client, + request, + ) + .await + .map_err(AsyncActivityError::from_status)?; } Ok(()) }) @@ -375,52 +394,46 @@ impl AsyncActivityHandle { Box::pin(async move { let (identifier, details, rpc_options) = input.into_parts(); let details = encode_optional_value(details, client.data_converter()).await?; - match identifier { - ActivityIdentifier::TaskToken(token) => { - let mut request = RecordActivityTaskHeartbeatRequest { - task_token: token.0, - details, - identity: client.identity(), - namespace: client.namespace(), - resource_id: Default::default(), - } - .into_request(); - rpc_options.apply_to(&mut request); - let response = WorkflowService::record_activity_task_heartbeat( + if let ActivityIdentifier::TaskToken(token) = identifier { + let mut request = RecordActivityTaskHeartbeatRequest { + task_token: token.into_inner(), + details, + identity: client.identity(), + namespace: client.namespace(), + resource_id: Default::default(), + } + .into_request(); + rpc_options.apply_to(&mut request); + let response = WorkflowService::record_activity_task_heartbeat( + &mut client, + request, + ) + .await + .map_err(AsyncActivityError::from_status)? + .into_inner(); + Ok(ActivityHeartbeatResponse::from(response)) + } else { + let (workflow_id, run_id, activity_id) = identifier.into_parts().unwrap(); + let mut request = RecordActivityTaskHeartbeatByIdRequest { + namespace: client.namespace(), + workflow_id, + run_id, + activity_id, + details, + identity: client.identity(), + resource_id: Default::default(), + } + .into_request(); + rpc_options.apply_to(&mut request); + let response = + WorkflowService::record_activity_task_heartbeat_by_id( &mut client, request, ) .await .map_err(AsyncActivityError::from_status)? .into_inner(); - Ok(ActivityHeartbeatResponse::from(response)) - } - ActivityIdentifier::ById { - workflow_id, - run_id, - activity_id, - } => { - let mut request = RecordActivityTaskHeartbeatByIdRequest { - namespace: client.namespace(), - workflow_id, - run_id, - activity_id, - details, - identity: client.identity(), - resource_id: Default::default(), - } - .into_request(); - rpc_options.apply_to(&mut request); - let response = - WorkflowService::record_activity_task_heartbeat_by_id( - &mut client, - request, - ) - .await - .map_err(AsyncActivityError::from_status)? - .into_inner(); - Ok(ActivityHeartbeatResponse::from(response)) - } + Ok(ActivityHeartbeatResponse::from(response)) } }) } @@ -432,6 +445,7 @@ impl AsyncActivityHandle { /// Response from a heartbeat call. #[derive(Debug, Clone)] +#[non_exhaustive] pub struct ActivityHeartbeatResponse { /// True if the activity has been asked to cancel itself. pub cancel_requested: bool, diff --git a/crates/client/src/envconfig.rs b/crates/client/src/envconfig.rs index c1360bf69..4eebbf8e0 100644 --- a/crates/client/src/envconfig.rs +++ b/crates/client/src/envconfig.rs @@ -117,6 +117,7 @@ impl TryFrom for ConnectionOptions { tls, codec: _, grpc_meta, + .. } = profile; let has_api_key = api_key.is_some(); @@ -147,6 +148,10 @@ fn resolve_datasource(data_source: DataSource) -> Result, std::io::Error match data_source { DataSource::Path(path) => fs::read(path), DataSource::Data(data) => Ok(data), + _ => Err(std::io::Error::new( + std::io::ErrorKind::Unsupported, + "unsupported envconfig data source", + )), } } @@ -184,21 +189,17 @@ mod tests { #[case] expected: &str, ) { let tls = enable_tls.then(ClientConfigTLS::default); - let profile = ClientConfigProfile { - address: address.map(str::to_string), - tls, - ..Default::default() - }; + let profile = ClientConfigProfile::builder() + .maybe_address(address.map(str::to_string)) + .maybe_tls(tls) + .build(); let conn: ConnectionOptions = profile.try_into().unwrap(); assert_eq!(conn.target.as_str(), expected); } #[test] fn invalid_address_errors() { - let profile = ClientConfigProfile { - address: Some("://bad".to_string()), - ..Default::default() - }; + let profile = ClientConfigProfile::builder().address("://bad").build(); assert!(ConnectionOptions::try_from(profile).is_err()); } @@ -233,20 +234,16 @@ mod tests { let mut meta = HashMap::new(); meta.insert("x-custom".to_string(), "value".to_string()); meta.insert("another".to_string(), "header".to_string()); - let profile = ClientConfigProfile { - grpc_meta: meta.clone(), - ..Default::default() - }; + let profile = ClientConfigProfile::builder() + .grpc_meta(meta.clone()) + .build(); let conn: ConnectionOptions = profile.try_into().unwrap(); assert_eq!(conn.headers.unwrap(), meta); } #[test] fn api_key_populates_field() { - let profile = ClientConfigProfile { - api_key: Some("my-key".to_string()), - ..Default::default() - }; + let profile = ClientConfigProfile::builder().api_key("my-key").build(); let conn: ConnectionOptions = profile.try_into().unwrap(); assert_eq!(conn.api_key.as_deref(), Some("my-key")); } @@ -264,28 +261,27 @@ mod tests { #[case] api_key: Option<&str>, #[case] expect_tls: bool, ) { - let profile = ClientConfigProfile { - api_key: api_key.map(str::to_string), - tls: tls_disabled.map(|disabled| ClientConfigTLS { - disabled, - ..Default::default() - }), - ..Default::default() - }; + let profile = ClientConfigProfile::builder() + .maybe_api_key(api_key.map(str::to_string)) + .maybe_tls( + tls_disabled + .map(|disabled| ClientConfigTLS::builder().maybe_disabled(disabled).build()), + ) + .build(); let conn: ConnectionOptions = profile.try_into().unwrap(); assert_eq!(conn.tls_options.is_some(), expect_tls); } #[test] fn data_source_certs() { - let profile = ClientConfigProfile { - tls: Some(ClientConfigTLS { - client_cert: Some(DataSource::Data(b"cert-data".to_vec())), - client_key: Some(DataSource::Data(b"key-data".to_vec())), - ..Default::default() - }), - ..Default::default() - }; + let profile = ClientConfigProfile::builder() + .tls( + ClientConfigTLS::builder() + .client_cert(DataSource::Data(b"cert-data".to_vec())) + .client_key(DataSource::Data(b"key-data".to_vec())) + .build(), + ) + .build(); let conn: ConnectionOptions = profile.try_into().unwrap(); let tls = conn.tls_options.unwrap(); let mtls = tls.client_tls_options.unwrap(); @@ -300,14 +296,14 @@ mod tests { std::fs::write(&cert_path, b"file-cert").unwrap(); std::fs::write(&key_path, b"file-key").unwrap(); - let profile = ClientConfigProfile { - tls: Some(ClientConfigTLS { - client_cert: Some(DataSource::Path(cert_path.to_str().unwrap().to_string())), - client_key: Some(DataSource::Path(key_path.to_str().unwrap().to_string())), - ..Default::default() - }), - ..Default::default() - }; + let profile = ClientConfigProfile::builder() + .tls( + ClientConfigTLS::builder() + .client_cert(DataSource::Path(cert_path.to_str().unwrap().to_string())) + .client_key(DataSource::Path(key_path.to_str().unwrap().to_string())) + .build(), + ) + .build(); let conn: ConnectionOptions = profile.try_into().unwrap(); let tls = conn.tls_options.unwrap(); let mtls = tls.client_tls_options.unwrap(); @@ -317,13 +313,13 @@ mod tests { #[test] fn server_ca_cert() { - let profile = ClientConfigProfile { - tls: Some(ClientConfigTLS { - server_ca_cert: Some(DataSource::Data(b"ca-data".to_vec())), - ..Default::default() - }), - ..Default::default() - }; + let profile = ClientConfigProfile::builder() + .tls( + ClientConfigTLS::builder() + .server_ca_cert(DataSource::Data(b"ca-data".to_vec())) + .build(), + ) + .build(); let conn: ConnectionOptions = profile.try_into().unwrap(); let tls = conn.tls_options.unwrap(); assert_eq!(tls.server_root_ca_cert.unwrap(), b"ca-data"); @@ -331,13 +327,13 @@ mod tests { #[test] fn server_name_sni() { - let profile = ClientConfigProfile { - tls: Some(ClientConfigTLS { - server_name: Some("my.server.com".to_string()), - ..Default::default() - }), - ..Default::default() - }; + let profile = ClientConfigProfile::builder() + .tls( + ClientConfigTLS::builder() + .server_name("my.server.com") + .build(), + ) + .build(); let conn: ConnectionOptions = profile.try_into().unwrap(); let tls = conn.tls_options.unwrap(); assert_eq!(tls.domain.as_deref(), Some("my.server.com")); @@ -350,14 +346,14 @@ mod tests { #[case] client_cert: Option, #[case] client_key: Option, ) { - let profile = ClientConfigProfile { - tls: Some(ClientConfigTLS { - client_cert, - client_key, - ..Default::default() - }), - ..Default::default() - }; + let profile = ClientConfigProfile::builder() + .tls( + ClientConfigTLS::builder() + .maybe_client_cert(client_cert) + .maybe_client_key(client_key) + .build(), + ) + .build(); assert!(ConnectionOptions::try_from(profile).is_err()); } diff --git a/crates/client/src/errors.rs b/crates/client/src/errors.rs index 1bc16a725..e958511bc 100644 --- a/crates/client/src/errors.rs +++ b/crates/client/src/errors.rs @@ -1,10 +1,24 @@ //! Contains errors that can be returned by clients. -use crate::{PluginApplyError, WorkflowExecutionStatus, workflow_handle::WorkflowResultDetails}; +#[cfg(feature = "experimental")] +use crate::PluginApplyError; +use crate::{WorkflowExecutionStatus, workflow_handle::WorkflowResultDetails}; use http::uri::InvalidUri; use temporalio_common::{ - data_converters::PayloadConversionError, error::IncomingError, - protos::temporal::api::failure::v1::Failure, + data_converters::{DecodablePayloads, PayloadConversionError}, + error::{IncomingError, TimeoutType}, + protos::{ + google::rpc::Status as RpcStatus, + temporal::api::{ + errordetails::v1::{ + ActivityExecutionAlreadyStartedFailure, MultiOperationExecutionFailure, + WorkflowExecutionAlreadyStartedFailure, + multi_operation_execution_failure::OperationStatus, + }, + failure::v1::Failure, + }, + utilities::{decode_status_detail, encode_status_details}, + }, }; use tonic::Code; @@ -13,6 +27,7 @@ use tonic::Code; #[non_exhaustive] pub enum ClientConnectError { /// A plugin failed while configuring connection options. + #[cfg(feature = "experimental")] #[error(transparent)] Plugin(#[from] PluginApplyError), /// Invalid URI. Configuration error, fatal. @@ -45,6 +60,7 @@ pub enum ClientConnectError { impl From for ClientConnectError { fn from(value: ClientNewError) -> Self { match value { + #[cfg(feature = "experimental")] ClientNewError::Plugin(err) => Self::Plugin(err), } } @@ -103,6 +119,22 @@ pub enum WorkflowStartError { Rpc(#[from] tonic::Status), } +impl WorkflowStartError { + pub(crate) fn from_status(status: tonic::Status) -> Self { + if status.code() == Code::AlreadyExists { + let run_id = + decode_status_detail::(status.details()) + .map(|failure| failure.run_id); + Self::AlreadyStarted { + run_id, + source: status, + } + } else { + Self::Rpc(status) + } + } +} + /// Errors returned by query operations on [crate::WorkflowHandle]. #[derive(Debug, thiserror::Error)] #[non_exhaustive] @@ -176,6 +208,79 @@ impl WorkflowUpdateError { } } +/// Errors returned by update-with-start operations +/// (see [crate::Client::start_update_with_start_workflow]). +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum WorkflowUpdateWithStartError { + /// The start operation failed. + #[error("Workflow start failed: {0}")] + Start(#[source] WorkflowStartError), + + /// The update operation failed, or waiting for the update result failed. + #[error("Workflow update failed: {0}")] + Update(#[source] WorkflowUpdateError), + + /// Error serializing the workflow input or update arguments. + #[error("Payload conversion error: {0}")] + PayloadConversion(#[from] PayloadConversionError), + + /// An RPC error from the server that could not be attributed to either operation. + #[error("Server error: {0}")] + Rpc(tonic::Status), + + /// Other errors. + #[error(transparent)] + Other(#[from] Box), +} + +const MULTI_OPERATION_ABORTED_NAME: &str = "temporal.api.failure.v1.MultiOperationExecutionAborted"; + +/// Reconstruct a standalone gRPC status from a multi-operation `OperationStatus`, re-encoding +/// its details so the operation-specific failure information stays available to callers. +fn operation_status_to_tonic(op_status: OperationStatus) -> tonic::Status { + let code = Code::from(op_status.code); + let details = encode_status_details(&RpcStatus { + code: op_status.code, + message: op_status.message.clone(), + details: op_status.details, + }); + tonic::Status::with_details(code, op_status.message, details.into()) +} + +impl WorkflowUpdateWithStartError { + /// A multi-operation failure carries one status per operation; all operations except the + /// failed one are marked aborted. Attribute the error to the operation that actually failed + /// (index 0 is the start operation, index 1 the update). + pub(crate) fn from_status(status: tonic::Status) -> Self { + let Some(failure) = + decode_status_detail::(status.details()) + else { + return Self::Rpc(status); + }; + let culprit = failure + .statuses + .into_iter() + .enumerate() + .find(|(_, op_status)| { + op_status.code != Code::Ok as i32 + && !op_status + .details + .iter() + .any(|detail| detail.type_url.ends_with(MULTI_OPERATION_ABORTED_NAME)) + }); + match culprit { + Some((0, op_status)) => Self::Start(WorkflowStartError::from_status( + operation_status_to_tonic(op_status), + )), + Some((_, op_status)) => Self::Update(WorkflowUpdateError::from_status( + operation_status_to_tonic(op_status), + )), + None => Self::Rpc(status), + } + } +} + /// Errors returned by workflow get_result operations. #[derive(Debug, thiserror::Error)] #[non_exhaustive] @@ -324,6 +429,282 @@ impl AsyncActivityError { #[non_exhaustive] pub enum ClientNewError { /// A plugin failed while configuring client options. + #[cfg(feature = "experimental")] #[error(transparent)] Plugin(#[from] PluginApplyError), } + +/// Errors returned by methods on [crate::ActivityHandle] that don't need more specific error types. +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum ActivityInteractionError { + /// The activity was not found. + #[error("Activity not found")] + NotFound(#[source] tonic::Status), + + /// Error deserializing output. + #[error("Payload conversion error: {0}")] + PayloadConversion(#[from] PayloadConversionError), + + /// An uncategorized RPC error from the server. + #[error("Server error: {0}")] + Rpc(#[source] tonic::Status), + + /// Other errors. + #[error(transparent)] + Other(#[from] Box), +} + +impl From for ActivityInteractionError { + fn from(status: tonic::Status) -> Self { + if status.code() == Code::NotFound { + Self::NotFound(status) + } else { + Self::Rpc(status) + } + } +} + +/// Errors that can occur when starting a standalone activity. +#[allow(clippy::large_enum_variant)] +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum StartActivityError { + /// There's a conflicting activity execution with the same ID according to chosen ID reuse + /// policy and ID conflict policy. + #[error("Activity already started with run_id={run_id}")] + AlreadyStarted { + /// Run ID of the existing execution with the same activity ID. + run_id: String, + /// Raw error from the server. + #[source] + source: tonic::Status, + }, + + /// Error serializing input. + #[error("Payload conversion error: {0}")] + PayloadConversion(#[from] PayloadConversionError), + + /// An uncategorized RPC error from the server. + #[error("Server error: {0}")] + Rpc(#[source] tonic::Status), + + /// Other errors. + #[error(transparent)] + Other(#[from] Box), +} + +impl From for StartActivityError { + fn from(status: tonic::Status) -> Self { + if status.code() == tonic::Code::AlreadyExists + && let Some(details) = + decode_status_detail::(status.details()) + { + StartActivityError::AlreadyStarted { + run_id: details.run_id, + source: status, + } + } else { + StartActivityError::Rpc(status) + } + } +} + +/// Errors returned by [`crate::ActivityHandle::result`]. +#[allow(clippy::large_enum_variant)] +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum ActivityResultError { + /// Activity execution did not complete successfully. + #[error("Activity failed: {0}")] + ActivityFailed(#[source] IncomingError), + + /// The activity was canceled. + #[error("Activity canceled")] + Cancelled { + /// Details provided at cancellation time. + details: DecodablePayloads, + }, + + /// The workflow was terminated. + #[error("Activity terminated")] + Terminated, + + /// The activity timed out. + #[error("Activity timed out: {0:?}")] + TimedOut(TimeoutType), + + /// The activity was not found. + #[error("Activity not found")] + NotFound(#[source] tonic::Status), + + /// Error deserializing output. + #[error("Payload conversion error: {0}")] + PayloadConversion(#[from] PayloadConversionError), + + /// An uncategorized RPC error from the server. + #[error("Server error: {0}")] + Rpc(#[source] tonic::Status), + + /// Other errors. + #[error(transparent)] + Other(#[from] Box), +} + +impl From for ActivityResultError { + fn from(status: tonic::Status) -> Self { + if status.code() == Code::NotFound { + Self::NotFound(status) + } else { + Self::Rpc(status) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use assert_matches::assert_matches; + use prost::Message; + use temporalio_common::protos::{ + temporal::api::{ + errordetails::v1::NotFoundFailure, failure::v1::MultiOperationExecutionAborted, + }, + utilities::pack_any, + }; + + fn multi_op_status(code: Code, statuses: Vec) -> tonic::Status { + let failure = MultiOperationExecutionFailure { statuses }; + let rpc_status = RpcStatus { + code: code as i32, + message: "multi-op failure".to_owned(), + details: vec![ + pack_any( + "type.googleapis.com/temporal.api.errordetails.v1.MultiOperationExecutionFailure" + .to_owned(), + &failure, + ) + .unwrap(), + ], + }; + tonic::Status::with_details(code, "multi-op failure", rpc_status.encode_to_vec().into()) + } + + fn aborted_status() -> OperationStatus { + OperationStatus { + code: Code::Aborted as i32, + message: "aborted".to_owned(), + details: vec![ + pack_any( + "type.googleapis.com/temporal.api.failure.v1.MultiOperationExecutionAborted" + .to_owned(), + &MultiOperationExecutionAborted {}, + ) + .unwrap(), + ], + } + } + + #[test] + fn update_with_start_error_attributes_start_already_started() { + let status = multi_op_status( + Code::AlreadyExists, + vec![ + OperationStatus { + code: Code::AlreadyExists as i32, + message: "already started".to_owned(), + details: vec![ + pack_any( + "type.googleapis.com/temporal.api.errordetails.v1.WorkflowExecutionAlreadyStartedFailure" + .to_owned(), + &WorkflowExecutionAlreadyStartedFailure { + run_id: "existing-run".to_owned(), + ..Default::default() + }, + ) + .unwrap(), + ], + }, + aborted_status(), + ], + ); + + let err = WorkflowUpdateWithStartError::from_status(status); + assert_matches!( + err, + WorkflowUpdateWithStartError::Start(WorkflowStartError::AlreadyStarted { + run_id: Some(run_id), + .. + }) if run_id == "existing-run" + ); + } + + #[test] + fn update_with_start_error_attributes_update_failure() { + let status = multi_op_status( + Code::NotFound, + vec![ + aborted_status(), + OperationStatus { + code: Code::NotFound as i32, + message: "no such workflow".to_owned(), + details: vec![ + pack_any( + "type.googleapis.com/temporal.api.errordetails.v1.NotFoundFailure" + .to_owned(), + &NotFoundFailure { + current_cluster: "here".to_owned(), + ..Default::default() + }, + ) + .unwrap(), + ], + }, + ], + ); + + let err = WorkflowUpdateWithStartError::from_status(status); + let inner = assert_matches!( + err, + WorkflowUpdateWithStartError::Update(WorkflowUpdateError::NotFound(status)) => status + ); + assert_eq!(inner.message(), "no such workflow"); + // The operation's own failure details must survive reconstruction of the inner status. + let detail = decode_status_detail::(inner.details()) + .expect("operation details must be preserved"); + assert_eq!(detail.current_cluster, "here"); + } + + #[test] + fn update_with_start_error_skips_successful_start() { + let status = multi_op_status( + Code::NotFound, + vec![ + OperationStatus { + code: Code::Ok as i32, + message: String::new(), + details: vec![], + }, + OperationStatus { + code: Code::NotFound as i32, + message: "update failed".to_owned(), + details: vec![], + }, + ], + ); + + let err = WorkflowUpdateWithStartError::from_status(status); + assert_matches!( + err, + WorkflowUpdateWithStartError::Update(WorkflowUpdateError::NotFound(status)) + if status.message() == "update failed" + ); + } + + #[test] + fn update_with_start_error_without_details_is_rpc() { + let err = + WorkflowUpdateWithStartError::from_status(tonic::Status::new(Code::Internal, "boom")); + assert_matches!(err, WorkflowUpdateWithStartError::Rpc(status) if status.code() == Code::Internal); + } +} diff --git a/crates/client/src/grpc.rs b/crates/client/src/grpc.rs index 59ba46875..9aa1c8306 100644 --- a/crates/client/src/grpc.rs +++ b/crates/client/src/grpc.rs @@ -193,7 +193,7 @@ fn req_cloner(cloneme: &Request) -> Request { /// `*_warn` are the connection's configured warn thresholds; per-call error limits ride a /// [`PayloadErrorLimits`] extension. On an error-level violation, returns a [`Status`] carrying -/// the [`PayloadLimitViolation`] as its source (extract via [crate::payload_limit_violation_from]). +/// the payload limit violation as its source. fn validate_request_payload_limits( req: &Request, blob_warn: usize, @@ -1387,8 +1387,22 @@ proxier! { ExecuteMultiOperationRequest, ExecuteMultiOperationResponse, |r| { - let labels = namespaced_request!(r); - r.extensions_mut().insert(labels); + let mut labels = namespaced_request!(r); + if let Some(execute_multi_operation_request::operation::Operation::StartWorkflow( + start_req, + )) = r + .get_ref() + .operations + .first() + .and_then(|op| op.operation.as_ref()) + { + labels.task_q(start_req.task_queue.clone()); + } + let exts = r.extensions_mut(); + exts.insert(labels); + // Update-with-start blocks until the update reaches the requested wait stage, so it + // must be retried/timed out like other user long-polls. + exts.insert(IsUserLongPoll); } ); ( @@ -2086,8 +2100,11 @@ mod tests { req.extensions_mut() .insert(PayloadErrorLimits { blob: 10, memo: 10 }); let err = validate_request_payload_limits(&req, 1, 1).unwrap_err(); - let violation = - crate::payload_limit_violation_from(&err).expect("violation carried on status"); + let violation = std::error::Error::source(&err) + .and_then(|source| { + source.downcast_ref::() + }) + .expect("violation carried on status"); assert_eq!(violation.path, "input"); assert_eq!( violation.class, @@ -2338,10 +2355,12 @@ mod tests { } } - let deployment_opts = WorkerDeploymentOptions::new(WorkerDeploymentVersion { - deployment_name: "test-deployment".to_string(), - build_id: "test-build-123".to_string(), - }) + let deployment_opts = WorkerDeploymentOptions::new( + WorkerDeploymentVersion::builder() + .deployment_name("test-deployment".to_string()) + .build_id("test-build-123".to_string()) + .build(), + ) .use_worker_versioning(use_worker_versioning) .build(); diff --git a/crates/client/src/interceptors.rs b/crates/client/src/interceptors.rs index f9b4570f5..abaf291d5 100644 --- a/crates/client/src/interceptors.rs +++ b/crates/client/src/interceptors.rs @@ -4,10 +4,10 @@ use crate::{ ActivityHeartbeatResponse, ActivityIdentifier, WorkflowCancelOptions, WorkflowCountOptions, WorkflowDescribeOptions, WorkflowFetchHistoryOptions, WorkflowQueryOptions, WorkflowSignalOptions, WorkflowStartError, WorkflowStartOptions, WorkflowStartUpdateOptions, - WorkflowTerminateOptions, + WorkflowTerminateOptions, WorkflowUpdateWithStartOptions, errors::{ AsyncActivityError, ClientError, WorkflowInteractionError, WorkflowQueryError, - WorkflowUpdateError, + WorkflowUpdateError, WorkflowUpdateWithStartError, }, schedules::{ CreateScheduleOptions, ScheduleBackfill, ScheduleError, ScheduleOverlapPolicy, @@ -187,6 +187,106 @@ impl StartWorkflowInput { impl_with_args!(StartWorkflowInput); +/// Input to [`ClientInterceptor::signal_with_start_workflow`]. +#[non_exhaustive] +#[derive(derive_more::Debug)] +pub struct SignalWithStartWorkflowInput { + /// The workflow type sent to the server. + pub workflow_type: String, + /// The signal name sent to the workflow. + pub signal_name: String, + /// Options for the workflow start. + pub options: WorkflowStartOptions, + /// Controls for the signal-with-start RPC. + pub rpc_options: crate::RpcOptions, + // These remain type-erased until after interception so interceptors can replace either value + // before the client's payload converter and codec run. + #[debug(skip)] + workflow_args: Box, + #[debug(skip)] + signal_args: Box, +} + +impl SignalWithStartWorkflowInput { + pub(crate) fn new( + workflow_type: String, + workflow_args: W, + signal_name: String, + signal_args: S, + mut options: WorkflowStartOptions, + ) -> Self + where + W: TemporalSerializable + Send + 'static, + S: TemporalSerializable + Send + 'static, + { + let rpc_options = std::mem::take(&mut options.rpc_options); + Self { + workflow_type, + signal_name, + options, + rpc_options, + workflow_args: Box::new(workflow_args), + signal_args: Box::new(signal_args), + } + } + + pub(crate) fn into_parts( + self, + ) -> ( + String, + Box, + String, + Box, + WorkflowStartOptions, + crate::RpcOptions, + ) { + ( + self.workflow_type, + self.workflow_args, + self.signal_name, + self.signal_args, + self.options, + self.rpc_options, + ) + } + + /// Attempt to access the workflow arguments as a concrete type. + pub fn workflow_args_ref(&self) -> Option<&T> { + self.workflow_args.as_any().downcast_ref() + } + + /// Attempt to access the signal arguments as a concrete type. + pub fn signal_args_ref(&self) -> Option<&T> { + self.signal_args.as_any().downcast_ref() + } + + /// Attempt to mutably access the workflow arguments as a concrete type. + pub fn workflow_args_mut(&mut self) -> Option<&mut T> { + self.workflow_args.as_any_mut().downcast_mut() + } + + /// Attempt to mutably access the signal arguments as a concrete type. + pub fn signal_args_mut(&mut self) -> Option<&mut T> { + self.signal_args.as_any_mut().downcast_mut() + } + + /// Replace the workflow arguments before serialization. + pub fn replace_workflow_args(&mut self, args: T) + where + T: TemporalSerializable + Send + 'static, + { + self.workflow_args = Box::new(args); + } + + /// Replace the signal arguments before serialization. + pub fn replace_signal_args(&mut self, args: T) + where + T: TemporalSerializable + Send + 'static, + { + self.signal_args = Box::new(args); + } +} + /// Result of a successful intercepted workflow start. #[non_exhaustive] #[derive(Clone, Debug, PartialEq, Eq)] @@ -535,6 +635,114 @@ impl StartWorkflowUpdateOutput { } } +/// Input to [`ClientInterceptor::update_with_start_workflow`]. +#[non_exhaustive] +#[derive(derive_more::Debug)] +pub struct UpdateWithStartWorkflowInput { + /// The workflow type sent to the server. + pub workflow_type: String, + /// Update name sent to the workflow. + pub update_name: String, + /// Options for the atomic start-and-update operation. + pub options: WorkflowUpdateWithStartOptions, + /// Controls for the multi-operation RPC. + pub rpc_options: crate::RpcOptions, + #[debug(skip)] + pub(crate) workflow_args: Box, + #[debug(skip)] + pub(crate) update_args: Box, +} + +impl UpdateWithStartWorkflowInput { + pub(crate) fn new( + workflow_type: String, + workflow_args: WA, + update_name: String, + update_args: UA, + mut options: WorkflowUpdateWithStartOptions, + ) -> Self + where + WA: TemporalSerializable + Send + 'static, + UA: TemporalSerializable + Send + 'static, + { + let rpc_options = std::mem::take(&mut options.rpc_options); + Self { + workflow_type, + update_name, + options, + rpc_options, + workflow_args: Box::new(workflow_args), + update_args: Box::new(update_args), + } + } + + /// Attempt to access the workflow start arguments as a concrete type. + pub fn workflow_args_ref(&self) -> Option<&T> { + self.workflow_args.as_any().downcast_ref() + } + + /// Attempt to mutably access the workflow start arguments as a concrete type. + pub fn workflow_args_mut(&mut self) -> Option<&mut T> { + self.workflow_args.as_any_mut().downcast_mut() + } + + /// Replace the workflow start arguments with another serializable value. + pub fn replace_workflow_args(&mut self, args: T) + where + T: TemporalSerializable + Send + 'static, + { + self.workflow_args = Box::new(args); + } + + /// Attempt to access the update arguments as a concrete type. + pub fn update_args_ref(&self) -> Option<&T> { + self.update_args.as_any().downcast_ref() + } + + /// Attempt to mutably access the update arguments as a concrete type. + pub fn update_args_mut(&mut self) -> Option<&mut T> { + self.update_args.as_any_mut().downcast_mut() + } + + /// Replace the update arguments with another serializable value. + pub fn replace_update_args(&mut self, args: T) + where + T: TemporalSerializable + Send + 'static, + { + self.update_args = Box::new(args); + } +} + +/// Result of an intercepted update-with-start operation. +#[non_exhaustive] +#[derive(Clone, Debug)] +pub struct UpdateWithStartWorkflowOutput { + /// Workflow ID used by the operation. + pub workflow_id: String, + /// Update ID used by the operation. + pub update_id: String, + /// Run ID associated with the update, when available. + pub run_id: Option, + /// Outcome returned when the requested wait stage completed the update. + pub known_outcome: Option, +} + +impl UpdateWithStartWorkflowOutput { + pub(crate) fn new( + workflow_id: impl Into, + update_id: impl Into, + run_id: Option, + known_outcome: Option, + ) -> Self { + Self { + workflow_id: workflow_id.into(), + update_id: update_id.into(), + run_id, + known_outcome, + } + } +} + /// Input to [`ClientInterceptor::poll_workflow_update`]. #[non_exhaustive] #[derive(Clone, Debug)] @@ -1062,6 +1270,19 @@ pub trait ClientInterceptor: Send + Sync + 'static { next.run(input) } + /// Intercept a `signal_with_start_workflow` operation. + fn signal_with_start_workflow<'a>( + &'a self, + input: SignalWithStartWorkflowInput, + next: Next< + 'a, + SignalWithStartWorkflowInput, + BoxFuture<'a, Result>, + >, + ) -> BoxFuture<'a, Result> { + next.run(input) + } + /// Intercept a `list_workflows_page` operation. fn list_workflows_page<'a>( &'a self, @@ -1149,6 +1370,19 @@ pub trait ClientInterceptor: Send + Sync + 'static { next.run(input) } + /// Intercept an `update_with_start_workflow` operation. + fn update_with_start_workflow<'a>( + &'a self, + input: UpdateWithStartWorkflowInput, + next: Next< + 'a, + UpdateWithStartWorkflowInput, + BoxFuture<'a, Result>, + >, + ) -> BoxFuture<'a, Result> { + next.run(input) + } + /// Intercept a `poll_workflow_update` operation. fn poll_workflow_update<'a>( &'a self, @@ -1351,6 +1585,13 @@ interceptor_chain!( BoxFuture<'a, Result> ); +interceptor_chain!( + call_signal_with_start_workflow, + signal_with_start_workflow, + SignalWithStartWorkflowInput, + BoxFuture<'a, Result> +); + interceptor_chain!( call_list_workflows_page, list_workflows_page, @@ -1400,6 +1641,13 @@ interceptor_chain!( BoxFuture<'a, Result> ); +interceptor_chain!( + call_update_with_start_workflow, + update_with_start_workflow, + UpdateWithStartWorkflowInput, + BoxFuture<'a, Result> +); + interceptor_chain!( call_poll_workflow_update, poll_workflow_update, diff --git a/crates/client/src/lib.rs b/crates/client/src/lib.rs index 19cf2b89b..4c80ce5b0 100644 --- a/crates/client/src/lib.rs +++ b/crates/client/src/lib.rs @@ -1,3 +1,4 @@ +#![cfg_attr(docsrs, feature(doc_cfg))] #![warn(missing_docs)] // error if there are missing docs //! This crate contains client implementations that can be used to contact the Temporal service. @@ -7,6 +8,7 @@ #[macro_use] extern crate tracing; +mod activity; mod async_activity_handle; pub mod callback_based; mod dns; @@ -19,11 +21,10 @@ pub mod grpc; pub mod interceptors; mod metrics; mod options_structs; +#[cfg(feature = "experimental")] /// Experimental APIs for configuring clients with reusable plugins. pub mod plugins; -/// Visible only for tests -#[doc(hidden)] -pub mod proxy; +mod proxy; mod replaceable; pub mod request_extensions; mod retry; @@ -36,14 +37,12 @@ pub mod worker; mod workflow_handle; mod workflow_status; -pub use crate::{ - proxy::HttpConnectProxyOptions, - request_extensions::PayloadErrorLimits, - retry::{CallType, RETRYABLE_ERROR_CODES}, -}; +pub use crate::{proxy::HttpConnectProxyOptions, request_extensions::PayloadErrorLimits}; +pub use activity::*; pub use async_activity_handle::{ ActivityHeartbeatResponse, ActivityIdentifier, AsyncActivityHandle, }; +pub(crate) use retry::CallType; #[doc(hidden)] pub use retry::jittered; @@ -56,12 +55,14 @@ pub use interceptors::{ ListSchedulesPageOutput, ListWorkflowsPageInput, ListWorkflowsPageOutput, Next, PauseScheduleInput, PollWorkflowUpdateInput, PollWorkflowUpdateOutput, QueryWorkflowInput, QueryWorkflowOutput, ReportAsyncActivityCancellationInput, SendScheduleUpdateInput, - SignalWorkflowInput, StartWorkflowInput, StartWorkflowOutput, StartWorkflowUpdateInput, - StartWorkflowUpdateOutput, TemporalClientValue, TerminateWorkflowInput, TriggerScheduleInput, - UnpauseScheduleInput, UpdateScheduleInput, + SignalWithStartWorkflowInput, SignalWorkflowInput, StartWorkflowInput, StartWorkflowOutput, + StartWorkflowUpdateInput, StartWorkflowUpdateOutput, TemporalClientValue, + TerminateWorkflowInput, TriggerScheduleInput, UnpauseScheduleInput, UpdateScheduleInput, + UpdateWithStartWorkflowInput, UpdateWithStartWorkflowOutput, }; pub use metrics::{LONG_REQUEST_LATENCY_HISTOGRAM_NAME, REQUEST_LATENCY_HISTOGRAM_NAME}; pub use options_structs::*; +#[cfg(feature = "experimental")] pub use plugins::{ ClientPlugin, ErasedClientPlugin, PluginApplyError, PluginError, PluginTarget, WorkerPluginData, }; @@ -99,7 +100,7 @@ pub use tonic; pub use workflow_handle::{ UntypedQuery, UntypedSignal, UntypedUpdate, UntypedWorkflow, UntypedWorkflowHandle, WorkflowExecutionDescription, WorkflowExecutionInfo, WorkflowExecutionResult, WorkflowHandle, - WorkflowHistory, WorkflowHistoryJsonError, WorkflowResultDetails, WorkflowUpdateHandle, + WorkflowHistory, WorkflowHistoryError, WorkflowResultDetails, WorkflowUpdateHandle, }; pub use workflow_status::WorkflowExecutionStatus; @@ -113,11 +114,16 @@ use crate::{ worker::ClientWorkerSet, }; use errors::*; -use futures_util::{future::BoxFuture, stream, stream::Stream}; +use futures_util::{ + future::{BoxFuture, try_join}, + stream, + stream::Stream, +}; use http::Uri; use parking_lot::RwLock; use std::{ collections::{HashMap, VecDeque}, + error::Error, fmt::Debug, pin::Pin, str::FromStr, @@ -126,10 +132,10 @@ use std::{ time::{Duration, SystemTime}, }; use temporalio_common::{ - HasWorkflowDefinition, + ActivityDefinition, HasWorkflowDefinition, SignalDefinition, UntypedActivity, UpdateDefinition, data_converters::{ - DataConverter, GenericPayloadConverter, PayloadConverter, SerializationContext, - SerializationContextData, + ActivitySerializationContext, DataConverter, SerializationContext, + SerializationContextData, WorkflowSerializationContext, }, payload_visitor::decode_payloads, protos::{ @@ -138,20 +144,26 @@ use temporalio_common::{ proto_ts_to_system_time, temporal::api::{ cloud::cloudservice::v1::cloud_service_client::CloudServiceClient, - common::v1::WorkflowType, - enums::v1::TaskQueueKind, - errordetails::v1::WorkflowExecutionAlreadyStartedFailure, + common::v1::{ActivityType, Memo as ProtoMemo, Payloads, WorkflowType}, + enums::v1::{ + ActivityIdConflictPolicy as ProtoActivityIdConflictPolicy, + ActivityIdReusePolicy as ProtoActivityIdReusePolicy, TaskQueueKind, + UpdateWorkflowExecutionLifecycleStage, + WorkflowIdConflictPolicy as ProtoWorkflowIdConflictPolicy, + WorkflowIdReusePolicy as ProtoWorkflowIdReusePolicy, + }, operatorservice::v1::operator_service_client::OperatorServiceClient, sdk::v1::UserMetadata, taskqueue::v1::TaskQueue, testservice::v1::test_service_client::TestServiceClient, workflow::v1 as workflow, workflowservice::v1::{ - count_workflow_executions_response, workflow_service_client::WorkflowServiceClient, - *, + count_workflow_executions_response, + execute_multi_operation_request::operation::Operation as MultiOperationRequest, + execute_multi_operation_response::response::Response as MultiOperationResponse, + workflow_service_client::WorkflowServiceClient, *, }, }, - utilities::decode_status_detail, }, search_attributes::{SearchAttributeError, SearchAttributeValue, SearchAttributes}, }; @@ -179,14 +191,6 @@ static TEMPORAL_NAMESPACE_HEADER_KEY: &str = "temporal-namespace"; /// Key used to communicate when a GRPC message is too large pub static MESSAGE_TOO_LARGE_KEY: &str = "message-too-large"; #[doc(hidden)] -/// Returns the violation, if `status` is the client proactively rejecting an outbound request for exceeding a -/// payload/memo error size limit. -pub fn payload_limit_violation_from( - status: &tonic::Status, -) -> Option<&temporalio_common::payload_limits::PayloadLimitViolation> { - std::error::Error::source(status).and_then(|src| src.downcast_ref()) -} -#[doc(hidden)] /// Key used to indicate a error was returned by the retryer because of the short-circuit predicate pub static ERROR_RETURNED_DUE_TO_SHORT_CIRCUIT: &str = "short-circuit"; @@ -420,6 +424,14 @@ impl Connection { } else { None }; + #[cfg(feature = "experimental")] + let payloads_warn_size = options.payload_limits.payloads_warn_size; + #[cfg(not(feature = "experimental"))] + let payloads_warn_size = options_structs::DEFAULT_PAYLOADS_WARN_SIZE; + #[cfg(feature = "experimental")] + let memo_warn_size = options.payload_limits.memo_warn_size; + #[cfg(not(feature = "experimental"))] + let memo_warn_size = options_structs::DEFAULT_MEMO_WARN_SIZE; Ok(Self { inner: Arc::new(ConnectionInner { service: svc_client, @@ -433,12 +445,9 @@ impl Connection { _dns_task: dns_task, payloads_warn_size: resolve_warn_threshold( "payloads_warn_size", - options.payload_limits.payloads_warn_size, - ), - memo_warn_size: resolve_warn_threshold( - "memo_warn_size", - options.payload_limits.memo_warn_size, + payloads_warn_size, ), + memo_warn_size: resolve_warn_threshold("memo_warn_size", memo_warn_size), }), }) } @@ -1085,9 +1094,12 @@ impl Client { /// Connect to a Temporal service and create a namespace-bound client, applying registered /// plugins to connection and client options in registration order. pub async fn connect( - mut connection_options: ConnectionOptions, + connection_options: ConnectionOptions, client_options: ClientOptions, ) -> Result { + #[cfg(feature = "experimental")] + let mut connection_options = connection_options; + #[cfg(feature = "experimental")] plugins::apply_connection_plugins(&client_options, &mut connection_options)?; let connection = Connection::connect(connection_options).await?; Ok(Self::new(connection, client_options)?) @@ -1097,7 +1109,10 @@ impl Client { /// /// Registered client plugins are applied here. Connection plugin hooks only run when using /// [`Client::connect`]. - pub fn new(connection: Connection, mut options: ClientOptions) -> Result { + pub fn new(connection: Connection, options: ClientOptions) -> Result { + #[cfg(feature = "experimental")] + let mut options = options; + #[cfg(feature = "experimental")] plugins::apply_client_plugins(&mut options)?; Ok(Client { connection, @@ -1150,6 +1165,92 @@ impl Client { WorkflowClientTrait::start_workflow(self, workflow, input, options).await } + /// Atomically signal a workflow as it starts. + /// + /// The workflow receives the signal before its first workflow task. + pub async fn signal_with_start_workflow( + &self, + workflow: W, + workflow_input: W::Input, + signal: S, + signal_input: S::Input, + options: WorkflowStartOptions, + ) -> Result, WorkflowStartError> + where + W: HasWorkflowDefinition, + W::Input: Send, + S: SignalDefinition, + S::Input: Send, + { + WorkflowClientTrait::signal_with_start_workflow( + self, + workflow, + workflow_input, + signal, + signal_input, + options, + ) + .await + } + + /// Start a workflow and send it an update as a single atomic operation. + /// + /// Returns once the update has been accepted by the workflow, yielding a + /// [`WorkflowUpdateHandle`] that can be used to wait for the update result. + pub async fn start_update_with_start_workflow( + &self, + workflow: W, + workflow_input: W::Input, + update: U, + update_input: U::Input, + options: WorkflowUpdateWithStartOptions, + ) -> Result, WorkflowUpdateWithStartError> + where + W: HasWorkflowDefinition, + W::Input: Send, + U: UpdateDefinition, + U::Input: Send, + { + WorkflowClientTrait::start_update_with_start_workflow( + self, + workflow, + workflow_input, + update, + update_input, + options, + ) + .await + } + + /// Start a workflow and send it an update as a single atomic operation, waiting for the + /// update to complete and returning its result. + /// + /// See [Client::start_update_with_start_workflow] for details on option requirements. + pub async fn execute_update_with_start_workflow( + &self, + workflow: W, + workflow_input: W::Input, + update: U, + update_input: U::Input, + options: WorkflowUpdateWithStartOptions, + ) -> Result + where + W: HasWorkflowDefinition, + W::Input: Send, + U: UpdateDefinition, + U::Input: Send, + { + WorkflowClientTrait::execute_update_with_start_workflow( + self, + workflow, + workflow_input, + update, + update_input, + options, + ) + .await + } + /// Get a handle to an existing workflow. /// /// For untyped access, use `get_workflow_handle::(...)`. @@ -1184,12 +1285,93 @@ impl Client { /// Get a handle to complete an activity asynchronously. /// /// An activity returning `ActivityError::WillCompleteAsync` can be completed with this handle. + /// + /// To get a handle to a standalone activity that can be used to wait for result and manage + /// the execution, see [`get_activity_handle`](Self::get_activity_handle). pub fn get_async_activity_handle( &self, identifier: ActivityIdentifier, ) -> AsyncActivityHandle { WorkflowClientTrait::get_async_activity_handle(self, identifier) } + + /// Start a standalone activity. + /// + /// Returns [`ActivityHandle`] that can be used to wait for result or to perform other + /// operations on the activity. + pub async fn start_activity( + &self, + activity: A, + input: A::Input, + options: ActivityStartOptions, + ) -> Result, StartActivityError> + where + A: ActivityDefinition, + { + WorkflowClientTrait::start_activity(self, activity, input, options).await + } + + /// Get a handle to an existing standalone activity execution. If `run_id` is not specified, + /// the handle always targets the latest execution with matching ID. + /// + /// Note that the validity of the handle is not checked until a method is called on it. + /// If invalid ID or run ID is used, the method will return `NotFound` error. + /// + /// To get an untyped handle, use [`get_untyped_activity_handle`](Self::get_untyped_activity_handle). + /// + /// To get a handle that can be used to complete an activity asynchronously, + /// see [`get_async_activity_handle`](Self::get_async_activity_handle). + pub fn get_activity_handle( + &self, + activity: A, + id: impl Into, + run_id: Option, + ) -> ActivityHandle + where + Self: Sized, + A: ActivityDefinition, + { + WorkflowClientTrait::get_activity_handle(self, activity, id, run_id) + } + + /// Get an untyped handle to an existing standalone activity execution. If `run_id` is not + /// specified, the handle always targets the latest execution with matching ID. + /// + /// Note that the validity of the handle is not checked until a method is called on it. + /// If invalid ID or run ID is used, the method will return `NotFound` error. + /// + /// To get a typed handle, use [`get_activity_handle`](Self::get_activity_handle). + /// + /// To get a handle that can be used to complete an activity asynchronously, + /// see [`get_async_activity_handle`](Self::get_async_activity_handle). + pub fn get_untyped_activity_handle( + &self, + id: impl Into, + run_id: Option, + ) -> ActivityHandle + where + Self: Sized, + { + WorkflowClientTrait::get_untyped_activity_handle(self, id, run_id) + } + + /// List activities matching a query. Returns a stream that lazily paginates through results. + pub fn list_activities( + &self, + query: impl Into, + options: ActivityListOptions, + ) -> ListActivitiesStream { + WorkflowClientTrait::list_activities(self, query, options) + } + + /// Count activities matching a query. + pub async fn count_activities( + &self, + query: impl Into, + options: ActivityCountOptions, + ) -> Result { + WorkflowClientTrait::count_activities(self, query, options).await + } } impl NamespacedClient for Client { @@ -1234,6 +1416,56 @@ pub(crate) trait WorkflowClientTrait: NamespacedClient { W: HasWorkflowDefinition, W::Input: Send; + /// Start a workflow and atomically send it a signal. + fn signal_with_start_workflow( + &self, + workflow: W, + workflow_input: W::Input, + signal: S, + signal_input: S::Input, + options: WorkflowStartOptions, + ) -> impl Future, WorkflowStartError>> + where + Self: Sized, + W: HasWorkflowDefinition, + W::Input: Send, + S: SignalDefinition, + S::Input: Send; + + /// Start a workflow and send it an update as a single atomic operation, returning once the + /// update reaches the requested wait stage. + fn start_update_with_start_workflow( + &self, + workflow: W, + workflow_input: W::Input, + update: U, + update_input: U::Input, + options: WorkflowUpdateWithStartOptions, + ) -> impl Future, WorkflowUpdateWithStartError>> + where + Self: Sized, + W: HasWorkflowDefinition, + W::Input: Send, + U: UpdateDefinition, + U::Input: Send; + + /// Start a workflow and send it an update as a single atomic operation, waiting for the + /// update to complete and returning its result. + fn execute_update_with_start_workflow( + &self, + workflow: W, + workflow_input: W::Input, + update: U, + update_input: U::Input, + options: WorkflowUpdateWithStartOptions, + ) -> impl Future> + where + Self: Sized, + W: HasWorkflowDefinition, + W::Input: Send, + U: UpdateDefinition, + U::Input: Send; + /// Get a handle to an existing workflow. `run_id` may be left blank to specify the most recent /// execution having the provided `workflow_id`. /// @@ -1272,6 +1504,51 @@ pub(crate) trait WorkflowClientTrait: NamespacedClient { ) -> AsyncActivityHandle where Self: Sized; + + /// Start a standalone activity. + fn start_activity( + &self, + activity: A, + input: A::Input, + options: ActivityStartOptions, + ) -> impl Future, StartActivityError>> + where + Self: Sized, + A: ActivityDefinition; + + /// Get a handle to a previously started standalone activity. + fn get_activity_handle( + &self, + activity: A, + id: impl Into, + run_id: Option, + ) -> ActivityHandle + where + Self: Sized, + A: ActivityDefinition; + + /// Get an untyped handle to a previously started standalone activity. + fn get_untyped_activity_handle( + &self, + id: impl Into, + run_id: Option, + ) -> ActivityHandle + where + Self: Sized; + + /// List activities matching a query. Returns a stream that lazily paginates through results. + fn list_activities( + &self, + query: impl Into, + _options: ActivityListOptions, + ) -> ListActivitiesStream; + + /// Count activities matching a query. + fn count_activities( + &self, + query: impl Into, + _options: ActivityCountOptions, + ) -> impl Future>; } /// A client that is bound to a namespace @@ -1388,7 +1665,7 @@ impl WorkflowExecution { Memo::from_raw( self.raw.memo.clone(), self.data_converter.payload_converter().clone(), - SerializationContextData::Workflow, + SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) } @@ -1524,6 +1801,59 @@ impl WorkflowCountAggregationGroup { } } +// Keep the common fields used by start RPC variants in one place so their option handling does +// not drift as new fields are added. +fn build_start_workflow_request( + client: &impl NamespacedClient, + workflow_type: String, + input: Option, + memo: Option, + options: WorkflowStartOptions, +) -> StartWorkflowExecutionRequest { + let user_metadata = options.user_metadata(); + let request_eager_execution = options.enable_eager_workflow_start; + StartWorkflowExecutionRequest { + namespace: client.namespace(), + input, + workflow_id: options.workflow_id, + workflow_type: Some(WorkflowType { + name: workflow_type, + }), + task_queue: Some(TaskQueue { + name: options.task_queue, + kind: TaskQueueKind::Unspecified as i32, + normal_name: String::new(), + }), + identity: client.identity(), + request_id: Uuid::new_v4().to_string(), + workflow_id_reuse_policy: ProtoWorkflowIdReusePolicy::from(options.id_reuse_policy) as i32, + workflow_id_conflict_policy: ProtoWorkflowIdConflictPolicy::from(options.id_conflict_policy) + as i32, + workflow_execution_timeout: options + .execution_timeout + .and_then(|duration| duration.try_into().ok()), + workflow_run_timeout: options + .run_timeout + .and_then(|duration| duration.try_into().ok()), + workflow_task_timeout: options + .task_timeout + .and_then(|duration| duration.try_into().ok()), + search_attributes: options + .search_attributes + .map(|attributes| attributes.into_proto()), + cron_schedule: options.cron_schedule.unwrap_or_default(), + request_eager_execution, + retry_policy: options.retry_policy.map(Into::into), + links: options.links, + completion_callbacks: options.completion_callbacks, + priority: Some(options.priority.into()), + memo, + header: options.header, + user_metadata, + ..Default::default() + } +} + impl WorkflowClientTrait for T where T: WorkflowService + NamespacedClient + Clone + Send + Sync + 'static, @@ -1554,155 +1884,168 @@ where let data_converter = client.data_converter().clone(); let unencoded_payloads = { let payload_converter = data_converter.payload_converter(); - let context = SerializationContext { - data: &SerializationContextData::Workflow, - converter: payload_converter, - }; + let context_data = SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ); + let context = + SerializationContext::new(&context_data, payload_converter); args.serialize_payloads(&context) }; drop(args); let payloads = data_converter .codec() - .encode(&SerializationContextData::Workflow, unencoded_payloads?) + .encode(&SerializationContextData::Workflow(WorkflowSerializationContext::new()), unencoded_payloads?) .await?; - let namespace = client.namespace(); let workflow_id = options.workflow_id.clone(); - let task_queue_name = options.task_queue.clone(); - - let user_metadata = if options.static_summary.is_some() - || options.static_details.is_some() - { - let payload_converter = PayloadConverter::default(); - let context = SerializationContext { - data: &SerializationContextData::Workflow, - converter: &payload_converter, - }; - Some(UserMetadata { - summary: options.static_summary.map(|summary| { - payload_converter.to_payload(&context, &summary).expect( - "String-to-JSON payload serialization is infallible", - ) - }), - details: options.static_details.map(|details| { - payload_converter.to_payload(&context, &details).expect( - "String-to-JSON payload serialization is infallible", - ) - }), - }) - } else { - None - }; - - let run_id = if let Some(start_signal) = options.start_signal { - let mut request = SignalWithStartWorkflowExecutionRequest { - namespace, - workflow_id: workflow_id.clone(), - workflow_type: Some(WorkflowType { - name: workflow_type, - }), - task_queue: Some(TaskQueue { - name: task_queue_name, - kind: TaskQueueKind::Normal as i32, - normal_name: String::new(), - }), - input: payloads.into_payloads(), - signal_name: start_signal.signal_name, - signal_input: start_signal.input, - identity: client.identity(), - request_id: Uuid::new_v4().to_string(), - workflow_id_reuse_policy: options.id_reuse_policy as i32, - workflow_id_conflict_policy: options.id_conflict_policy as i32, - workflow_execution_timeout: options - .execution_timeout - .and_then(|duration| duration.try_into().ok()), - workflow_run_timeout: options - .run_timeout - .and_then(|duration| duration.try_into().ok()), - workflow_task_timeout: options - .task_timeout - .and_then(|duration| duration.try_into().ok()), - search_attributes: options - .search_attributes - .map(|attributes| attributes.into_proto()), - cron_schedule: options.cron_schedule.unwrap_or_default(), - retry_policy: options.retry_policy.map(Into::into), - header: options.header.or(start_signal.header), - user_metadata, - ..Default::default() - } - .into_request(); - rpc_options.apply_to(&mut request); - WorkflowService::signal_with_start_workflow_execution( - &mut client, - request, - ) - .await? + let memo = options.encoded_memo(&data_converter).await?; + let mut request = build_start_workflow_request( + &client, + workflow_type, + payloads.into_payloads(), + memo, + options, + ) + .into_request(); + rpc_options.apply_to(&mut request); + let run_id = client + .start_workflow_execution(request) + .await + .map_err(WorkflowStartError::from_status)? .into_inner() - .run_id - } else { - let mut request = StartWorkflowExecutionRequest { - namespace, - input: payloads.into_payloads(), - workflow_id: workflow_id.clone(), - workflow_type: Some(WorkflowType { - name: workflow_type, - }), - task_queue: Some(TaskQueue { - name: task_queue_name, - kind: TaskQueueKind::Unspecified as i32, - normal_name: String::new(), - }), - request_id: Uuid::new_v4().to_string(), - workflow_id_reuse_policy: options.id_reuse_policy as i32, - workflow_id_conflict_policy: options.id_conflict_policy as i32, - workflow_execution_timeout: options - .execution_timeout - .and_then(|duration| duration.try_into().ok()), - workflow_run_timeout: options - .run_timeout - .and_then(|duration| duration.try_into().ok()), - workflow_task_timeout: options - .task_timeout - .and_then(|duration| duration.try_into().ok()), - search_attributes: options - .search_attributes - .map(|attributes| attributes.into_proto()), - cron_schedule: options.cron_schedule.unwrap_or_default(), - request_eager_execution: options.enable_eager_workflow_start, - retry_policy: options.retry_policy.map(Into::into), - links: options.links, - completion_callbacks: options.completion_callbacks, - priority: Some(options.priority.into()), - header: options.header, - user_metadata, - ..Default::default() - } - .into_request(); - rpc_options.apply_to(&mut request); - client - .start_workflow_execution(request) - .await - .map_err(|status| { - if status.code() == Code::AlreadyExists { - let run_id = decode_status_detail::< - WorkflowExecutionAlreadyStartedFailure, - >( - status.details() - ) - .map(|failure| failure.run_id); - WorkflowStartError::AlreadyStarted { - run_id, - source: status, - } - } else { - WorkflowStartError::Rpc(status) - } - })? - .into_inner() - .run_id - }; + .run_id; + + Ok(StartWorkflowOutput::new(workflow_id, run_id)) + }) + } + }), + ) + .await?; + let StartWorkflowOutput { + workflow_id, + run_id, + } = interceptor_output; + + Ok(WorkflowHandle::new( + self.clone(), + WorkflowExecutionInfo { + namespace, + workflow_id, + run_id: Some(run_id.clone()), + first_execution_run_id: Some(run_id), + }, + )) + } + async fn signal_with_start_workflow( + &self, + workflow: W, + workflow_input: W::Input, + signal: S, + signal_input: S::Input, + options: WorkflowStartOptions, + ) -> Result, WorkflowStartError> + where + W: HasWorkflowDefinition, + W::Input: Send, + S: SignalDefinition, + S::Input: Send, + { + let namespace = self.namespace(); + let interceptor_output = interceptors::call_signal_with_start_workflow( + self.client_interceptors(), + SignalWithStartWorkflowInput::new( + workflow.name().to_owned(), + workflow_input, + signal.name().to_owned(), + signal_input, + options, + ), + Next::new({ + let client = (*self).clone(); + move |input: SignalWithStartWorkflowInput| -> BoxFuture< + '_, + Result, + > { + let mut client = client; + Box::pin(async move { + let ( + workflow_type, + workflow_args, + signal_name, + signal_args, + options, + rpc_options, + ) = input.into_parts(); + let data_converter = client.data_converter().clone(); + let payload_converter = data_converter.payload_converter(); + let context_data = SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ); + let context = SerializationContext::new(&context_data, payload_converter); + let workflow_payloads = workflow_args.serialize_payloads(&context); + let signal_payloads = signal_args.serialize_payloads(&context); + drop(workflow_args); + drop(signal_args); + let workflow_payloads = data_converter + .codec() + .encode(&SerializationContextData::Workflow(WorkflowSerializationContext::new()), workflow_payloads?) + .await?; + let signal_payloads = data_converter + .codec() + .encode(&SerializationContextData::Workflow(WorkflowSerializationContext::new()), signal_payloads?) + .await?; + let workflow_id = options.workflow_id.clone(); + let memo = options.encoded_memo(&data_converter).await?; + let mut start_request = build_start_workflow_request( + &client, + workflow_type, + workflow_payloads.into_payloads(), + memo, + options, + ); + if let Some(task_queue) = &mut start_request.task_queue { + task_queue.kind = TaskQueueKind::Normal as i32; + } + let mut request = SignalWithStartWorkflowExecutionRequest { + namespace: start_request.namespace, + workflow_id: start_request.workflow_id, + workflow_type: start_request.workflow_type, + task_queue: start_request.task_queue, + input: start_request.input, + workflow_execution_timeout: start_request.workflow_execution_timeout, + workflow_run_timeout: start_request.workflow_run_timeout, + workflow_task_timeout: start_request.workflow_task_timeout, + identity: start_request.identity, + request_id: start_request.request_id, + workflow_id_reuse_policy: start_request.workflow_id_reuse_policy, + workflow_id_conflict_policy: start_request.workflow_id_conflict_policy, + signal_name, + signal_input: Some(Payloads { + payloads: signal_payloads, + }), + retry_policy: start_request.retry_policy, + cron_schedule: start_request.cron_schedule, + memo: start_request.memo, + search_attributes: start_request.search_attributes, + header: start_request.header, + workflow_start_delay: start_request.workflow_start_delay, + user_metadata: start_request.user_metadata, + links: start_request.links, + versioning_override: start_request.versioning_override, + priority: start_request.priority, + time_skipping_config: start_request.time_skipping_config, + ..Default::default() + } + .into_request(); + rpc_options.apply_to(&mut request); + let run_id = WorkflowService::signal_with_start_workflow_execution( + &mut client, + request, + ) + .await? + .into_inner() + .run_id; Ok(StartWorkflowOutput::new(workflow_id, run_id)) }) } @@ -1725,6 +2068,217 @@ where )) } + async fn start_update_with_start_workflow( + &self, + workflow: W, + workflow_input: W::Input, + update: U, + update_input: U::Input, + options: WorkflowUpdateWithStartOptions, + ) -> Result, WorkflowUpdateWithStartError> + where + W: HasWorkflowDefinition, + W::Input: Send, + U: UpdateDefinition, + U::Input: Send, + { + let output = interceptors::call_update_with_start_workflow( + self.client_interceptors(), + UpdateWithStartWorkflowInput::new( + workflow.name().to_owned(), + workflow_input, + update.name().to_owned(), + update_input, + options, + ), + Next::new({ + let client = (*self).clone(); + move |input: UpdateWithStartWorkflowInput| -> BoxFuture< + '_, + Result, + > { + let mut client = client; + Box::pin(async move { + let UpdateWithStartWorkflowInput { + workflow_type, + update_name, + options, + rpc_options, + workflow_args, + update_args, + } = input; + let (start_options, update_id, update_header) = options.into_parts(); + + let data_converter = client.data_converter().clone(); + let (unencoded_workflow_payloads, unencoded_update_payloads) = { + let payload_converter = data_converter.payload_converter(); + let context_data = SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ); + let context = + SerializationContext::new(&context_data, payload_converter); + ( + workflow_args.serialize_payloads(&context), + update_args.serialize_payloads(&context), + ) + }; + drop(workflow_args); + drop(update_args); + // The codec may do expensive work per call (e.g. remote encryption), so + // encode both payload sets concurrently. + let (workflow_payloads, update_payloads) = try_join( + data_converter.codec().encode( + &SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ), + unencoded_workflow_payloads?, + ), + data_converter.codec().encode( + &SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ), + unencoded_update_payloads?, + ), + ) + .await?; + + let namespace = client.namespace(); + let workflow_id = start_options.workflow_id.clone(); + let memo = start_options.encoded_memo(&data_converter).await?; + let start_request = build_start_workflow_request( + &client, + workflow_type, + workflow_payloads.into_payloads(), + memo, + start_options, + ); + + let update_id = update_id.unwrap_or_else(|| Uuid::new_v4().to_string()); + let update_request = workflow_handle::build_update_workflow_request( + namespace.clone(), + client.identity(), + workflow_id.clone(), + String::new(), + update_id.clone(), + update_name, + update_header, + update_payloads, + ); + + let request = ExecuteMultiOperationRequest { + namespace, + operations: vec![ + execute_multi_operation_request::Operation { + operation: Some(MultiOperationRequest::StartWorkflow( + start_request, + )), + }, + execute_multi_operation_request::Operation { + operation: Some(MultiOperationRequest::UpdateWorkflow( + update_request, + )), + }, + ], + resource_id: workflow_id.clone(), + }; + + let (start_response, update_response) = loop { + let mut rpc_request = request.clone().into_request(); + rpc_options.apply_to(&mut rpc_request); + let response = + WorkflowService::execute_multi_operation(&mut client, rpc_request) + .await + .map_err(WorkflowUpdateWithStartError::from_status)? + .into_inner(); + + let [start_response, update_response]: [_; 2] = + response.responses.try_into().map_err(|_| { + WorkflowUpdateWithStartError::Other( + "Server response did not include exactly two operation \ + responses" + .into(), + ) + })?; + let ( + Some(MultiOperationResponse::StartWorkflow(start_response)), + Some(MultiOperationResponse::UpdateWorkflow(update_response)), + ) = (start_response.response, update_response.response) + else { + return Err(WorkflowUpdateWithStartError::Other( + "Server response did not include start and update operation \ + responses in request order" + .into(), + )); + }; + + if update_response.stage + < UpdateWorkflowExecutionLifecycleStage::Accepted as i32 + { + continue; + } + break (start_response, update_response); + }; + + let run_id = update_response + .update_ref + .as_ref() + .and_then(|reference| reference.workflow_execution.as_ref()) + .map(|execution| execution.run_id.clone()) + .filter(|run_id| !run_id.is_empty()) + .or_else(|| { + (!start_response.run_id.is_empty()).then_some(start_response.run_id) + }); + Ok(UpdateWithStartWorkflowOutput::new( + workflow_id, + update_id, + run_id, + update_response.outcome, + )) + }) + } + }), + ) + .await?; + Ok(WorkflowUpdateHandle::new( + self.clone(), + output.update_id, + output.workflow_id, + output.run_id, + output.known_outcome, + )) + } + + async fn execute_update_with_start_workflow( + &self, + workflow: W, + workflow_input: W::Input, + update: U, + update_input: U::Input, + options: WorkflowUpdateWithStartOptions, + ) -> Result + where + W: HasWorkflowDefinition, + W::Input: Send, + U: UpdateDefinition, + U::Input: Send, + { + let rpc_options = options.rpc_options.clone(); + let update_handle = WorkflowClientTrait::start_update_with_start_workflow( + self, + workflow, + workflow_input, + update, + update_input, + options, + ) + .await?; + let result = update_handle + .get_result(rpc_options) + .await + .map_err(WorkflowUpdateWithStartError::Update)?; + Ok(result) + } + fn get_workflow_handle( &self, workflow_id: impl Into, @@ -1830,7 +2384,9 @@ where && let Err(err) = decode_payloads( memo, data_converter.codec(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ), ) .await { @@ -1913,17 +2469,205 @@ where { AsyncActivityHandle::new(self.clone(), identifier) } -} -macro_rules! dbg_panic { - ($($arg:tt)*) => { - use tracing::error; + async fn start_activity( + &self, + activity: A, + input: A::Input, + options: ActivityStartOptions, + ) -> Result, StartActivityError> + where + Self: Sized, + A: ActivityDefinition, + { + let mut client = self.clone(); + let dc = client.data_converter(); + let sc = &SerializationContextData::Activity(ActivitySerializationContext::new()); + + let user_metadata = { + let summary = match &options.summary { + Some(summary) => Some(dc.to_payload(sc, summary).await?), + None => None, + }; + let details = match &options.static_details { + Some(details) => Some(dc.to_payload(sc, details).await?), + None => None, + }; + (summary.is_some() || details.is_some()).then_some(UserMetadata { summary, details }) + }; + + let resp = client + .start_activity_execution( + StartActivityExecutionRequest { + namespace: client.namespace(), + identity: client.identity(), + request_id: Uuid::new_v4().to_string(), + activity_id: options.id.clone(), + activity_type: Some(ActivityType { + name: activity.name().to_string(), + }), + task_queue: Some(TaskQueue { + name: options.task_queue, + kind: TaskQueueKind::Normal.into(), + normal_name: "".to_string(), + }), + schedule_to_close_timeout: try_into_or_box_err( + options.close_timeouts.schedule_to_close(), + StartActivityError::Other, + )?, + schedule_to_start_timeout: try_into_or_box_err( + options.schedule_to_start_timeout, + StartActivityError::Other, + )?, + start_to_close_timeout: try_into_or_box_err( + options.close_timeouts.start_to_close(), + StartActivityError::Other, + )?, + heartbeat_timeout: try_into_or_box_err( + options.heartbeat_timeout, + StartActivityError::Other, + )?, + retry_policy: options.retry_policy.map(Into::into), + input: dc.to_payloads(sc, &input).await?.into_payloads(), + id_reuse_policy: ProtoActivityIdReusePolicy::from(options.id_reuse_policy) + .into(), + id_conflict_policy: ProtoActivityIdConflictPolicy::from( + options.id_conflict_policy, + ) + .into(), + search_attributes: options.search_attributes.map(SearchAttributes::into_proto), + header: options.header, + user_metadata, + priority: Some(options.priority.into()), + start_delay: try_into_or_box_err( + options.start_delay, + StartActivityError::Other, + )?, + ..Default::default() + } + .into_request(), + ) + .await? + .into_inner(); + + Ok(ActivityHandle::new( + client, + options.id, + (!resp.run_id.is_empty()).then_some(resp.run_id), + )) + } + + fn get_activity_handle( + &self, + _activity: A, + id: impl Into, + run_id: Option, + ) -> ActivityHandle + where + Self: Sized, + A: ActivityDefinition, + { + ActivityHandle::new(self.clone(), id.into(), run_id) + } + + fn get_untyped_activity_handle( + &self, + id: impl Into, + run_id: Option, + ) -> ActivityHandle + where + Self: Sized, + { + ActivityHandle::new(self.clone(), id.into(), run_id) + } + + fn list_activities( + &self, + query: impl Into, + _options: ActivityListOptions, + ) -> ListActivitiesStream { + let client = self.clone(); + let namespace = client.namespace(); + let query = query.into(); + + ListActivitiesStream::new(stream::unfold( + Some(vec![]), // empty token for initial query, None if done + move |next_page_token| { + let mut client = client.clone(); + let namespace = namespace.clone(); + let query = query.clone(); + + async move { + // making it more visible that we're terminating stream here + #[allow(clippy::question_mark)] + let Some(token): Option> = next_page_token else { + return None; + }; + + match WorkflowService::list_activity_executions( + &mut client, + ListActivityExecutionsRequest { + namespace, + page_size: 0, // Use server default + next_page_token: token.clone(), + query, + } + .into_request(), + ) + .await + .map(|r| r.into_inner()) + { + Ok(resp) => Some(( + Ok(resp.executions), + (!resp.next_page_token.is_empty()).then_some(resp.next_page_token), + )), + Err(e) => Some((Err(e.into()), Some(token))), + } + } + }, + )) + } + + async fn count_activities( + &self, + query: impl Into, + _options: ActivityCountOptions, + ) -> Result { + let mut client = self.clone(); + let resp = client + .count_activity_executions( + CountActivityExecutionsRequest { + namespace: client.namespace(), + query: query.into(), + } + .into_request(), + ) + .await? + .into_inner(); + Ok(ActivityExecutionCount::from_response(resp)) + } +} + +macro_rules! dbg_panic { + ($($arg:tt)*) => { + use tracing::error; error!($($arg)*); debug_assert!(false, $($arg)*); }; } pub(crate) use dbg_panic; +fn try_into_or_box_err(val: Option, map_err: MapErr) -> Result, E> +where + A: TryInto, + >::Error: Error + Send + Sync + 'static, + MapErr: FnOnce(Box) -> E, +{ + val.map(TryInto::try_into) + .transpose() + .map_err(|e| map_err(Box::from(e))) +} + #[cfg(test)] mod tests { use super::*; @@ -2531,39 +3275,53 @@ mod tests { mod start_workflow_interceptor_tests { use super::*; - use crate::request_extensions::RetryConfigForCall; + use crate::{request_extensions::RetryConfigForCall, test_helpers::XorCodec}; use parking_lot::Mutex; use std::sync::atomic::{AtomicUsize, Ordering}; use temporalio_common::{ - HasWorkflowDefinition, WorkflowDefinition, + MemoValues, SignalDefinition, data_converters::{ - DefaultFailureConverter, PayloadCodec, PayloadConversionError, - SerializationContext, SerializationContextData, TemporalSerializable, + DefaultFailureConverter, PayloadCodec, PayloadConversionError, PayloadConverter, + SerializationContext, SerializationContextData, TemporalDeserializable, + TemporalSerializable, + }, + protos::temporal::api::common::v1::{ + Link, Memo as ProtoMemo, Payload, Priority as ProtoPriority, }, - protos::temporal::api::common::v1::Payload, }; + use temporalio_macros::{workflow, workflow_methods}; + use temporalio_workflow::{SyncWorkflowContext, WorkflowContext, WorkflowResult}; use tonic::{Request, Response}; + #[workflow] + #[derive(Default)] struct TestWorkflow; - impl WorkflowDefinition for TestWorkflow { - type Input = Vec; - type Output = (); - - fn name(&self) -> &str { - "test-workflow" + #[workflow_methods] + impl TestWorkflow { + #[run] + async fn run( + _ctx: &mut WorkflowContext, + _input: Vec, + ) -> WorkflowResult<()> { + Ok(()) } - } - impl HasWorkflowDefinition for TestWorkflow { - type Run = Self; + #[signal] + fn test_signal(&mut self, _ctx: &mut SyncWorkflowContext, _input: Vec) {} } #[derive(Default)] struct RecordedStart { calls: usize, workflow_type: String, + memo: Option, payloads: Vec, + signal_name: String, + signal_payloads: Vec, + identity: String, + links: Vec, + priority: Option, ascii_metadata: Option, binary_metadata: Option>, grpc_timeout: Option, @@ -2647,7 +3405,11 @@ mod tests { let mut recorded = self.recorded.lock(); recorded.calls += 1; recorded.workflow_type = request.workflow_type.unwrap().name; + recorded.memo = request.memo; recorded.payloads = request.input.unwrap_or_default().payloads; + recorded.identity = request.identity; + recorded.links = request.links; + recorded.priority = request.priority; recorded.ascii_metadata = ascii_metadata; recorded.binary_metadata = binary_metadata; recorded.grpc_timeout = grpc_timeout; @@ -2688,7 +3450,13 @@ mod tests { let mut recorded = self.recorded.lock(); recorded.calls += 1; recorded.workflow_type = request.workflow_type.unwrap().name; + recorded.memo = request.memo; recorded.payloads = request.input.unwrap_or_default().payloads; + recorded.signal_name = request.signal_name; + recorded.signal_payloads = request.signal_input.unwrap_or_default().payloads; + recorded.identity = request.identity; + recorded.links = request.links; + recorded.priority = request.priority; recorded.ascii_metadata = ascii_metadata; recorded.binary_metadata = binary_metadata; recorded.grpc_timeout = grpc_timeout; @@ -2852,6 +3620,56 @@ mod tests { } } + struct ReplacingSignalWithStartInterceptor; + + impl ClientInterceptor for ReplacingSignalWithStartInterceptor { + fn signal_with_start_workflow<'a>( + &'a self, + mut input: SignalWithStartWorkflowInput, + next: Next< + 'a, + SignalWithStartWorkflowInput, + BoxFuture<'a, Result>, + >, + ) -> BoxFuture<'a, Result> { + assert_eq!( + input.workflow_args_ref::>().unwrap(), + &["workflow".to_owned()] + ); + assert_eq!( + input.signal_args_ref::>().unwrap(), + &["signal".to_owned()] + ); + input.replace_workflow_args(vec!["replaced-workflow".to_owned()]); + input.replace_signal_args(vec!["replaced-signal".to_owned()]); + next.run(input) + } + } + + struct FailingSignal; + + impl SignalDefinition for FailingSignal { + type Workflow = test_workflow::Run; + type Input = FailingSignalInput; + + fn name(&self) -> &str { + "failing-signal" + } + } + + struct FailingSignalInput; + + impl TemporalDeserializable for FailingSignalInput {} + + impl TemporalSerializable for FailingSignalInput { + fn to_payloads( + &self, + _context: &SerializationContext<'_>, + ) -> Result, PayloadConversionError> { + Err(PayloadConversionError::WrongEncoding) + } + } + fn mock_client( interceptors: Vec>, encode_calls: Arc, @@ -2859,7 +3677,7 @@ mod tests { let recorded = Arc::new(Mutex::new(RecordedStart::default())); let data_converter = DataConverter::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), CountingCodec { encode_calls: encode_calls.clone(), }, @@ -2876,6 +3694,153 @@ mod tests { ) } + /// A mock client whose data converter uses `codec`, for asserting on what reaches the + /// wire. + fn mock_client_with_codec( + codec: impl PayloadCodec + Send + Sync + 'static, + ) -> (MockStartWorkflowClient, Arc>) { + let recorded = Arc::new(Mutex::new(RecordedStart::default())); + let data_converter = DataConverter::new( + PayloadConverter::default(), + DefaultFailureConverter::default(), + codec, + ); + ( + MockStartWorkflowClient { + recorded: recorded.clone(), + data_converter, + }, + recorded, + ) + } + + /// Decode a sent memo the same way `describe`/`list` do, and read it back. + async fn read_back(sent: ProtoMemo) -> Memo { + let mut sent = sent; + decode_payloads( + &mut sent, + &XorCodec, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + ) + .await + .unwrap(); + Memo::from_raw( + Some(sent), + PayloadConverter::default(), + SerializationContextData::Workflow(WorkflowSerializationContext::new()), + ) + } + + #[tokio::test] + async fn start_workflow_encodes_memo_with_payload_converter_and_codec() { + let (client, recorded) = mock_client_with_codec(XorCodec); + let mut memo = MemoValues::new(); + memo.insert("memo-key", "memo-value".to_owned()); + + client + .start_workflow( + TestWorkflow::run, + vec!["initial".to_owned()], + WorkflowStartOptions::new("task-queue", "workflow-id") + .memo(memo) + .build(), + ) + .await + .unwrap(); + + let sent = recorded.lock().memo.clone().expect("memo should be sent"); + assert_eq!( + read_back(sent).await.get::("memo-key").unwrap(), + Some("memo-value".to_owned()) + ); + } + + #[tokio::test] + async fn signal_with_start_workflow_encodes_memo() { + let (client, recorded) = mock_client_with_codec(XorCodec); + let mut memo = MemoValues::new(); + memo.insert("memo-key", "memo-value".to_owned()); + + client + .signal_with_start_workflow( + TestWorkflow::run, + vec!["initial".to_owned()], + TestWorkflow::test_signal, + vec!["signal".to_owned()], + WorkflowStartOptions::new("task-queue", "workflow-id") + .memo(memo) + .build(), + ) + .await + .unwrap(); + + let sent = recorded.lock().memo.clone().expect("memo should be sent"); + assert_eq!( + read_back(sent).await.get::("memo-key").unwrap(), + Some("memo-value".to_owned()) + ); + } + + #[tokio::test] + async fn start_workflow_without_memo_sends_none() { + let (client, recorded) = mock_client_with_codec(XorCodec); + + client + .start_workflow( + TestWorkflow::run, + vec!["initial".to_owned()], + WorkflowStartOptions::new("task-queue", "workflow-id").build(), + ) + .await + .unwrap(); + + assert_eq!(recorded.lock().memo, None); + } + + #[tokio::test] + async fn start_workflow_reports_memo_serialization_errors() { + #[derive(Debug)] + struct FailingMemoValue; + + impl TemporalSerializable for FailingMemoValue { + fn to_payload( + &self, + _ctx: &SerializationContext<'_>, + ) -> Result { + Err(PayloadConversionError::EncodingError( + std::io::Error::other("memo serialization failure").into(), + )) + } + } + + let (client, recorded) = mock_client_with_codec(XorCodec); + let mut memo = MemoValues::new(); + memo.insert("invalid", FailingMemoValue); + + let err = client + .start_workflow( + TestWorkflow::run, + vec!["initial".to_owned()], + WorkflowStartOptions::new("task-queue", "workflow-id") + .memo(memo) + .build(), + ) + .await + .map(|_| ()) + .expect_err("memo serialization errors should be surfaced"); + + assert!( + matches!(err, WorkflowStartError::PayloadConversion(_)), + "expected a payload conversion error, got {err:?}" + ); + assert!( + err.to_string().contains("memo serialization failure"), + "error should surface the underlying cause, got {err}" + ); + // The request must not have been sent. + assert_eq!(recorded.lock().calls, 0); + } + #[tokio::test] async fn interceptors_order_mutate_replace_and_defer_conversion() { let events = Arc::new(Mutex::new(Vec::new())); @@ -2896,7 +3861,7 @@ mod tests { let handle = client .start_workflow( - TestWorkflow, + TestWorkflow::run, vec!["initial".to_owned()], WorkflowStartOptions::new("task-queue", "workflow-id").build(), ) @@ -2917,7 +3882,10 @@ mod tests { }; let replacement: String = client .data_converter() - .from_payloads(&SerializationContextData::Workflow, payloads) + .from_payloads( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + payloads, + ) .await .unwrap(); assert_eq!(replacement, "replacement"); @@ -2932,7 +3900,7 @@ mod tests { ); let handle = client .start_workflow( - TestWorkflow, + TestWorkflow::run, vec!["initial".to_owned()], WorkflowStartOptions::new("task-queue", "ignored-workflow-id").build(), ) @@ -2952,7 +3920,7 @@ mod tests { let recorded = Arc::new(Mutex::new(RecordedStart::default())); let data_converter = DataConverter::new( PayloadConverter::UseWrappers, - DefaultFailureConverter, + DefaultFailureConverter::default(), CountingCodec { encode_calls: encode_calls.clone(), }, @@ -2969,7 +3937,7 @@ mod tests { client .start_workflow( - TestWorkflow, + TestWorkflow::run, vec!["initial".to_owned()], WorkflowStartOptions::new("task-queue", "workflow-id").build(), ) @@ -2992,7 +3960,7 @@ mod tests { client .start_workflow( - TestWorkflow, + TestWorkflow::run, vec!["initial".to_owned()], WorkflowStartOptions::new("task-queue", "workflow-id").build(), ) @@ -3018,7 +3986,7 @@ mod tests { options.rpc_options = rpc_options.clone(); client - .start_workflow(TestWorkflow, vec!["initial".to_owned()], options) + .start_workflow(TestWorkflow::run, vec!["initial".to_owned()], options) .await .unwrap(); @@ -3031,10 +3999,15 @@ mod tests { } let mut options = WorkflowStartOptions::new("task-queue", "signal-workflow-id").build(); - options.start_signal = Some(WorkflowStartSignal::new("signal-name").build()); options.rpc_options = rpc_options; let handle = client - .start_workflow(TestWorkflow, vec!["initial".to_owned()], options) + .signal_with_start_workflow( + TestWorkflow::run, + vec!["initial".to_owned()], + TestWorkflow::test_signal, + vec!["signal".to_owned()], + options, + ) .await .unwrap(); @@ -3044,9 +4017,77 @@ mod tests { assert_eq!(recorded.binary_metadata.as_deref(), Some(&[0, 255][..])); assert_eq!(recorded.grpc_timeout.as_deref(), Some("250000u")); assert_eq!(recorded.retry_options, Some(RetryOptions::no_retries())); + assert_eq!(recorded.signal_name, "test_signal"); + assert_eq!(recorded.signal_payloads.len(), 1); assert_eq!(handle.run_id(), Some("signal-server-run-id")); } + #[tokio::test] + async fn signal_with_start_interceptor_can_replace_both_argument_sets() { + let (client, recorded) = mock_client( + vec![Arc::new(ReplacingSignalWithStartInterceptor)], + Arc::new(AtomicUsize::new(0)), + ); + + client + .signal_with_start_workflow( + TestWorkflow::run, + vec!["workflow".to_owned()], + TestWorkflow::test_signal, + vec!["signal".to_owned()], + WorkflowStartOptions::new("task-queue", "workflow-id").build(), + ) + .await + .unwrap(); + + let data_converter = DataConverter::default(); + let (workflow_payloads, signal_payloads) = { + let recorded = recorded.lock(); + (recorded.payloads.clone(), recorded.signal_payloads.clone()) + }; + assert_eq!( + data_converter + .from_payloads::>( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + workflow_payloads, + ) + .await + .unwrap(), + vec!["replaced-workflow".to_owned()] + ); + assert_eq!( + data_converter + .from_payloads::>( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + signal_payloads, + ) + .await + .unwrap(), + vec!["replaced-signal".to_owned()] + ); + } + + #[tokio::test] + async fn signal_with_start_payload_conversion_failure_does_not_call_service() { + let (client, recorded) = mock_client(Vec::new(), Arc::new(AtomicUsize::new(0))); + + let result = client + .signal_with_start_workflow( + TestWorkflow::run, + vec!["workflow".to_owned()], + FailingSignal, + FailingSignalInput, + WorkflowStartOptions::new("task-queue", "workflow-id").build(), + ) + .await; + + assert!(matches!( + result, + Err(WorkflowStartError::PayloadConversion(_)) + )); + assert_eq!(recorded.lock().calls, 0); + } + #[test] fn rpc_metadata_combines_with_and_overrides_connection_defaults() { let headers = Arc::new(RwLock::new(ClientHeaders { @@ -3119,13 +4160,453 @@ mod tests { } } + mod update_with_start_tests { + use super::*; + use assert_matches::assert_matches; + use parking_lot::Mutex; + use std::collections::VecDeque; + use temporalio_common::{ + UpdateDefinition, WorkflowDefinition, + data_converters::{GenericPayloadConverter, PayloadConverter}, + protos::temporal::api::{ + common::v1::{ + Header, Payload, Payloads, WorkflowExecution as ProtoWorkflowExecution, + }, + enums::v1::{ + UpdateWorkflowExecutionLifecycleStage, + WorkflowIdConflictPolicy as ProtoWorkflowIdConflictPolicy, + }, + update::v1::{ + Input as UpdateInput, Meta as UpdateMeta, Outcome, Request as UpdateRequest, + UpdateRef, WaitPolicy, outcome, + }, + }, + }; + use tonic::{Request, Response}; + + struct TestWorkflow; + + impl WorkflowDefinition for TestWorkflow { + type Input = String; + type Output = (); + + fn name(&self) -> &str { + "test-workflow" + } + } + + impl HasWorkflowDefinition for TestWorkflow { + type Run = Self; + } + + struct TestUpdate; + + impl UpdateDefinition for TestUpdate { + type Workflow = TestWorkflow; + type Input = String; + type Output = String; + + fn name(&self) -> &str { + "test-update" + } + } + + fn successful_multi_operation_response( + stage: UpdateWorkflowExecutionLifecycleStage, + ) -> ExecuteMultiOperationResponse { + let outcome = (stage == UpdateWorkflowExecutionLifecycleStage::Completed).then(|| { + let payload_converter = PayloadConverter::default(); + let result_payloads = + payload_converter + .to_payloads( + &SerializationContext::new( + &SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ), + &payload_converter, + ), + &"update-result".to_owned(), + ) + .unwrap(); + Outcome { + value: Some(outcome::Value::Success(Payloads { + payloads: result_payloads, + })), + } + }); + ExecuteMultiOperationResponse { + responses: vec![ + execute_multi_operation_response::Response { + response: Some(MultiOperationResponse::StartWorkflow( + StartWorkflowExecutionResponse { + run_id: "started-run-id".to_owned(), + first_execution_run_id: "first-run-id".to_owned(), + started: true, + ..Default::default() + }, + )), + }, + execute_multi_operation_response::Response { + response: Some(MultiOperationResponse::UpdateWorkflow( + UpdateWorkflowExecutionResponse { + update_ref: Some(UpdateRef { + workflow_execution: Some(ProtoWorkflowExecution { + workflow_id: "workflow-id".to_owned(), + run_id: "update-run-id".to_owned(), + }), + update_id: "server-update-id".to_owned(), + }), + outcome, + stage: stage as i32, + ..Default::default() + }, + )), + }, + ], + } + } + + #[derive(Clone)] + struct MockMultiOperationClient { + recorded: Arc>>, + responses: Arc>>, + call_count: Arc>, + interceptors: Vec>, + } + + impl MockMultiOperationClient { + fn new( + interceptors: Vec>, + responses: impl IntoIterator, + ) -> Self { + Self { + recorded: Arc::new(Mutex::new(None)), + responses: Arc::new(Mutex::new(responses.into_iter().collect())), + call_count: Arc::new(Mutex::new(0)), + interceptors, + } + } + } + + impl NamespacedClient for MockMultiOperationClient { + fn namespace(&self) -> String { + "test-namespace".to_owned() + } + + fn identity(&self) -> String { + "test-identity".to_owned() + } + + fn client_interceptors(&self) -> &[Arc] { + &self.interceptors + } + } + + impl WorkflowService for MockMultiOperationClient { + fn execute_multi_operation( + &mut self, + request: Request, + ) -> futures_util::future::BoxFuture< + '_, + Result, tonic::Status>, + > { + *self.recorded.lock() = Some(request.into_inner()); + *self.call_count.lock() += 1; + let response = self.responses.lock().pop_front().unwrap_or_else(|| { + successful_multi_operation_response( + UpdateWorkflowExecutionLifecycleStage::Completed, + ) + }); + Box::pin(async { Ok(Response::new(response)) }) + } + } + + fn update_with_start_options( + conflict_policy: WorkflowIdConflictPolicy, + ) -> WorkflowUpdateWithStartOptions { + WorkflowUpdateWithStartOptions::new("task-queue", "workflow-id", conflict_policy) + .build() + } + + #[tokio::test] + async fn update_with_start_builds_multi_operation_request() { + let client = MockMultiOperationClient::new(Vec::new(), []); + let recorded = client.recorded.clone(); + + let start_header = Header { + fields: HashMap::from([("start-header".to_owned(), Payload::default())]), + }; + let update_header = Header { + fields: HashMap::from([("update-header".to_owned(), Payload::default())]), + }; + let update_handle = client + .start_update_with_start_workflow( + TestWorkflow, + "workflow-input".to_owned(), + TestUpdate, + "update-input".to_owned(), + WorkflowUpdateWithStartOptions::new( + "task-queue", + "workflow-id", + WorkflowIdConflictPolicy::UseExisting, + ) + .update_id("my-update-id".to_owned()) + .start_header(start_header.clone()) + .update_header(update_header.clone()) + .build(), + ) + .await + .unwrap(); + + let payload_converter = PayloadConverter::default(); + let context_data = + SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let context = SerializationContext::new(&context_data, &payload_converter); + let workflow_payloads = payload_converter + .to_payloads(&context, &"workflow-input".to_owned()) + .unwrap(); + let update_payloads = payload_converter + .to_payloads(&context, &"update-input".to_owned()) + .unwrap(); + + let request = recorded.lock().take().unwrap(); + let request_id = assert_matches!( + &request.operations[0].operation, + Some(execute_multi_operation_request::operation::Operation::StartWorkflow(r)) => r + ) + .request_id + .clone(); + assert_eq!( + request, + ExecuteMultiOperationRequest { + namespace: "test-namespace".to_owned(), + operations: vec![ + execute_multi_operation_request::Operation { + operation: Some(MultiOperationRequest::StartWorkflow( + StartWorkflowExecutionRequest { + namespace: "test-namespace".to_owned(), + workflow_id: "workflow-id".to_owned(), + workflow_type: Some(WorkflowType { + name: "test-workflow".to_owned(), + }), + task_queue: Some(TaskQueue { + name: "task-queue".to_owned(), + ..Default::default() + }), + input: Some(Payloads { + payloads: workflow_payloads, + }), + request_id, + identity: "test-identity".to_owned(), + workflow_id_conflict_policy: + ProtoWorkflowIdConflictPolicy::UseExisting as i32, + header: Some(start_header), + priority: Some(Default::default()), + ..Default::default() + }, + )), + }, + execute_multi_operation_request::Operation { + operation: Some(MultiOperationRequest::UpdateWorkflow( + UpdateWorkflowExecutionRequest { + namespace: "test-namespace".to_owned(), + workflow_execution: Some(ProtoWorkflowExecution { + workflow_id: "workflow-id".to_owned(), + run_id: String::new(), + }), + wait_policy: Some(WaitPolicy { + lifecycle_stage: + UpdateWorkflowExecutionLifecycleStage::Accepted as i32, + }), + request: Some(UpdateRequest { + meta: Some(UpdateMeta { + update_id: "my-update-id".to_owned(), + identity: "test-identity".to_owned(), + }), + input: Some(UpdateInput { + header: Some(update_header), + name: "test-update".to_owned(), + args: Some(Payloads { + payloads: update_payloads, + }), + }), + ..Default::default() + }), + ..Default::default() + }, + )), + }, + ], + resource_id: "workflow-id".to_owned(), + } + ); + + assert_eq!(update_handle.id(), "my-update-id"); + assert_eq!(update_handle.workflow_run_id(), Some("update-run-id")); + // The outcome came back with the multi-operation response, so no poll RPC is needed + // (the mock would fail it). + let result: String = update_handle + .get_result(RpcOptions::default()) + .await + .unwrap(); + assert_eq!(result, "update-result"); + } + + #[tokio::test] + async fn update_with_start_retries_until_update_is_accepted() { + let client = MockMultiOperationClient::new( + Vec::new(), + [ + successful_multi_operation_response( + UpdateWorkflowExecutionLifecycleStage::Unspecified, + ), + successful_multi_operation_response( + UpdateWorkflowExecutionLifecycleStage::Accepted, + ), + ], + ); + let call_count = client.call_count.clone(); + + let update_handle = client + .start_update_with_start_workflow( + TestWorkflow, + "workflow-input".to_owned(), + TestUpdate, + "update-input".to_owned(), + update_with_start_options(WorkflowIdConflictPolicy::Fail), + ) + .await + .unwrap(); + + assert_eq!(*call_count.lock(), 2); + assert_eq!(update_handle.workflow_run_id(), Some("update-run-id")); + } + + #[tokio::test] + async fn update_with_start_rejects_malformed_operation_responses() { + let mut missing_response = successful_multi_operation_response( + UpdateWorkflowExecutionLifecycleStage::Accepted, + ); + missing_response.responses[0] = execute_multi_operation_response::Response::default(); + let mut extra_response = successful_multi_operation_response( + UpdateWorkflowExecutionLifecycleStage::Accepted, + ); + extra_response + .responses + .push(execute_multi_operation_response::Response::default()); + let mut wrong_order = successful_multi_operation_response( + UpdateWorkflowExecutionLifecycleStage::Accepted, + ); + wrong_order.responses.swap(0, 1); + + for response in [missing_response, extra_response, wrong_order] { + let client = MockMultiOperationClient::new(Vec::new(), [response]); + let result = client + .start_update_with_start_workflow( + TestWorkflow, + "workflow-input".to_owned(), + TestUpdate, + "update-input".to_owned(), + update_with_start_options(WorkflowIdConflictPolicy::Fail), + ) + .await; + assert!(matches!( + result, + Err(WorkflowUpdateWithStartError::Other(_)) + )); + } + } + + #[tokio::test] + async fn update_with_start_interceptor_can_mutate_args() { + struct ReplaceArgsInterceptor; + + impl ClientInterceptor for ReplaceArgsInterceptor { + fn update_with_start_workflow<'a>( + &'a self, + mut input: UpdateWithStartWorkflowInput, + next: Next< + 'a, + UpdateWithStartWorkflowInput, + BoxFuture< + 'a, + Result, + >, + >, + ) -> BoxFuture< + 'a, + Result, + > { + assert_eq!( + input.workflow_args_ref::().unwrap(), + "workflow-input" + ); + input.replace_workflow_args("replaced-workflow-input".to_owned()); + *input.update_args_mut::().unwrap() = + "replaced-update-input".to_owned(); + next.run(input) + } + } + + let client = MockMultiOperationClient::new(vec![Arc::new(ReplaceArgsInterceptor)], []); + let recorded = client.recorded.clone(); + + client + .start_update_with_start_workflow( + TestWorkflow, + "workflow-input".to_owned(), + TestUpdate, + "update-input".to_owned(), + update_with_start_options(WorkflowIdConflictPolicy::Fail), + ) + .await + .unwrap(); + + let request = recorded.lock().take().unwrap(); + let start_request = assert_matches!( + &request.operations[0].operation, + Some(execute_multi_operation_request::operation::Operation::StartWorkflow(r)) => r + ); + let workflow_input: String = client + .data_converter() + .from_payloads( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + start_request.input.clone().unwrap().payloads, + ) + .await + .unwrap(); + assert_eq!(workflow_input, "replaced-workflow-input"); + let update_request = assert_matches!( + &request.operations[1].operation, + Some(execute_multi_operation_request::operation::Operation::UpdateWorkflow(r)) => r + ); + let update_input: String = client + .data_converter() + .from_payloads( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + update_request + .request + .clone() + .unwrap() + .input + .unwrap() + .args + .unwrap() + .payloads, + ) + .await + .unwrap(); + assert_eq!(update_input, "replaced-update-input"); + } + } + mod list_workflows_tests { use super::*; use crate::test_helpers::{FailingCodec, XorCodec}; use futures_util::{FutureExt, StreamExt}; use std::sync::atomic::{AtomicUsize, Ordering}; use temporalio_common::{ - data_converters::DefaultFailureConverter, + data_converters::{DefaultFailureConverter, PayloadConverter}, protos::temporal::api::common::v1::{ Memo as ProtoMemo, Payload, WorkflowExecution as ProtoWorkflowExecution, }, @@ -3334,12 +4815,12 @@ mod tests { async fn list_workflows_exposes_typed_memo() { let data_converter = DataConverter::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), XorCodec, ); let memo_payload = data_converter .to_payload( - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), &"memo-value".to_owned(), ) .await @@ -3374,7 +4855,7 @@ mod tests { total_workflows: 1, data_converter: DataConverter::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), FailingCodec, ), memo_payload: Some(Payload::default()), diff --git a/crates/client/src/options_structs.rs b/crates/client/src/options_structs.rs index 9dd6f4605..e0b8f6cb4 100644 --- a/crates/client/src/options_structs.rs +++ b/crates/client/src/options_structs.rs @@ -1,22 +1,33 @@ use crate::{ - ClientInterceptor, ClientPlugin, ErasedClientPlugin, HttpConnectProxyOptions, RetryOptions, - RpcOptions, VERSION, callback_based, + ClientInterceptor, HttpConnectProxyOptions, RetryOptions, RpcOptions, VERSION, callback_based, }; +#[cfg(feature = "experimental")] +use crate::{ClientPlugin, ErasedClientPlugin}; use http::Uri; use std::{collections::HashMap, sync::Arc, time::Duration}; use temporalio_common::{ - RetryPolicy, - data_converters::DataConverter, + ActivityCloseTimeouts, MemoValues, RetryPolicy, + data_converters::{ + DataConverter, GenericPayloadConverter, PayloadConversionError, PayloadConverter, + SerializationContext, SerializationContextData, WorkflowSerializationContext, + }, + payload_visitor::encode_payloads, protos::temporal::api::{ common::{ self, - v1::{Header, Payloads}, + v1::{Header, Memo as ProtoMemo, Payloads}, }, enums::v1::{ - ArchivalState, HistoryEventFilterType, QueryRejectCondition, WorkflowIdConflictPolicy, - WorkflowIdReusePolicy, + ActivityIdConflictPolicy as ProtoActivityIdConflictPolicy, + ActivityIdReusePolicy as ProtoActivityIdReusePolicy, + ArchivalState as ProtoArchivalState, + HistoryEventFilterType as ProtoHistoryEventFilterType, + QueryRejectCondition as ProtoQueryRejectCondition, + WorkflowIdConflictPolicy as ProtoWorkflowIdConflictPolicy, + WorkflowIdReusePolicy as ProtoWorkflowIdReusePolicy, }, replication::v1::ClusterReplicationConfig, + sdk::v1::UserMetadata, workflowservice::v1::RegisterNamespaceRequest, }, search_attributes::SearchAttributes, @@ -27,6 +38,9 @@ use tokio_rustls::rustls::client::ResolvesClientCert; use tokio_rustls::rustls::client::danger::ServerCertVerifier; use url::Url; +pub(crate) const DEFAULT_PAYLOADS_WARN_SIZE: u64 = 512 * 1024; +pub(crate) const DEFAULT_MEMO_WARN_SIZE: u64 = 2 * 1024; + /// Options for [crate::Connection::connect]. #[derive(bon::Builder, Clone, Debug)] #[non_exhaustive] @@ -99,6 +113,14 @@ pub struct ConnectionOptions { /// Payload size limit options for this connection. Defaults to the standard warning thresholds; /// disable an individual warning by setting its threshold to `0`. /// NOTE: Experimental + #[cfg(feature = "experimental")] + #[cfg_attr( + docsrs, + builder(setters( + some_fn(name = payload_limits_impl, vis = "pub(crate)"), + option_fn(name = maybe_payload_limits_impl, vis = "pub(crate)") + )) + )] #[builder(default)] pub payload_limits: PayloadLimitsOptions, @@ -121,6 +143,35 @@ pub struct ConnectionOptions { pub(crate) client_version: String, } +// Bon does not propagate `doc(cfg)` to generated setters, so these docs-only methods forward to +// renamed generated implementations. +#[cfg(all(feature = "experimental", docsrs))] +impl ConnectionOptionsBuilder { + /// Set the payload size limit options for this connection. + #[doc(cfg(feature = "experimental"))] + pub fn payload_limits( + self, + value: PayloadLimitsOptions, + ) -> ConnectionOptionsBuilder> + where + S::PayloadLimits: connection_options_builder::IsUnset, + { + self.payload_limits_impl(value) + } + + /// Set the payload size limit options for this connection from an optional value. + #[doc(cfg(feature = "experimental"))] + pub fn maybe_payload_limits( + self, + value: Option, + ) -> ConnectionOptionsBuilder> + where + S::PayloadLimits: connection_options_builder::IsUnset, + { + self.maybe_payload_limits_impl(value) + } +} + // Setters/getters for fields that should only be touched by SDK implementers. #[cfg(feature = "core-based-sdk")] impl ConnectionOptions { @@ -153,10 +204,12 @@ pub struct ClientOptions { #[builder(field)] #[debug(skip)] + #[cfg(feature = "experimental")] plugins: Vec, #[builder(field)] #[debug(skip)] + #[cfg(feature = "experimental")] client_plugins_applied: bool, /// The data converter used for serializing/deserializing payloads. @@ -168,6 +221,7 @@ pub struct ClientOptions { pub client_interceptors: Vec>, } +#[cfg(feature = "experimental")] impl ClientOptionsBuilder { /// Register a type-erased client plugin. /// @@ -204,14 +258,17 @@ impl ClientOptions { /// This is intended for SDK integrations that propagate worker plugin registrations. /// /// **Experimental:** This API may change or be removed. + #[cfg(feature = "experimental")] pub fn plugins(&self) -> &[ErasedClientPlugin] { &self.plugins } + #[cfg(feature = "experimental")] pub(crate) fn client_plugins_applied(&self) -> bool { self.client_plugins_applied } + #[cfg(feature = "experimental")] pub(crate) fn mark_client_plugins_applied(&mut self) { self.client_plugins_applied = true; } @@ -346,19 +403,21 @@ impl Default for DnsLoadBalancingOptions { /// Payload size limit options for a connection. /// NOTE: Experimental +#[cfg(feature = "experimental")] #[derive(Clone, Debug, PartialEq, bon::Builder)] #[non_exhaustive] pub struct PayloadLimitsOptions { /// Warning threshold (bytes) for the size of an outbound payload-bearing field; over-threshold /// fields are logged but still sent to server. Defaults to 512 KiB. Set to `0` to disable. - #[builder(default = 512 * 1024)] + #[builder(default = DEFAULT_PAYLOADS_WARN_SIZE)] pub payloads_warn_size: u64, /// Warning threshold (bytes) for outbound memo sizes; over-threshold memos are logged but still /// sent to server. Defaults to 2 KiB. Set to `0` to disable. - #[builder(default = 2 * 1024)] + #[builder(default = DEFAULT_MEMO_WARN_SIZE)] pub memo_warn_size: u64, } +#[cfg(feature = "experimental")] impl Default for PayloadLimitsOptions { fn default() -> Self { Self::builder().build() @@ -372,6 +431,59 @@ impl std::fmt::Debug for ClientTlsOptions { } } +/// Controls whether a closed workflow ID may be reused. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)] +#[non_exhaustive] +pub enum WorkflowIdReusePolicy { + /// Use the server's default policy. + #[default] + Unspecified, + /// Allow starting a workflow using the same workflow ID. + AllowDuplicate, + /// Allow reuse only when the previous execution did not complete successfully. + AllowDuplicateFailedOnly, + /// Reject reuse of the workflow ID. + RejectDuplicate, +} + +impl From for ProtoWorkflowIdReusePolicy { + fn from(value: WorkflowIdReusePolicy) -> Self { + match value { + WorkflowIdReusePolicy::Unspecified => Self::Unspecified, + WorkflowIdReusePolicy::AllowDuplicate => Self::AllowDuplicate, + WorkflowIdReusePolicy::AllowDuplicateFailedOnly => Self::AllowDuplicateFailedOnly, + WorkflowIdReusePolicy::RejectDuplicate => Self::RejectDuplicate, + } + } +} + +/// Controls how starting a workflow resolves a conflict with a running workflow using the same +/// workflow ID. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)] +#[non_exhaustive] +pub enum WorkflowIdConflictPolicy { + /// Use the server's default policy. + #[default] + Unspecified, + /// Do not start a new workflow and return an already-started error. + Fail, + /// Do not start a new workflow and return a handle for the running workflow. + UseExisting, + /// Terminate the running workflow before starting a new one. + TerminateExisting, +} + +impl From for ProtoWorkflowIdConflictPolicy { + fn from(value: WorkflowIdConflictPolicy) -> Self { + match value { + WorkflowIdConflictPolicy::Unspecified => Self::Unspecified, + WorkflowIdConflictPolicy::Fail => Self::Fail, + WorkflowIdConflictPolicy::UseExisting => Self::UseExisting, + WorkflowIdConflictPolicy::TerminateExisting => Self::TerminateExisting, + } + } +} + /// Options for starting a workflow execution. #[derive(Debug, Clone, bon::Builder)] #[builder(start_fn = new, on(String, into))] @@ -410,8 +522,7 @@ pub struct WorkflowStartOptions { /// Additional search attributes for the workflow. pub search_attributes: Option, - /// Optionally enable Eager Workflow Start, a latency optimization using local workers - /// NOTE: Experimental + /// Optionally enable Eager Workflow Start, a latency optimization using local workers. #[builder(default)] pub enable_eager_workflow_start: bool, @@ -419,10 +530,6 @@ pub struct WorkflowStartOptions { #[builder(into)] pub retry_policy: Option, - /// If set, send a signal to the workflow atomically with start. - /// The workflow will receive this signal before its first task. - pub start_signal: Option, - /// Links to associate with the workflow. Ex: References to a nexus operation. #[builder(default)] pub links: Vec, @@ -439,6 +546,9 @@ pub struct WorkflowStartOptions { /// Headers to include with the start request. pub header: Option
, + /// Non-indexed values attached to the workflow, serialized with the client's data converter. + pub memo: Option, + /// Single-line static summary for the workflow, shown in the Temporal UI. pub static_summary: Option, @@ -450,19 +560,184 @@ pub struct WorkflowStartOptions { pub rpc_options: RpcOptions, } -/// A signal to send atomically when starting a workflow. -/// Use with `WorkflowStartOptions::start_signal` to achieve signal-with-start behavior. +impl WorkflowStartOptions { + pub(crate) async fn encoded_memo( + &self, + data_converter: &DataConverter, + ) -> Result, PayloadConversionError> { + let Some(memo) = &self.memo else { + return Ok(None); + }; + + let payload_converter = data_converter.payload_converter(); + let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let context = SerializationContext::new(&context_data, payload_converter); + let mut memo = ProtoMemo { + fields: memo + .iter() + .map(|(key, value)| { + payload_converter + .to_payload(&context, value) + .map(|payload| (key.to_owned(), payload)) + }) + .collect::>()?, + }; + encode_payloads( + &mut memo, + data_converter.codec(), + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + ) + .await?; + Ok(Some(memo)) + } + + pub(crate) fn user_metadata(&self) -> Option { + (self.static_summary.is_some() || self.static_details.is_some()).then(|| { + let payload_converter = PayloadConverter::default(); + let context_data = + SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let context = SerializationContext::new(&context_data, &payload_converter); + UserMetadata { + summary: self.static_summary.as_ref().map(|summary| { + payload_converter + .to_payload(&context, summary) + .expect("String-to-JSON payload serialization is infallible") + }), + details: self.static_details.as_ref().map(|details| { + payload_converter + .to_payload(&context, details) + .expect("String-to-JSON payload serialization is infallible") + }), + } + }) + } +} + +/// Options for starting a workflow and sending it an update in one atomic operation. +/// +/// See [crate::Client::start_update_with_start_workflow] and +/// [crate::Client::execute_update_with_start_workflow]. #[derive(Debug, Clone, bon::Builder)] #[builder(start_fn = new, on(String, into))] #[non_exhaustive] -pub struct WorkflowStartSignal { - /// Name of the signal to send. +pub struct WorkflowUpdateWithStartOptions { + /// The task queue to run the workflow on. #[builder(start_fn)] - pub signal_name: String, - /// Payload for the signal. - pub input: Option, - /// Headers for the signal. - pub header: Option
, + pub task_queue: String, + + /// The workflow ID. + #[builder(start_fn)] + pub workflow_id: String, + + /// How to resolve a conflict with an already-running workflow. This is required so callers + /// explicitly choose whether an update may attach to an existing workflow. + #[builder(start_fn)] + pub id_conflict_policy: WorkflowIdConflictPolicy, + + /// The policy for reusing the workflow ID after a workflow closes. + #[builder(default)] + pub id_reuse_policy: WorkflowIdReusePolicy, + + /// The workflow execution timeout. + pub execution_timeout: Option, + + /// The workflow run timeout. + pub run_timeout: Option, + + /// The workflow task timeout. + pub task_timeout: Option, + + /// Search attributes for the workflow. + pub search_attributes: Option, + + /// The workflow retry policy. + #[builder(into)] + pub retry_policy: Option, + + /// Links to associate with the workflow. + #[builder(default)] + pub links: Vec, + + /// Callbacks invoked when the workflow completes. + #[builder(default)] + pub completion_callbacks: Vec, + + /// Priority for the workflow. Defaults to all-inherited (empty). + #[builder(default)] + pub priority: Priority, + + /// Headers to include with the start operation. + pub start_header: Option
, + + /// Headers to include with the update operation. + pub update_header: Option
, + + /// Non-indexed values attached to the workflow, serialized with the client's data converter. + pub memo: Option, + + /// Single-line static summary for the workflow, shown in the Temporal UI. + pub static_summary: Option, + + /// Multi-line static details for the workflow, shown in the Temporal UI. + pub static_details: Option, + + /// Update ID for idempotency. If not provided, a UUID will be generated. + pub update_id: Option, + + /// Controls for the multi-operation RPC and, when executing the update, subsequent polling. + #[builder(default)] + pub rpc_options: RpcOptions, +} + +impl WorkflowUpdateWithStartOptions { + pub(crate) fn into_parts(self) -> (WorkflowStartOptions, Option, Option
) { + let Self { + task_queue, + workflow_id, + id_conflict_policy, + id_reuse_policy, + execution_timeout, + run_timeout, + task_timeout, + search_attributes, + retry_policy, + links, + completion_callbacks, + priority, + start_header, + update_header, + memo, + static_summary, + static_details, + update_id, + rpc_options: _, + } = self; + ( + WorkflowStartOptions { + task_queue, + workflow_id, + id_reuse_policy, + id_conflict_policy, + execution_timeout, + run_timeout, + task_timeout, + cron_schedule: None, + search_attributes, + enable_eager_workflow_start: false, + retry_policy, + links, + completion_callbacks, + priority, + header: start_header, + memo, + static_summary, + static_details, + rpc_options: RpcOptions::default(), + }, + update_id, + update_header, + ) + } } pub use temporalio_common::Priority; @@ -514,6 +789,32 @@ pub struct WorkflowSignalOptions { pub rpc_options: RpcOptions, } +/// Controls when a workflow query should be rejected based on workflow state. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)] +#[non_exhaustive] +pub enum QueryRejectCondition { + /// Use the server's default condition. + #[default] + Unspecified, + /// Do not reject the query based on workflow state. + None, + /// Reject the query if the workflow is not open. + NotOpen, + /// Reject the query if the workflow did not complete successfully. + NotCompletedCleanly, +} + +impl From for ProtoQueryRejectCondition { + fn from(value: QueryRejectCondition) -> Self { + match value { + QueryRejectCondition::Unspecified => Self::Unspecified, + QueryRejectCondition::None => Self::None, + QueryRejectCondition::NotOpen => Self::NotOpen, + QueryRejectCondition::NotCompletedCleanly => Self::NotCompletedCleanly, + } + } +} + /// Options for querying a workflow. #[derive(Debug, Clone, Default, bon::Builder)] #[non_exhaustive] @@ -570,6 +871,29 @@ pub struct WorkflowDescribeOptions { /// Default workflow execution retention for a Namespace is 3 days const DEFAULT_WORKFLOW_EXECUTION_RETENTION_PERIOD: Duration = Duration::from_secs(60 * 60 * 24 * 3); +/// Controls whether archival is enabled for a namespace. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)] +#[non_exhaustive] +pub enum ArchivalState { + /// Use the server's default archival state. + #[default] + Unspecified, + /// Disable archival. + Disabled, + /// Enable archival. + Enabled, +} + +impl From for ProtoArchivalState { + fn from(value: ArchivalState) -> Self { + match value { + ArchivalState::Unspecified => Self::Unspecified, + ArchivalState::Disabled => Self::Disabled, + ArchivalState::Enabled => Self::Enabled, + } + } +} + /// Helper struct for `register_namespace`. #[derive(Clone, Debug, bon::Builder)] #[builder(on(String, into))] @@ -629,14 +953,38 @@ impl From for RegisterNamespaceRequest { data: val.data, security_token: val.security_token, is_global_namespace: val.is_global_namespace, - history_archival_state: val.history_archival_state as i32, + history_archival_state: ProtoArchivalState::from(val.history_archival_state) as i32, history_archival_uri: val.history_archival_uri, - visibility_archival_state: val.visibility_archival_state as i32, + visibility_archival_state: ProtoArchivalState::from(val.visibility_archival_state) + as i32, visibility_archival_uri: val.visibility_archival_uri, } } } +/// Selects which workflow history events are returned when fetching history. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)] +#[non_exhaustive] +pub enum HistoryEventFilterType { + /// Use the server's default filter. + #[default] + Unspecified, + /// Return all history events. + AllEvent, + /// Return only the workflow's close event. + CloseEvent, +} + +impl From for ProtoHistoryEventFilterType { + fn from(value: HistoryEventFilterType) -> Self { + match value { + HistoryEventFilterType::Unspecified => Self::Unspecified, + HistoryEventFilterType::AllEvent => Self::AllEvent, + HistoryEventFilterType::CloseEvent => Self::CloseEvent, + } + } +} + /// Options for fetching workflow history. #[derive(Debug, Clone, Default, bon::Builder)] #[non_exhaustive] @@ -668,6 +1016,17 @@ pub struct WorkflowStartUpdateOptions { pub rpc_options: RpcOptions, } +impl From for WorkflowStartUpdateOptions { + /// Execute-update is start-update followed by waiting for the update result. + fn from(options: WorkflowExecuteUpdateOptions) -> Self { + Self::builder() + .maybe_update_id(options.update_id) + .maybe_header(options.header) + .rpc_options(options.rpc_options) + .build() + } +} + /// Options for listing workflows. #[derive(Debug, Clone, Default, bon::Builder)] #[non_exhaustive] @@ -688,3 +1047,179 @@ pub struct WorkflowCountOptions { #[builder(default)] pub rpc_options: RpcOptions, } + +/// Options for starting a standalone activity. +#[derive(Clone, Debug, bon::Builder)] +#[builder(start_fn = new, on(String, into))] +#[non_exhaustive] +pub struct ActivityStartOptions { + /// Task queue to run this activity on. + #[builder(start_fn)] + pub task_queue: String, + /// Activity ID of the started activity. It's recommended to use a meaningful business ID. + #[builder(start_fn)] + pub id: String, + /// Timeouts for activity completion. + /// + /// See [`ActivityCloseTimeouts`] for the meaning of each timeout variant. + #[builder(start_fn)] + pub close_timeouts: ActivityCloseTimeouts, + /// If set, specifies maximum time the activity can wait in the task queue before being picked + /// up by a worker. This timeout is non-retryable. + pub schedule_to_start_timeout: Option, + /// If set, specifies maximum time between successful heartbeats. + pub heartbeat_timeout: Option, + /// Controls how Activity is retried. If not set, the server will assign default retry policy. + #[builder(into)] + pub retry_policy: Option, + /// Priority to use when starting this activity. + #[builder(default)] + pub priority: Priority, + /// Specifies behavior if there's a *closed* activity with the same ID. + #[builder(default)] + pub id_reuse_policy: ActivityIdReusePolicy, + /// Specifies behavior if there's a *running* activity with the same ID. Note that there can + /// only be one running activity for each Activity ID. + #[builder(default)] + pub id_conflict_policy: ActivityIdConflictPolicy, + /// Search attributes for the activity. + pub search_attributes: Option, + /// Headers to include with the start request. + pub header: Option
, + /// Single-line static summary for the activity, shown in the Temporal UI. + pub summary: Option, + /// Multi-line static details for the activity, shown in the Temporal UI. + pub static_details: Option, + /// Time to wait before dispatching the first activity task. + /// This delay is not applied to retry attempts. + pub start_delay: Option, +} + +impl ActivityStartOptions { + /// Returns a builder with `close_timeouts` set to [`ActivityCloseTimeouts::StartToClose`]. + pub fn with_start_to_close_timeout( + task_queue: impl Into, + activity_id: impl Into, + start_to_close_timeout: Duration, + ) -> ActivityStartOptionsBuilder { + Self::new( + task_queue, + activity_id, + ActivityCloseTimeouts::StartToClose(start_to_close_timeout), + ) + } + + /// Returns a builder with `close_timeouts` set to [`ActivityCloseTimeouts::ScheduleToClose`]. + pub fn with_schedule_to_close_timeout( + task_queue: impl Into, + activity_id: impl Into, + schedule_to_close_timeout: Duration, + ) -> ActivityStartOptionsBuilder { + Self::new( + task_queue, + activity_id, + ActivityCloseTimeouts::ScheduleToClose(schedule_to_close_timeout), + ) + } +} + +/// Specifies behavior when starting a standalone activity if there's a *closed* activity with +/// the same ID. See [`ActivityStartOptions::id_reuse_policy`]. +#[non_exhaustive] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)] +pub enum ActivityIdReusePolicy { + #[default] + /// Always allow starting an activity using the same activity ID. This is the default. + AllowDuplicate, + /// Allow starting an activity using the same ID only when the last execution did not complete + /// successfully. + AllowDuplicateFailedOnly, + /// Do not permit re-use of the ID for this activity. + RejectDuplicate, +} + +impl From for ProtoActivityIdReusePolicy { + fn from(value: ActivityIdReusePolicy) -> Self { + match value { + ActivityIdReusePolicy::AllowDuplicate => Self::AllowDuplicate, + ActivityIdReusePolicy::AllowDuplicateFailedOnly => Self::AllowDuplicateFailedOnly, + ActivityIdReusePolicy::RejectDuplicate => Self::RejectDuplicate, + } + } +} + +/// Specifies behavior when starting a standalone activity if there's a *running* activity with +/// the same ID. See [`ActivityStartOptions::id_conflict_policy`]. +#[non_exhaustive] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)] +pub enum ActivityIdConflictPolicy { + #[default] + /// Don't start a new activity; instead return + /// [`StartActivityError::AlreadyStarted`](crate::errors::StartActivityError::AlreadyStarted). + Fail, + /// Don't start a new activity; instead return a handle for the running activity. + UseExisting, +} + +impl From for ProtoActivityIdConflictPolicy { + fn from(value: ActivityIdConflictPolicy) -> Self { + match value { + ActivityIdConflictPolicy::Fail => Self::Fail, + ActivityIdConflictPolicy::UseExisting => Self::UseExisting, + } + } +} + +/// Options for listing activities. +#[derive(Debug, Clone, Default, bon::Builder)] +#[non_exhaustive] +pub struct ActivityListOptions {} + +/// Options for counting activities. +#[derive(Debug, Clone, Default, bon::Builder)] +#[non_exhaustive] +pub struct ActivityCountOptions {} + +/// Controls which optional fields will be requested in +/// [`ActivityHandle::describe`](crate::ActivityHandle::describe) operation. The fields will be +/// present in returned [`ActivityExecutionDescription`](crate::ActivityExecutionDescription), +/// subject to data availability and server support. +/// +/// Note that these fields contain payloads that can be arbitrarily large. It's recommended not to +/// include them unless they're needed. +#[derive(Debug, Clone, Default, bon::Builder)] +#[non_exhaustive] +pub struct ActivityDescribeOptions { + /// If set and the activity received input, the input will be included. + #[builder(default)] + pub include_input: bool, + /// If set and the activity is closed, the activity outcome will be included. + #[builder(default)] + pub include_outcome: bool, + /// If set and the activity sent heartbeat details, the heartbeat details will be included. + #[builder(default)] + pub include_heartbeat_details: bool, + /// If set and the activity has a failed attempt, the last failure will be included. + #[builder(default)] + pub include_last_failure: bool, +} + +/// Options for [`ActivityHandle::cancel`](crate::ActivityHandle::cancel). +#[derive(Debug, Clone, Default, bon::Builder)] +#[builder(on(String, into))] +#[non_exhaustive] +pub struct ActivityCancelOptions { + /// Reason for cancellation. Can be empty. + #[builder(default)] + pub reason: String, +} + +/// Options for [`ActivityHandle::terminate`](crate::ActivityHandle::terminate). +#[derive(Debug, Clone, Default, bon::Builder)] +#[builder(on(String, into))] +#[non_exhaustive] +pub struct ActivityTerminateOptions { + /// Reason for termination. Can be empty. + #[builder(default)] + pub reason: String, +} diff --git a/crates/client/src/plugins.rs b/crates/client/src/plugins.rs index 09716f079..f5e7d2fcc 100644 --- a/crates/client/src/plugins.rs +++ b/crates/client/src/plugins.rs @@ -174,7 +174,7 @@ pub(crate) fn apply_client_plugins(options: &mut ClientOptions) -> Result<(), Pl Ok(()) } -#[cfg(test)] +#[cfg(all(test, feature = "experimental"))] mod tests { use super::*; use std::sync::atomic::{AtomicUsize, Ordering}; diff --git a/crates/client/src/proxy.rs b/crates/client/src/proxy.rs index df1fd3eca..09ae6ae13 100644 --- a/crates/client/src/proxy.rs +++ b/crates/client/src/proxy.rs @@ -108,9 +108,7 @@ impl Service for OverrideAddrConnector { } } -/// Visible only for tests -#[doc(hidden)] -pub enum ProxyStream { +enum ProxyStream { Tcp(TcpStream), #[cfg(unix)] Unix(UnixStream), @@ -234,11 +232,214 @@ fn ensure_connect_authority_port(uri: tonic::transport::Uri) -> tonic::transport #[cfg(test)] mod tests { - use super::*; + use super::{HttpConnectProxyOptions, ProxyStream}; + use crate::{ + Client, ClientOptions, Connection as TemporalConnection, ConnectionOptions, RetryOptions, + grpc::WorkflowService, + }; + use base64::prelude::*; + use futures_util::{FutureExt, future::BoxFuture}; + use http::{Request, Response}; + use http_body_util::Empty; + use hyper::{ + body::{Bytes, Incoming}, + server::conn::http1, + service::service_fn, + }; + use hyper_util::rt::TokioIo; + use std::{ + convert::Infallible, + io, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + task::{Context, Poll}, + }; + use temporalio_common::protos::temporal::api::workflowservice::v1::ListNamespacesRequest; + #[cfg(unix)] + use tokio::net::UnixListener; use tokio::{ io::{AsyncBufReadExt, AsyncWriteExt, BufReader}, - net::TcpListener, + net::{TcpListener, TcpStream}, + sync::oneshot, }; + use tokio_stream::wrappers::TcpListenerStream; + use tonic::{IntoRequest, body::Body, server::NamedService, transport::Server}; + use tower::Service; + use tracing::warn; + use url::Url; + + #[derive(Clone)] + struct FakeWorkflowService(F); + + impl Service> for FakeWorkflowService + where + F: FnMut(Request) -> BoxFuture<'static, Response>, + { + type Response = Response; + type Error = Infallible; + type Future = BoxFuture<'static, Result>; + + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn call(&mut self, request: Request) -> Self::Future { + let response = (self.0)(request); + async move { Ok(response.await) }.boxed() + } + } + + impl NamedService for FakeWorkflowService { + const NAME: &'static str = "temporal.api.workflowservice.v1.WorkflowService"; + } + + struct FakeServer { + addr: std::net::SocketAddr, + shutdown_tx: oneshot::Sender<()>, + } + + async fn fake_server(response_maker: F) -> FakeServer + where + F: FnMut(Request) -> BoxFuture<'static, Response> + + Clone + + Send + + Sync + + 'static, + { + let (shutdown_tx, shutdown_rx) = oneshot::channel(); + let listener = TcpListener::bind("[::]:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { + Server::builder() + .add_service(FakeWorkflowService(response_maker)) + .serve_with_incoming_shutdown(TcpListenerStream::new(listener), async move { + let _ = shutdown_rx.await; + }) + .await + .unwrap(); + }); + FakeServer { addr, shutdown_tx } + } + + struct HttpProxy { + proxy_hits: Arc, + shutdown_tx: oneshot::Sender<()>, + } + + impl HttpProxy { + fn spawn_tcp(listener: TcpListener) -> Self { + Self::spawn(ProxyListener::Tcp(listener)) + } + + #[cfg(unix)] + fn spawn_unix(listener: UnixListener) -> Self { + Self::spawn(ProxyListener::Unix(listener)) + } + + fn spawn(listener: ProxyListener) -> Self { + let (shutdown_tx, mut shutdown_rx) = oneshot::channel(); + let proxy_hits = Arc::new(AtomicUsize::new(0)); + let proxy_hits_for_task = proxy_hits.clone(); + tokio::spawn(async move { + loop { + let proxy_hits = proxy_hits_for_task.clone(); + tokio::select! { + _ = &mut shutdown_rx => break, + stream = listener.accept() => { + let stream = match stream { + Ok(stream) => stream, + Err(error) => { + warn!(%error, "Proxy accept failed"); + continue; + } + }; + tokio::spawn(async move { + if let Err(error) = http1::Builder::new() + .serve_connection( + TokioIo::new(stream), + service_fn(move |request| { + handle_connect(request, proxy_hits.clone()) + }), + ) + .with_upgrades() + .await + { + warn!(%error, "Proxy connection failed"); + } + }); + } + } + } + }); + Self { + proxy_hits, + shutdown_tx, + } + } + + fn hit_count(&self) -> usize { + self.proxy_hits.load(Ordering::SeqCst) + } + + fn shutdown(self) { + let _ = self.shutdown_tx.send(()); + } + } + + async fn handle_connect( + request: Request, + counter: Arc, + ) -> Result>, hyper::Error> { + if request.method() != hyper::Method::CONNECT { + return Ok(Response::builder() + .status(hyper::StatusCode::METHOD_NOT_ALLOWED) + .body(Empty::new()) + .unwrap()); + } + + counter.fetch_add(1, Ordering::SeqCst); + tokio::spawn(async move { + if let Some(addr) = request + .uri() + .authority() + .map(|authority| authority.as_str()) + && let Ok(mut server_stream) = TcpStream::connect(addr).await + && let Ok(upgraded) = hyper::upgrade::on(request).await + { + let mut upgraded = TokioIo::new(upgraded); + let _ = tokio::io::copy_bidirectional(&mut upgraded, &mut server_stream).await; + } + }); + + Ok(Response::builder() + .status(hyper::StatusCode::OK) + .body(Empty::new()) + .unwrap()) + } + + enum ProxyListener { + Tcp(TcpListener), + #[cfg(unix)] + Unix(UnixListener), + } + + impl ProxyListener { + async fn accept(&self) -> io::Result { + match self { + ProxyListener::Tcp(listener) => listener + .accept() + .await + .map(|(stream, _)| ProxyStream::Tcp(stream)), + #[cfg(unix)] + ProxyListener::Unix(listener) => listener + .accept() + .await + .map(|(stream, _)| ProxyStream::Unix(stream)), + } + } + } struct CapturedConnect { request_line: String, @@ -311,4 +512,77 @@ mod tests { format!("proxy-authorization: Basic {creds}") ); } + + #[tokio::test] + async fn connection_uses_http_connect_proxy() { + let call_count = Arc::new(AtomicUsize::new(0)); + let call_count_for_server = call_count.clone(); + let server = fake_server(move |_| { + call_count_for_server.fetch_add(1, Ordering::SeqCst); + async { Response::new(Body::empty()) }.boxed() + }) + .await; + + let tcp_proxy_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let tcp_proxy_addr = tcp_proxy_listener.local_addr().unwrap(); + let tcp_proxy = HttpProxy::spawn_tcp(tcp_proxy_listener); + + let mut options = ConnectionOptions::new( + Url::parse(&format!("http://[::1]:{}", server.addr.port())).unwrap(), + ) + .retry_options(RetryOptions::no_retries()) + .skip_get_system_info(true) + .build(); + + let connection = TemporalConnection::connect(options.clone()).await.unwrap(); + let client_options = ClientOptions::new("my-namespace").build(); + let client = Client::new(connection, client_options).unwrap(); + let _ = WorkflowService::list_namespaces( + &mut client.clone(), + ListNamespacesRequest::default().into_request(), + ) + .await; + assert_eq!(call_count.load(Ordering::SeqCst), 1); + assert_eq!(tcp_proxy.hit_count(), 0); + + options.http_connect_proxy = + Some(HttpConnectProxyOptions::new(tcp_proxy_addr.to_string()).build()); + options.dns_load_balancing = None; + let connection = TemporalConnection::connect(options.clone()).await.unwrap(); + let client_options = ClientOptions::new("my-namespace").build(); + let proxied_client = Client::new(connection, client_options).unwrap(); + let _ = WorkflowService::list_namespaces( + &mut proxied_client.clone(), + ListNamespacesRequest::default().into_request(), + ) + .await; + assert_eq!(call_count.load(Ordering::SeqCst), 2); + assert_eq!(tcp_proxy.hit_count(), 1); + + #[cfg(unix)] + { + let socket_dir = tempfile::tempdir().unwrap(); + let socket_path = socket_dir.path().join("http-proxy.sock"); + let unix_proxy = HttpProxy::spawn_unix(UnixListener::bind(&socket_path).unwrap()); + + options.http_connect_proxy = Some( + HttpConnectProxyOptions::new(format!("unix:{}", socket_path.display())).build(), + ); + let connection = TemporalConnection::connect(options).await.unwrap(); + let client_options = ClientOptions::new("my-namespace").build(); + let proxied_client = Client::new(connection, client_options).unwrap(); + let _ = WorkflowService::list_namespaces( + &mut proxied_client.clone(), + ListNamespacesRequest::default().into_request(), + ) + .await; + assert_eq!(call_count.load(Ordering::SeqCst), 3); + assert_eq!(unix_proxy.hit_count(), 1); + + unix_proxy.shutdown(); + } + + let _ = server.shutdown_tx.send(()); + tcp_proxy.shutdown(); + } } diff --git a/crates/client/src/request_extensions.rs b/crates/client/src/request_extensions.rs index 6f6f614f2..3e4904648 100644 --- a/crates/client/src/request_extensions.rs +++ b/crates/client/src/request_extensions.rs @@ -7,7 +7,7 @@ use crate::RetryOptions; use std::time::Duration; /// A request extension that, when set, should make the retry behavior consider this call to be a -/// [CallType::TaskLongPoll](crate::CallType::TaskLongPoll) +/// worker task long poll. #[derive(Copy, Clone, Debug)] pub struct IsWorkerTaskLongPoll; diff --git a/crates/client/src/retry.rs b/crates/client/src/retry.rs index f655bf689..9ac2cdd81 100644 --- a/crates/client/src/retry.rs +++ b/crates/client/src/retry.rs @@ -14,8 +14,7 @@ use std::{ use tonic::Code; /// List of gRPC error codes that client will retry. -#[doc(hidden)] -pub const RETRYABLE_ERROR_CODES: [Code; 7] = [ +const RETRYABLE_ERROR_CODES: [Code; 7] = [ Code::DataLoss, Code::Internal, Code::Unknown, @@ -252,9 +251,8 @@ pub(crate) struct CallInfo { retry_short_circuit: Option, } -#[doc(hidden)] #[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)] -pub enum CallType { +pub(crate) enum CallType { Normal, // A long poll but won't always retry timeouts/cancels. EX: Get workflow history UserLongPoll, @@ -373,12 +371,25 @@ fn is_transport_cancelled(status: &tonic::Status) -> bool { #[cfg(test)] mod tests { use super::*; + use crate::{ + Client, ClientOptions, Connection, ConnectionOptions, + callback_based::{CallbackBasedGrpcService, GrpcSuccessResponse}, + }; use assert_matches::assert_matches; - use std::time::Instant; + use prost::Message; + use std::{ + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + time::Instant, + }; use temporalio_common::protos::temporal::api::workflowservice::v1::{ - PollActivityTaskQueueRequest, PollNexusTaskQueueRequest, PollWorkflowTaskQueueRequest, + CountWorkflowExecutionsResponse, PollActivityTaskQueueRequest, PollNexusTaskQueueRequest, + PollWorkflowTaskQueueRequest, }; use tonic::{IntoRequest, Status}; + use url::Url; /// Predefined retry configs with low durations to make unit tests faster const TEST_RETRY_CONFIG: RetryOptions = RetryOptions { @@ -394,6 +405,49 @@ mod tests { const POLL_ACTIVITY_METH_NAME: &str = "poll_activity_task_queue"; const POLL_NEXUS_METH_NAME: &str = "poll_nexus_task_queue"; + #[tokio::test] + async fn retryable_errors() { + // Resource exhausted has a separate retry policy and is covered below. + for code in RETRYABLE_ERROR_CODES + .iter() + .copied() + .filter(|code| code != &Code::ResourceExhausted) + { + let attempts = Arc::new(AtomicUsize::new(0)); + let callback_attempts = attempts.clone(); + let service_override = CallbackBasedGrpcService { + callback: Arc::new(move |request| { + assert_eq!(request.rpc, "CountWorkflowExecutions"); + let callback_attempts = callback_attempts.clone(); + Box::pin(async move { + if callback_attempts.fetch_add(1, Ordering::Relaxed) < 3 { + Err(Status::new(code, "retryable")) + } else { + Ok(GrpcSuccessResponse { + headers: Default::default(), + proto: CountWorkflowExecutionsResponse::default().encode_to_vec(), + }) + } + }) + }), + }; + let connection_options = + ConnectionOptions::new(Url::parse("http://localhost:7233").unwrap()) + .retry_options(TEST_RETRY_CONFIG) + .skip_get_system_info(true) + .service_override(service_override) + .dns_load_balancing(None) + .build(); + let connection = Connection::connect(connection_options).await.unwrap(); + let client = Client::new(connection, ClientOptions::new("ns").build()).unwrap(); + + let result = client.count_workflows("whatever", Default::default()).await; + + assert!(result.is_ok(), "{result:?}"); + assert_eq!(attempts.load(Ordering::Relaxed), 4); + } + } + #[tokio::test] async fn long_poll_non_retryable_errors() { for code in [ diff --git a/crates/client/src/schedules.rs b/crates/client/src/schedules.rs index 819f161f4..1e58e9f58 100644 --- a/crates/client/src/schedules.rs +++ b/crates/client/src/schedules.rs @@ -17,7 +17,7 @@ use temporalio_common::{ HasWorkflowDefinition, data_converters::{ DataConverter, PayloadConversionError, SerializationContextData, TemporalDeserializable, - TemporalSerializable, + TemporalSerializable, WorkflowSerializationContext, }, payload_visitor::decode_payloads, protos::{ @@ -102,7 +102,11 @@ impl ScheduleWorkflowInput { dc: &DataConverter, ) -> Result, PayloadConversionError> { let ScheduleWorkflowInputRepr::Deferred(v) = self.repr; - v.to_payloads(dc, &SerializationContextData::Workflow).await + v.to_payloads( + dc, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + ) + .await } } @@ -214,6 +218,7 @@ impl ScheduleAction { /// set here will use their proto defaults on the server. #[derive(Debug, Clone, Default, PartialEq, bon::Builder)] #[builder(on(String, into))] +#[non_exhaustive] pub struct ScheduleSpec { /// Interval-based triggers (e.g., every 1 hour). #[builder(default)] @@ -540,7 +545,10 @@ impl ScheduleDescriptionStartWorkflowAction { match &self.input { Some(input) => self .data_converter - .from_payloads(&SerializationContextData::Workflow, input.payloads.clone()) + .from_payloads( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + input.payloads.clone(), + ) .await .map(Some), None => Ok(None), @@ -725,7 +733,7 @@ impl ScheduleDescription { crate::Memo::from_raw( self.raw.memo.clone(), self.data_converter.payload_converter().clone(), - SerializationContextData::Workflow, + SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) } @@ -766,6 +774,7 @@ impl ScheduleDescription { /// Controls what happens when a scheduled workflow would overlap with a running one. #[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +#[non_exhaustive] pub enum ScheduleOverlapPolicy { /// Use the server default (currently Skip). #[default] @@ -977,7 +986,7 @@ impl ScheduleSummary { crate::Memo::from_raw( self.raw.memo.clone(), self.data_converter.payload_converter().clone(), - SerializationContextData::Workflow, + SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) } @@ -1086,7 +1095,7 @@ where decode_payloads( memo, self.client.data_converter().codec(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .await?; } @@ -1141,7 +1150,7 @@ where decode_payloads( &mut response, handle.client.data_converter().codec(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .await?; let description = ScheduleDescription::new( @@ -1606,7 +1615,9 @@ where && let Err(err) = decode_payloads( memo, data_converter.codec(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ), ) .await { @@ -1680,7 +1691,7 @@ mod tests { fn data_converter_with_codec() -> DataConverter { DataConverter::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), XorCodec, ) } @@ -1967,7 +1978,7 @@ mod tests { let data_converter = data_converter_with_codec(); let memo_payload = data_converter .to_payload( - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), &"memo-value".to_owned(), ) .await @@ -1998,7 +2009,7 @@ mod tests { let data_converter = DataConverter::default(); let memo_payload = data_converter .to_payload( - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), &"memo-value".to_owned(), ) .await @@ -2034,7 +2045,7 @@ mod tests { }, data_converter: DataConverter::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), FailingCodec, ), ..Default::default() @@ -2507,7 +2518,10 @@ mod tests { let data_converter = DataConverter::default(); let expected = MultiArgs2("hello".to_string(), 42i32); let payloads = data_converter - .to_payloads(&SerializationContextData::Workflow, &expected) + .to_payloads( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + &expected, + ) .await .unwrap(); let desc = ScheduleDescription::new( @@ -2540,7 +2554,10 @@ mod tests { let data_converter = DataConverter::default(); let expected: String = "not-an-int".to_string(); let payloads = data_converter - .to_payloads(&SerializationContextData::Workflow, &expected) + .to_payloads( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + &expected, + ) .await .unwrap(); let desc = schedule_description_from_response(describe_response_with_start_workflow(Some( diff --git a/crates/client/src/worker.rs b/crates/client/src/worker.rs index 5cd2f06fb..87b32180f 100644 --- a/crates/client/src/worker.rs +++ b/crates/client/src/worker.rs @@ -230,11 +230,11 @@ impl ClientWorkerSetImpl { }; shared_worker.register_callback( worker_instance_key, - WorkerCallbacks { - heartbeat: heartbeat_callback, - heartbeat_success: worker.heartbeat_success_callback(), - cancel_activity: worker.cancel_activity_callback(), - }, + WorkerCallbacks::new( + heartbeat_callback, + worker.heartbeat_success_callback(), + worker.cancel_activity_callback(), + ), ); } @@ -511,6 +511,7 @@ pub type HeartbeatSuccessCallback = Arc; pub type CancelActivityCallback = Arc bool + Send + Sync>; /// Bundles all per-worker callbacks registered with the SharedNamespaceWorker. +#[non_exhaustive] pub struct WorkerCallbacks { /// Callback to collect heartbeat data from the worker. pub heartbeat: HeartbeatCallback, @@ -520,6 +521,21 @@ pub struct WorkerCallbacks { pub cancel_activity: Option, } +impl WorkerCallbacks { + /// Creates a callback bundle for a worker. + pub fn new( + heartbeat: HeartbeatCallback, + heartbeat_success: Option, + cancel_activity: Option, + ) -> Self { + Self { + heartbeat, + heartbeat_success, + cancel_activity, + } + } +} + /// Represents a complete worker that can handle both slot management /// and worker heartbeat functionality. #[cfg_attr(test, mockall::automock)] @@ -688,10 +704,12 @@ mod tests { .expect_task_queue() .return_const(task_queue.clone()); failing_worker.expect_deployment_options().return_const( - WorkerDeploymentOptions::new(temporalio_common::worker::WorkerDeploymentVersion { - deployment_name: "test-deployment".to_string(), - build_id: "build-fail".to_string(), - }) + WorkerDeploymentOptions::new( + temporalio_common::worker::WorkerDeploymentVersion::builder() + .deployment_name("test-deployment".to_string()) + .build_id("build-fail".to_string()) + .build(), + ) .use_worker_versioning(true) .build(), ); @@ -722,13 +740,14 @@ mod tests { succeeding_worker .expect_task_queue() .return_const(task_queue.clone()); - let success_deployment_options = - WorkerDeploymentOptions::new(temporalio_common::worker::WorkerDeploymentVersion { - deployment_name: "test-deployment".to_string(), - build_id: "build-success".to_string(), - }) - .use_worker_versioning(true) - .build(); + let success_deployment_options = WorkerDeploymentOptions::new( + temporalio_common::worker::WorkerDeploymentVersion::builder() + .deployment_name("test-deployment".to_string()) + .build_id("build-success".to_string()) + .build(), + ) + .use_worker_versioning(true) + .build(); succeeding_worker .expect_deployment_options() .return_const(success_deployment_options.clone()); @@ -785,10 +804,12 @@ mod tests { .expect_task_queue() .return_const(task_queue.clone()); failing_worker.expect_deployment_options().return_const( - WorkerDeploymentOptions::new(temporalio_common::worker::WorkerDeploymentVersion { - deployment_name: "test-deployment".to_string(), - build_id: "build-fail".to_string(), - }) + WorkerDeploymentOptions::new( + temporalio_common::worker::WorkerDeploymentVersion::builder() + .deployment_name("test-deployment".to_string()) + .build_id("build-fail".to_string()) + .build(), + ) .use_worker_versioning(true) .build(), ); @@ -996,10 +1017,10 @@ mod tests { .returning(move || { build_id_for_closure.as_ref().map(|build_id| { WorkerDeploymentOptions::new( - temporalio_common::worker::WorkerDeploymentVersion { - deployment_name: deployment_name.clone(), - build_id: build_id.clone(), - }, + temporalio_common::worker::WorkerDeploymentVersion::builder() + .deployment_name(deployment_name.clone()) + .build_id(build_id.clone()) + .build(), ) .use_worker_versioning(true) .build() diff --git a/crates/client/src/workflow_handle.rs b/crates/client/src/workflow_handle.rs index e561ba367..51ad3f749 100644 --- a/crates/client/src/workflow_handle.rs +++ b/crates/client/src/workflow_handle.rs @@ -1,26 +1,33 @@ use crate::{ CancelWorkflowInput, DescribeWorkflowInput, DescribeWorkflowOutput, - FetchWorkflowHistoryPageInput, FetchWorkflowHistoryPageOutput, NamespacedClient, Next, - PollWorkflowUpdateInput, PollWorkflowUpdateOutput, QueryWorkflowInput, QueryWorkflowOutput, - RpcOptions, SignalWorkflowInput, StartWorkflowUpdateInput, StartWorkflowUpdateOutput, - TerminateWorkflowInput, WorkflowCancelOptions, WorkflowDescribeOptions, - WorkflowExecuteUpdateOptions, WorkflowExecutionStatus, WorkflowFetchHistoryOptions, - WorkflowGetResultOptions, WorkflowQueryOptions, WorkflowSignalOptions, - WorkflowStartUpdateOptions, WorkflowTerminateOptions, + FetchWorkflowHistoryPageInput, FetchWorkflowHistoryPageOutput, HistoryEventFilterType, + NamespacedClient, Next, PollWorkflowUpdateInput, PollWorkflowUpdateOutput, QueryWorkflowInput, + QueryWorkflowOutput, RpcOptions, SignalWorkflowInput, StartWorkflowUpdateInput, + StartWorkflowUpdateOutput, TerminateWorkflowInput, WorkflowCancelOptions, + WorkflowDescribeOptions, WorkflowExecuteUpdateOptions, WorkflowExecutionStatus, + WorkflowFetchHistoryOptions, WorkflowGetResultOptions, WorkflowQueryOptions, + WorkflowSignalOptions, WorkflowStartUpdateOptions, WorkflowTerminateOptions, errors::{ WorkflowGetResultError, WorkflowInteractionError, WorkflowQueryError, WorkflowUpdateError, }, grpc::WorkflowService, interceptors, }; -use futures_util::future::BoxFuture; -use std::{fmt::Debug, marker::PhantomData}; +use futures_util::{TryStreamExt, future::BoxFuture, stream, stream::Stream}; +use std::{ + collections::VecDeque, + fmt::Debug, + marker::PhantomData, + pin::Pin, + task::{Context, Poll}, +}; pub use temporalio_common::UntypedWorkflow; use temporalio_common::{ HasWorkflowDefinition, QueryDefinition, SignalDefinition, UpdateDefinition, WorkflowDefinition, data_converters::{ DataConverter, DecodablePayloads, GenericPayloadConverter, PayloadConversionError, PayloadConverter, RawValue, SerializationContext, SerializationContextData, + WorkflowSerializationContext, }, error::IncomingError, payload_visitor::decode_payloads, @@ -28,8 +35,12 @@ use temporalio_common::{ coresdk::FromPayloadsExt, proto_ts_to_system_time, temporal::api::{ - common::v1::{Payload, Payloads, WorkflowExecution as ProtoWorkflowExecution}, - enums::v1::{HistoryEventFilterType, UpdateWorkflowExecutionLifecycleStage}, + common::v1::{Header, Payload, Payloads, WorkflowExecution as ProtoWorkflowExecution}, + enums::v1::{ + HistoryEventFilterType as ProtoHistoryEventFilterType, + QueryRejectCondition as ProtoQueryRejectCondition, + UpdateWorkflowExecutionLifecycleStage, + }, history::{ self, v1::{History, HistoryEvent, history_event::Attributes}, @@ -63,10 +74,7 @@ fn decode_user_metadata( user_metadata: Option, ) -> Result { let payload_converter = PayloadConverter::default(); - let context = SerializationContext { - data: context, - converter: &payload_converter, - }; + let context = SerializationContext::new(context, &payload_converter); let (summary, details) = user_metadata .map(|metadata| (metadata.summary, metadata.details)) .unwrap_or_default(); @@ -96,13 +104,16 @@ impl WorkflowResultDetails { ) -> Result { let payloads = data_converter .codec() - .decode(&SerializationContextData::Workflow, payloads) + .decode( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + payloads, + ) .await?; Ok(Self { payloads: DecodablePayloads::new( payloads, data_converter.payload_converter().clone(), - SerializationContextData::Workflow, + SerializationContextData::Workflow(WorkflowSerializationContext::new()), ), }) } @@ -174,11 +185,13 @@ impl WorkflowExecutionDescription { decode_payloads( &mut raw_description, data_converter.codec(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .await?; - let decoded_metadata = - decode_user_metadata(&SerializationContextData::Workflow, raw_user_metadata)?; + let decoded_metadata = decode_user_metadata( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + raw_user_metadata, + )?; let history_length_raw = raw_description .workflow_execution_info .as_ref() @@ -258,7 +271,7 @@ impl WorkflowExecutionDescription { crate::Memo::from_raw( self.workflow_info().memo.clone(), self.data_converter.payload_converter().clone(), - SerializationContextData::Workflow, + SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) } @@ -331,53 +344,69 @@ impl WorkflowExecutionDescription { } } -// TODO [rust-sdk-branch]: Could implment stream a-la ListWorkflowsStream -/// Workflow execution history returned by `WorkflowHandle::fetch_history`. -#[derive(Debug, Clone)] +/// Workflow execution history returned by [`WorkflowHandle::fetch_history`]. +/// +/// Events and their containing pages are fetched lazily as this stream is polled. Use +/// [`into_events`](Self::into_events) to fetch and collect all events at once. +#[derive(derive_more::Debug)] pub struct WorkflowHistory { - events: Vec, + #[debug(skip)] + inner: Pin> + Send>>, workflow_id: Option, } -impl From for history::v1::History { - fn from(h: WorkflowHistory) -> Self { - Self { events: h.events } - } -} - -/// Error converting a workflow history to or from JSON. -#[derive(Debug, thiserror::Error)] -#[error("failed to convert workflow history JSON: {0}")] -pub struct WorkflowHistoryJsonError(#[from] serde_json::Error); -impl WorkflowHistory { - fn new(events: Vec, workflow_id: Option) -> Self { +impl From for WorkflowHistory { + fn from(history: history::v1::History) -> Self { + let workflow_id = + history + .events + .first() + .and_then(|event| match event.attributes.as_ref() { + Some(Attributes::WorkflowExecutionStartedEventAttributes(attributes)) + if !attributes.workflow_id.is_empty() => + { + Some(attributes.workflow_id.clone()) + } + _ => None, + }); Self { - events, + inner: Box::pin(stream::iter(history.events.into_iter().map(Ok))), workflow_id, } } +} + +impl Stream for WorkflowHistory { + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + self.inner.as_mut().poll_next(cx) + } +} + +/// Error fetching or converting a workflow history. +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum WorkflowHistoryError { + /// Fetching the workflow history failed. + #[error("failed to fetch workflow history: {0}")] + Fetch(#[from] WorkflowInteractionError), + /// Converting the workflow history JSON failed. + #[error("failed to convert workflow history JSON: {0}")] + Json(#[from] serde_json::Error), +} +impl WorkflowHistory { /// Decode a workflow history from JSON bytes. - pub fn from_json(bytes: &[u8]) -> Result { + pub fn from_json(bytes: &[u8]) -> Result { let history: History = serde_json::from_slice(bytes)?; - let workflow_id = history - .events - .first() - .and_then(|event| match event.attributes.as_ref() { - Some(Attributes::WorkflowExecutionStartedEventAttributes(attributes)) => { - Some(attributes) - } - _ => None, - }) - .map(|attributes| attributes.workflow_id.clone()) - .filter(|wfid| !wfid.is_empty()); - Ok(Self::new(history.events, workflow_id)) + Ok(history.into()) } - /// Encode this workflow history as JSON bytes. - pub fn to_json(&self) -> Result, WorkflowHistoryJsonError> { + /// Fetch all remaining events and encode this workflow history as JSON bytes. + pub async fn to_json(self) -> Result, WorkflowHistoryError> { Ok(serde_json::to_vec(&History { - events: self.events.clone(), + events: self.into_events().await?, })?) } @@ -386,14 +415,9 @@ impl WorkflowHistory { self.workflow_id.as_deref() } - /// The history events. - pub fn events(&self) -> &[HistoryEvent] { - &self.events - } - - /// Consume the history and return the events. - pub fn into_events(self) -> Vec { - self.events + /// Fetch all remaining history pages and collect their events. + pub async fn into_events(self) -> Result, WorkflowInteractionError> { + self.inner.try_collect().await } } @@ -415,7 +439,9 @@ impl WorkflowHandle { } /// Holds needed information to refer to a specific workflow run, or workflow execution chain -#[derive(Debug, Clone)] +#[derive(Debug, Clone, bon::Builder)] +#[builder(on(String, into), state_mod(vis = "pub"))] +#[non_exhaustive] pub struct WorkflowExecutionInfo { /// Namespace the workflow lives in. pub namespace: String, @@ -526,6 +552,45 @@ impl UpdateDefinition for UntypedUpdate { } } +/// Shared by [WorkflowHandle::start_update] and the client's update-with-start, which sends the +/// same update request as one of its operations. Update starts always wait for the update to be +/// accepted; results are waited on separately via the update handle. +#[allow(clippy::too_many_arguments)] +pub(crate) fn build_update_workflow_request( + namespace: String, + identity: String, + workflow_id: String, + run_id: String, + update_id: String, + update_name: String, + header: Option
, + payloads: Vec, +) -> UpdateWorkflowExecutionRequest { + UpdateWorkflowExecutionRequest { + namespace, + workflow_execution: Some(ProtoWorkflowExecution { + workflow_id, + run_id, + }), + wait_policy: Some(WaitPolicy { + lifecycle_stage: UpdateWorkflowExecutionLifecycleStage::Accepted.into(), + }), + request: Some(update::v1::Request { + meta: Some(update::v1::Meta { + update_id, + identity, + }), + input: Some(update::v1::Input { + header, + name: update_name, + args: Some(Payloads { payloads }), + }), + ..Default::default() + }), + ..Default::default() + } +} + impl WorkflowHandle where CT: WorkflowService + Clone, @@ -556,7 +621,7 @@ where opts: WorkflowGetResultOptions, ) -> Result where - CT: WorkflowService + NamespacedClient + Clone, + CT: WorkflowService + NamespacedClient + Clone + 'static, { let raw = self.get_result_raw(opts).await?; match raw { @@ -581,7 +646,7 @@ where opts: WorkflowGetResultOptions, ) -> Result, WorkflowInteractionError> where - CT: WorkflowService + NamespacedClient + Clone, + CT: WorkflowService + NamespacedClient + Clone + 'static, { let mut run_id = self.info.run_id.clone().unwrap_or_default(); let fetch_opts = WorkflowFetchHistoryOptions::builder() @@ -592,8 +657,8 @@ where .build(); loop { - let history = self.fetch_history_for_run(&run_id, &fetch_opts).await?; - let mut events = history.into_events(); + let history = self.fetch_history_for_run(&run_id, fetch_opts.clone()); + let mut events = history.into_events().await?; if events.is_empty() { continue; @@ -620,7 +685,7 @@ where .and_then(|p| p.payloads.into_iter().next()) .unwrap_or_default(); let result: W::Output = dc - .from_payload(&SerializationContextData::Workflow, payload) + .from_payload(&SerializationContextData::Workflow(WorkflowSerializationContext::new()), payload) .await?; Ok(WorkflowExecutionResult::Succeeded(result)) } @@ -630,13 +695,13 @@ where decode_payloads( &mut failure, dc.codec(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .await?; let error = dc.failure_converter().to_error( failure, dc.payload_converter(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), )?; Ok(WorkflowExecutionResult::Failed(error)) } @@ -715,16 +780,17 @@ where let data_converter = client.data_converter().clone(); let unencoded_payloads = { let payload_converter = data_converter.payload_converter(); - let context = SerializationContext { - data: &SerializationContextData::Workflow, - converter: payload_converter, - }; + let context_data = SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ); + let context = + SerializationContext::new(&context_data, payload_converter); args.serialize_payloads(&context) }; drop(args); let payloads = data_converter .codec() - .encode(&SerializationContextData::Workflow, unencoded_payloads?) + .encode(&SerializationContextData::Workflow(WorkflowSerializationContext::new()), unencoded_payloads?) .await?; let mut request = SignalWorkflowExecutionRequest { namespace: client.namespace(), @@ -786,16 +852,17 @@ where let data_converter = client.data_converter().clone(); let unencoded_payloads = { let payload_converter = data_converter.payload_converter(); - let context = SerializationContext { - data: &SerializationContextData::Workflow, - converter: payload_converter, - }; + let context_data = SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ); + let context = + SerializationContext::new(&context_data, payload_converter); args.serialize_payloads(&context) }; drop(args); let payloads = data_converter .codec() - .encode(&SerializationContextData::Workflow, unencoded_payloads?) + .encode(&SerializationContextData::Workflow(WorkflowSerializationContext::new()), unencoded_payloads?) .await?; let mut request = QueryWorkflowRequest { namespace: client.namespace(), @@ -810,8 +877,8 @@ where }), query_reject_condition: options .reject_condition - .map(|condition| condition as i32) - .unwrap_or(1), + .map(|condition| ProtoQueryRejectCondition::from(condition) as i32) + .unwrap_or(ProtoQueryRejectCondition::None as i32), } .into_request(); options.rpc_options.apply_to(&mut request); @@ -842,7 +909,10 @@ where self.client .data_converter() - .from_payloads(&SerializationContextData::Workflow, result_payloads) + .from_payloads( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + result_payloads, + ) .await .map_err(WorkflowQueryError::from) } @@ -861,17 +931,7 @@ where U::Output: 'static, { let rpc_options = options.rpc_options.clone(); - let handle = self - .start_update( - update, - input, - WorkflowStartUpdateOptions::builder() - .maybe_update_id(options.update_id) - .maybe_header(options.header) - .rpc_options(rpc_options.clone()) - .build(), - ) - .await?; + let handle = self.start_update(update, input, options.into()).await?; handle.get_result(rpc_options).await } @@ -900,87 +960,100 @@ where Next::new({ let mut client = self.client.clone(); move |input: StartWorkflowUpdateInput| -> BoxFuture< - '_, - Result, - > { - Box::pin(async move { - let (workflow_id, run_id, update_name, args, options) = input.into_parts(); - let data_converter = client.data_converter().clone(); - let unencoded_payloads = { - let payload_converter = data_converter.payload_converter(); - let context = SerializationContext { - data: &SerializationContextData::Workflow, - converter: payload_converter, + '_, + Result, + > { + Box::pin(async move { + let (workflow_id, run_id, update_name, args, options) = + input.into_parts(); + let data_converter = client.data_converter().clone(); + let unencoded_payloads = { + let payload_converter = data_converter.payload_converter(); + let context_data = SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ); + let context = + SerializationContext::new(&context_data, payload_converter); + args.serialize_payloads(&context) }; - args.serialize_payloads(&context) - }; - drop(args); - let payloads = data_converter - .codec() - .encode(&SerializationContextData::Workflow, unencoded_payloads?) - .await?; - let update_id = options - .update_id - .unwrap_or_else(|| Uuid::new_v4().to_string()); - let mut request = UpdateWorkflowExecutionRequest { - namespace: client.namespace(), - workflow_execution: Some(ProtoWorkflowExecution { - workflow_id: workflow_id.clone(), + drop(args); + let payloads = data_converter + .codec() + .encode( + &SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ), + unencoded_payloads?, + ) + .await?; + let update_id = options + .update_id + .unwrap_or_else(|| Uuid::new_v4().to_string()); + let mut request = build_update_workflow_request( + client.namespace(), + client.identity(), + workflow_id.clone(), run_id, - }), - wait_policy: Some(WaitPolicy { - lifecycle_stage: - UpdateWorkflowExecutionLifecycleStage::Accepted.into(), - }), - request: Some(update::v1::Request { - meta: Some(update::v1::Meta { - update_id: update_id.clone(), - identity: client.identity(), - }), - input: Some(update::v1::Input { - header: options.header, - name: update_name, - args: Some(Payloads { payloads }), - }), - ..Default::default() - }), - ..Default::default() - } - .into_request(); - options.rpc_options.apply_to(&mut request); - let response = WorkflowService::update_workflow_execution( - &mut client, - request, - ) - .await - .map_err(WorkflowUpdateError::from_status)? - .into_inner(); - let run_id = response - .update_ref - .as_ref() - .and_then(|reference| reference.workflow_execution.as_ref()) - .map(|execution| execution.run_id.clone()) - .filter(|run_id| !run_id.is_empty()); - Ok(StartWorkflowUpdateOutput::new( - update_id, - workflow_id, - run_id, - response.outcome, - )) - }) - } + update_id.clone(), + update_name, + options.header, + payloads, + ) + .into_request(); + options.rpc_options.apply_to(&mut request); + let response = + WorkflowService::update_workflow_execution(&mut client, request) + .await + .map_err(WorkflowUpdateError::from_status)? + .into_inner(); + let run_id = response + .update_ref + .as_ref() + .and_then(|reference| reference.workflow_execution.as_ref()) + .map(|execution| execution.run_id.clone()) + .filter(|run_id| !run_id.is_empty()); + Ok(StartWorkflowUpdateOutput::new( + update_id, + workflow_id, + run_id, + response.outcome, + )) + }) + } }), ) .await?; - Ok(WorkflowUpdateHandle { - client: self.client.clone(), - update_id: output.update_id, - workflow_id: output.workflow_id, - run_id: output.run_id.or_else(|| self.info().run_id.clone()), - known_outcome: output.known_outcome, - _output: PhantomData, - }) + Ok(WorkflowUpdateHandle::new( + self.client.clone(), + output.update_id, + output.workflow_id, + output.run_id.or_else(|| self.info().run_id.clone()), + output.known_outcome, + )) + } + + /// Get a handle to an existing update. + /// + /// The update definition determines the result type. The returned handle uses this workflow + /// handle's workflow and run IDs and does not validate the update ID until + /// [`get_result`](WorkflowUpdateHandle::get_result) is called. + pub fn get_update_handle( + &self, + update: U, + update_id: impl Into, + ) -> WorkflowUpdateHandle + where + U: UpdateDefinition, + { + let _ = update; + WorkflowUpdateHandle::new( + self.client.clone(), + update_id.into(), + self.info.workflow_id.clone(), + self.info.run_id.clone(), + None, + ) } /// Request cancellation of this workflow. @@ -1134,91 +1207,124 @@ where .await .map_err(WorkflowInteractionError::from) } - /// Fetch workflow execution history. - pub async fn fetch_history( - &self, - opts: WorkflowFetchHistoryOptions, - ) -> Result + /// Fetch workflow execution history as a lazy stream. + /// + /// No request is sent until the returned stream is polled. + pub fn fetch_history(&self, opts: WorkflowFetchHistoryOptions) -> WorkflowHistory where - CT: NamespacedClient, + CT: NamespacedClient + 'static, { let run_id = self.info.run_id.clone().unwrap_or_default(); - self.fetch_history_for_run(&run_id, &opts).await + self.fetch_history_for_run(&run_id, opts) } - /// Fetch history for a specific run_id, handling pagination. - async fn fetch_history_for_run( + fn fetch_history_for_run( &self, run_id: &str, - opts: &WorkflowFetchHistoryOptions, - ) -> Result + opts: WorkflowFetchHistoryOptions, + ) -> WorkflowHistory where - CT: NamespacedClient, + CT: NamespacedClient + 'static, { - let mut all_events = Vec::new(); - let mut next_page_token = vec![]; + let client = self.client.clone(); + let workflow_id = self.info.workflow_id.clone(); + let history_workflow_id = workflow_id.clone(); + let run_id = run_id.to_string(); + + let stream = stream::unfold( + (Vec::new(), VecDeque::new(), false), + move |(mut next_page_token, mut buffer, mut exhausted)| { + let client = client.clone(); + let workflow_id = workflow_id.clone(); + let run_id = run_id.clone(); + let opts = opts.clone(); + + async move { + loop { + if let Some(event) = buffer.pop_front() { + return Some((Ok(event), (next_page_token, buffer, exhausted))); + } - loop { - let output = interceptors::call_fetch_workflow_history_page( - self.client.client_interceptors(), - FetchWorkflowHistoryPageInput { - workflow_id: self.info.workflow_id.clone(), - run_id: run_id.to_string(), - next_page_token, - options: opts.clone(), - }, - Next::new({ - let mut client = self.client.clone(); - move |input: FetchWorkflowHistoryPageInput| -> BoxFuture< - '_, - Result, - > { - Box::pin(async move { - let mut request = GetWorkflowExecutionHistoryRequest { - namespace: client.namespace(), - execution: Some(ProtoWorkflowExecution { - workflow_id: input.workflow_id, - run_id: input.run_id, - }), - next_page_token: input.next_page_token, - skip_archival: input.options.skip_archival, - wait_new_event: input.options.wait_new_event, - history_event_filter_type: input.options.event_filter_type as i32, - ..Default::default() + if exhausted { + return None; + } + + let output = interceptors::call_fetch_workflow_history_page( + client.client_interceptors(), + FetchWorkflowHistoryPageInput { + workflow_id: workflow_id.clone(), + run_id: run_id.clone(), + next_page_token: next_page_token.clone(), + options: opts.clone(), + }, + Next::new({ + let mut rpc_client = client.clone(); + move |input: FetchWorkflowHistoryPageInput| -> BoxFuture< + '_, + Result< + FetchWorkflowHistoryPageOutput, + WorkflowInteractionError, + >, + > { + Box::pin(async move { + let mut request = GetWorkflowExecutionHistoryRequest { + namespace: rpc_client.namespace(), + execution: Some(ProtoWorkflowExecution { + workflow_id: input.workflow_id, + run_id: input.run_id, + }), + next_page_token: input.next_page_token, + skip_archival: input.options.skip_archival, + wait_new_event: input.options.wait_new_event, + history_event_filter_type: + ProtoHistoryEventFilterType::from( + input.options.event_filter_type, + ) + as i32, + ..Default::default() + } + .into_request(); + input.options.rpc_options.apply_to(&mut request); + let response = + WorkflowService::get_workflow_execution_history( + &mut rpc_client, + request, + ) + .await + .map_err(WorkflowInteractionError::from_status)? + .into_inner(); + Ok(FetchWorkflowHistoryPageOutput::new( + response + .history + .map(|history| history.events) + .unwrap_or_default(), + response.next_page_token, + )) + }) + } + }), + ) + .await; + + match output { + Ok(output) => { + exhausted = output.next_page_token.is_empty(); + next_page_token = output.next_page_token; + buffer = output.events.into(); } - .into_request(); - input.options.rpc_options.apply_to(&mut request); - let response = WorkflowService::get_workflow_execution_history( - &mut client, - request, - ) - .await - .map_err(WorkflowInteractionError::from_status)? - .into_inner(); - Ok(FetchWorkflowHistoryPageOutput::new( - response - .history - .map(|history| history.events) - .unwrap_or_default(), - response.next_page_token, - )) - }) + Err(error) => { + return Some((Err(error), (next_page_token, buffer, true))); + } + } } - }), - ) - .await?; + } + }, + ); - all_events.extend(output.events); - if output.next_page_token.is_empty() { - break; - } - next_page_token = output.next_page_token; + WorkflowHistory { + inner: Box::pin(stream), + workflow_id: Some(history_workflow_id), } - - Ok(WorkflowHistory::new( - all_events, - Some(self.info.workflow_id.clone()), - )) } } @@ -1236,6 +1342,23 @@ pub struct WorkflowUpdateHandle { } impl WorkflowUpdateHandle { + pub(crate) fn new( + client: CT, + update_id: String, + workflow_id: String, + run_id: Option, + known_outcome: Option, + ) -> Self { + Self { + client, + update_id, + workflow_id, + run_id, + known_outcome, + _output: PhantomData, + } + } + /// Get the update ID. pub fn id(&self) -> &str { &self.update_id @@ -1323,7 +1446,10 @@ where Some(update::v1::outcome::Value::Success(success)) => self .client .data_converter() - .from_payloads(&SerializationContextData::Workflow, success.payloads) + .from_payloads( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + success.payloads, + ) .await .map_err(WorkflowUpdateError::from), Some(update::v1::outcome::Value::Failure(failure)) => { @@ -1339,8 +1465,15 @@ where #[cfg(test)] mod tests { use super::*; - use crate::test_helpers::XorCodec; - use std::collections::HashMap; + use crate::{ClientInterceptor, test_helpers::XorCodec}; + use futures_util::{FutureExt, StreamExt}; + use std::{ + collections::{HashMap, VecDeque}, + sync::{ + Arc, Mutex, + atomic::{AtomicUsize, Ordering}, + }, + }; use temporalio_common::{ data_converters::DefaultFailureConverter, protos::temporal::api::{ @@ -1349,11 +1482,13 @@ mod tests { history::v1::WorkflowExecutionStartedEventAttributes, sdk::v1::UserMetadata, workflow::v1::WorkflowExecutionConfig, + workflowservice::v1::GetWorkflowExecutionHistoryResponse, }, }; + use tonic::{Request, Response}; - #[test] - fn workflow_history_workflow_id_roundtrips() { + #[tokio::test] + async fn workflow_history_workflow_id_roundtrips() { let event = HistoryEvent { event_id: 1, attributes: Some(Attributes::WorkflowExecutionStartedEventAttributes( @@ -1365,24 +1500,168 @@ mod tests { )), ..Default::default() }; - let history = WorkflowHistory::new(vec![event], None); + let history = WorkflowHistory { + inner: Box::pin(stream::iter(std::iter::once(Ok(event)))), + workflow_id: None, + }; - let bytes = history.to_json().unwrap(); + let bytes = history.to_json().await.unwrap(); let decoded = WorkflowHistory::from_json(&bytes).unwrap(); assert_eq!(decoded.workflow_id(), Some("workflow-id")); } + #[derive(Clone)] + struct MockHistoryClient { + responses: Arc>>>, + calls: Arc, + interceptors: Vec>, + } + + impl NamespacedClient for MockHistoryClient { + fn namespace(&self) -> String { + "test-namespace".to_owned() + } + + fn identity(&self) -> String { + "test-identity".to_owned() + } + + fn client_interceptors(&self) -> &[Arc] { + &self.interceptors + } + } + + impl WorkflowService for MockHistoryClient { + fn get_workflow_execution_history( + &mut self, + _request: Request, + ) -> BoxFuture<'_, Result, tonic::Status>> + { + self.calls.fetch_add(1, Ordering::SeqCst); + let response = self.responses.lock().unwrap().pop_front().unwrap(); + async move { response.map(Response::new) }.boxed() + } + } + + struct CountingHistoryInterceptor(Arc); + + impl ClientInterceptor for CountingHistoryInterceptor { + fn fetch_workflow_history_page<'a>( + &'a self, + input: FetchWorkflowHistoryPageInput, + next: Next< + 'a, + FetchWorkflowHistoryPageInput, + BoxFuture<'a, Result>, + >, + ) -> BoxFuture<'a, Result> + { + self.0.fetch_add(1, Ordering::SeqCst); + next.run(input) + } + } + + fn history_response( + event_ids: impl IntoIterator, + next_page_token: &[u8], + ) -> GetWorkflowExecutionHistoryResponse { + GetWorkflowExecutionHistoryResponse { + history: Some(History { + events: event_ids + .into_iter() + .map(|event_id| HistoryEvent { + event_id, + ..Default::default() + }) + .collect(), + }), + next_page_token: next_page_token.to_vec(), + ..Default::default() + } + } + + fn history_handle( + responses: impl IntoIterator>, + calls: Arc, + interceptors: Vec>, + ) -> WorkflowHandle { + WorkflowHandle::new( + MockHistoryClient { + responses: Arc::new(Mutex::new(responses.into_iter().collect())), + calls, + interceptors, + }, + WorkflowExecutionInfo { + namespace: "test-namespace".to_owned(), + workflow_id: "workflow-id".to_owned(), + run_id: Some("run-id".to_owned()), + first_execution_run_id: None, + }, + ) + } + + #[tokio::test] + async fn workflow_history_fetches_pages_lazily() { + let calls = Arc::new(AtomicUsize::new(0)); + let interceptor_calls = Arc::new(AtomicUsize::new(0)); + let handle = history_handle( + [ + Ok(history_response([], b"second-page")), + Ok(history_response([1, 2], b"third-page")), + Ok(history_response([3], b"")), + ], + calls.clone(), + vec![Arc::new(CountingHistoryInterceptor( + interceptor_calls.clone(), + ))], + ); + + let mut history = handle.fetch_history(WorkflowFetchHistoryOptions::default()); + assert_eq!(calls.load(Ordering::SeqCst), 0); + + assert_eq!(history.next().await.unwrap().unwrap().event_id, 1); + assert_eq!(calls.load(Ordering::SeqCst), 2); + assert_eq!(history.next().await.unwrap().unwrap().event_id, 2); + assert_eq!(calls.load(Ordering::SeqCst), 2); + assert_eq!(history.next().await.unwrap().unwrap().event_id, 3); + assert_eq!(calls.load(Ordering::SeqCst), 3); + assert!(history.next().await.is_none()); + assert_eq!(interceptor_calls.load(Ordering::SeqCst), 3); + } + + #[tokio::test] + async fn workflow_history_yields_page_error_then_ends() { + let calls = Arc::new(AtomicUsize::new(0)); + let handle = history_handle( + [ + Ok(history_response([1], b"second-page")), + Err(tonic::Status::unavailable("history unavailable")), + ], + calls.clone(), + Vec::new(), + ); + let mut history = handle.fetch_history(WorkflowFetchHistoryOptions::default()); + + assert_eq!(history.next().await.unwrap().unwrap().event_id, 1); + assert!(matches!( + history.next().await.unwrap(), + Err(WorkflowInteractionError::Rpc(status)) if status.code() == tonic::Code::Unavailable + )); + assert!(history.next().await.is_none()); + assert_eq!(calls.load(Ordering::SeqCst), 2); + } + #[tokio::test] async fn workflow_result_details_support_typed_decoding() { let converter = DataConverter::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), XorCodec, ); let payloads = converter .to_payloads( - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), &"workflow-result-details".to_owned(), ) .await @@ -1415,12 +1694,12 @@ mod tests { async fn workflow_description_memo_uses_saved_converter() { let converter = DataConverter::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), XorCodec, ); let encoded = converter .to_payload( - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), &"memo-value".to_owned(), ) .await @@ -1451,19 +1730,31 @@ mod tests { async fn workflow_description_accessors_expose_decoded_fields() { let converter = DataConverter::default(); let memo_payload = converter - .to_payload(&SerializationContextData::Workflow, &"memo-value") + .to_payload( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + &"memo-value", + ) .await .unwrap(); let search_attr_payload = converter - .to_payload(&SerializationContextData::Workflow, &"search-value") + .to_payload( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + &"search-value", + ) .await .unwrap(); let summary_payload = converter - .to_payload(&SerializationContextData::Workflow, &"workflow summary") + .to_payload( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + &"workflow summary", + ) .await .unwrap(); let details_payload = converter - .to_payload(&SerializationContextData::Workflow, &"workflow details") + .to_payload( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + &"workflow details", + ) .await .unwrap(); let description = WorkflowExecutionDescription::new( diff --git a/crates/client/tests/typed_signal_with_start.rs b/crates/client/tests/typed_signal_with_start.rs new file mode 100644 index 000000000..241b03dd1 --- /dev/null +++ b/crates/client/tests/typed_signal_with_start.rs @@ -0,0 +1,5 @@ +#[test] +fn typed_signal_with_start_build_tests() { + let tests = trybuild::TestCases::new(); + tests.compile_fail("tests/typed_signal_with_start/*_fail.rs"); +} diff --git a/crates/client/tests/typed_signal_with_start/mismatched_signal_fail.rs b/crates/client/tests/typed_signal_with_start/mismatched_signal_fail.rs new file mode 100644 index 000000000..f0f9dea5b --- /dev/null +++ b/crates/client/tests/typed_signal_with_start/mismatched_signal_fail.rs @@ -0,0 +1,45 @@ +use temporalio_client::{Client, WorkflowStartOptions}; +use temporalio_macros::{workflow, workflow_methods}; +use temporalio_workflow::{SyncWorkflowContext, WorkflowContext, WorkflowResult}; + +#[workflow] +#[derive(Default)] +struct FirstWorkflow; + +#[workflow_methods] +impl FirstWorkflow { + #[run] + async fn run(_ctx: &mut WorkflowContext) -> WorkflowResult<()> { + Ok(()) + } + + #[signal] + fn first_signal(&mut self, _ctx: &mut SyncWorkflowContext, _input: String) {} +} + +#[workflow] +#[derive(Default)] +struct SecondWorkflow; + +#[workflow_methods] +impl SecondWorkflow { + #[run] + async fn run(_ctx: &mut WorkflowContext) -> WorkflowResult<()> { + Ok(()) + } + + #[signal] + fn second_signal(&mut self, _ctx: &mut SyncWorkflowContext, _input: String) {} +} + +fn mismatched_signal(client: &Client) { + let _ = client.signal_with_start_workflow( + FirstWorkflow::run, + (), + SecondWorkflow::second_signal, + "signal".to_owned(), + WorkflowStartOptions::new("task-queue", "workflow-id").build(), + ); +} + +fn main() {} diff --git a/crates/client/tests/typed_signal_with_start/mismatched_signal_fail.stderr b/crates/client/tests/typed_signal_with_start/mismatched_signal_fail.stderr new file mode 100644 index 000000000..671664706 --- /dev/null +++ b/crates/client/tests/typed_signal_with_start/mismatched_signal_fail.stderr @@ -0,0 +1,34 @@ +error[E0271]: type mismatch resolving `::Workflow == Run` + --> tests/typed_signal_with_start/mismatched_signal_fail.rs:39:9 + | +36 | let _ = client.signal_with_start_workflow( + | -------------------------- required by a bound introduced by this call +... +39 | SecondWorkflow::second_signal, + | ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ type mismatch resolving `::Workflow == Run` + | +note: expected this to be `first_workflow::Run` + --> tests/typed_signal_with_start/mismatched_signal_fail.rs:24:1 + | +24 | #[workflow_methods] + | ^^^^^^^^^^^^^^^^^^^ + = note: `second_workflow::Run` and `first_workflow::Run` have similar names, but are actually distinct types +note: `second_workflow::Run` is defined in module `crate::second_workflow` of the current crate + --> tests/typed_signal_with_start/mismatched_signal_fail.rs:24:1 + | +24 | #[workflow_methods] + | ^^^^^^^^^^^^^^^^^^^ +note: `first_workflow::Run` is defined in module `crate::first_workflow` of the current crate + --> tests/typed_signal_with_start/mismatched_signal_fail.rs:9:1 + | + 9 | #[workflow_methods] + | ^^^^^^^^^^^^^^^^^^^ +note: required by a bound in `temporalio_client::Client::signal_with_start_workflow` + --> src/lib.rs + | + | pub async fn signal_with_start_workflow( + | -------------------------- required by a bound in this associated function +... + | S: SignalDefinition, + | ^^^^^^^^^^^^^^^^^ required by this bound in `Client::signal_with_start_workflow` + = note: this error originates in the attribute macro `workflow_methods` (in Nightly builds, run with -Z macro-backtrace for more info) diff --git a/crates/common-wasm/Cargo.toml b/crates/common-wasm/Cargo.toml index 488c56e4f..ed5ad5b8d 100644 --- a/crates/common-wasm/Cargo.toml +++ b/crates/common-wasm/Cargo.toml @@ -1,12 +1,13 @@ [package] name = "temporalio-common-wasm" -version = "0.6.0" +version = "1.0.0" edition = "2024" +rust-version = "1.88.0" authors = ["Temporal Technologies Inc. "] license-file = { workspace = true } description = "WASM-safe shared functionality for the Temporal Rust workflow surface" homepage = "https://temporal.io/" -repository = "https://github.com/temporalio/sdk-core" +repository = "https://github.com/temporalio/sdk-rust" keywords = ["temporal", "workflow"] categories = ["development-tools"] @@ -45,7 +46,7 @@ tracing-core = "0.1" url = "2.5" [dependencies.temporalio-protos] path = "../protos" -version = "0.6" +version = "~0.9.0" [lints] workspace = true diff --git a/crates/common-wasm/README.md b/crates/common-wasm/README.md new file mode 100644 index 000000000..722c68f06 --- /dev/null +++ b/crates/common-wasm/README.md @@ -0,0 +1,8 @@ +# `temporalio-common-wasm` + +[![crates.io](https://img.shields.io/crates/v/temporalio-common-wasm.svg)](https://crates.io/crates/temporalio-common-wasm) +[![docs.rs](https://docs.rs/temporalio-common-wasm/badge.svg)](https://docs.rs/temporalio-common-wasm) + +Part of [Temporal](https://temporal.io)'s [Rust SDK](https://github.com/temporalio/sdk-rust). + +WASM-safe shared types, serialization, and protobuf support for authoring Temporal Workflows. diff --git a/crates/common-wasm/src/activity_definition.rs b/crates/common-wasm/src/activity_definition.rs index 04e6d5e39..6fcfa80ee 100644 --- a/crates/common-wasm/src/activity_definition.rs +++ b/crates/common-wasm/src/activity_definition.rs @@ -44,6 +44,7 @@ impl ActivityDefinition for UntypedActivity { /// Returned as errors from activity functions. #[derive(Debug)] +#[non_exhaustive] pub enum ActivityError { /// Return this error to attach application-failure metadata to an activity failure. Application(Box), diff --git a/crates/common-wasm/src/data_converters.rs b/crates/common-wasm/src/data_converters.rs index ffd0aba56..3529e6c75 100644 --- a/crates/common-wasm/src/data_converters.rs +++ b/crates/common-wasm/src/data_converters.rs @@ -2,16 +2,22 @@ //! serialization related functionality. mod failure_converter; +mod well_known; pub use failure_converter::{ - ActivityExecutionDecodeHint, ChildWorkflowExecutionDecodeHint, ChildWorkflowStartDecodeHint, - DefaultFailureConverter, FailureConverter, FailureDecodeHint, WorkflowSignalDecodeHint, + ActivityExecutionDecodeHint, CancelExternalWorkflowDecodeHint, + ChildWorkflowExecutionDecodeHint, ChildWorkflowStartDecodeHint, CommonAttributes, + DefaultFailureConverter, FailureConverter, FailureDecodeHint, NoopDecodeHint, + WorkflowSignalDecodeHint, }; +use well_known::{BINARY_NULL_ENCODING_VAL, WellKnownType, binary_null_payload}; -use crate::protos::temporal::api::common::v1::Payload; +use crate::protos::{ENCODING_PAYLOAD_KEY, JSON_ENCODING_VAL, temporal::api::common::v1::Payload}; use futures::{FutureExt, future::BoxFuture}; use std::{collections::HashMap, sync::Arc}; +const PROTOBUF_ENCODING_VAL: &str = "binary/protobuf"; + /// Combines a [`PayloadConverter`], [`FailureConverter`], and [`PayloadCodec`] to handle all /// serialization needs for communicating with the Temporal server. #[derive(Clone)] @@ -50,10 +56,7 @@ impl DataConverter { data: &SerializationContextData, val: &T, ) -> Result { - let context = SerializationContext { - data, - converter: &self.payload_converter, - }; + let context = SerializationContext::new(data, &self.payload_converter); let payload = self.payload_converter.to_payload(&context, val)?; let encoded = self.codec.encode(data, vec![payload]).await?; encoded @@ -68,10 +71,7 @@ impl DataConverter { data: &SerializationContextData, payload: Payload, ) -> Result { - let context = SerializationContext { - data, - converter: &self.payload_converter, - }; + let context = SerializationContext::new(data, &self.payload_converter); let decoded = self.codec.decode(data, vec![payload]).await?; let payload = decoded .into_iter() @@ -86,10 +86,7 @@ impl DataConverter { data: &SerializationContextData, val: &T, ) -> Result, PayloadConversionError> { - let context = SerializationContext { - data, - converter: &self.payload_converter, - }; + let context = SerializationContext::new(data, &self.payload_converter); let payloads = self.payload_converter.to_payloads(&context, val)?; self.codec.encode(data, payloads).await } @@ -100,10 +97,7 @@ impl DataConverter { data: &SerializationContextData, payloads: Vec, ) -> Result { - let context = SerializationContext { - data, - converter: &self.payload_converter, - }; + let context = SerializationContext::new(data, &self.payload_converter); let decoded = self.codec.decode(data, payloads).await?; self.payload_converter.from_payloads(&context, decoded) } @@ -147,15 +141,61 @@ impl DataConverter { } } +/// Data available when serializing in a workflow context. +#[derive(Clone, Debug, PartialEq, Eq)] +#[non_exhaustive] +pub struct WorkflowSerializationContext {} + +#[allow(clippy::new_without_default)] +impl WorkflowSerializationContext { + /// Creates an empty workflow serialization context. + /// + /// **Experimental:** This constructor may change when workflow context data is added. + pub fn new() -> Self { + Self {} + } +} + +/// Data available when serializing in an activity context. +#[derive(Clone, Debug, PartialEq, Eq)] +#[non_exhaustive] +pub struct ActivitySerializationContext {} + +#[allow(clippy::new_without_default)] +impl ActivitySerializationContext { + /// Creates an empty activity serialization context. + /// + /// **Experimental:** This constructor may change when activity context data is added. + pub fn new() -> Self { + Self {} + } +} + +/// Data available when serializing in a Nexus context. +#[derive(Clone, Debug, PartialEq, Eq)] +#[non_exhaustive] +pub struct NexusSerializationContext {} + +#[allow(clippy::new_without_default)] +impl NexusSerializationContext { + /// Creates an empty Nexus serialization context. + /// + /// **Experimental:** This constructor may change when Nexus context data is added. + pub fn new() -> Self { + Self {} + } +} + /// Data about the serialization context, indicating where the serialization is occurring. -#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[derive(Clone, Debug, PartialEq, Eq)] +#[non_exhaustive] pub enum SerializationContextData { /// Serialization is occurring in a workflow context. - Workflow, + Workflow(WorkflowSerializationContext), /// Serialization is occurring in an activity context. - Activity, + Activity(ActivitySerializationContext), /// Serialization is occurring in a nexus context. - Nexus, + Nexus(NexusSerializationContext), /// No specific serialization context. None, } @@ -163,14 +203,24 @@ pub enum SerializationContextData { /// Context for serialization operations, including the kind of context and the /// payload converter for nested serialization. #[derive(Clone, Copy)] +#[non_exhaustive] pub struct SerializationContext<'a> { /// The kind of serialization context (workflow, activity, etc.). pub data: &'a SerializationContextData, /// Allows nested types to serialize their contents using the same converter. pub converter: &'a PayloadConverter, } + +impl<'a> SerializationContext<'a> { + /// Creates a serialization context for the given execution context and payload converter. + pub fn new(data: &'a SerializationContextData, converter: &'a PayloadConverter) -> Self { + Self { data, converter } + } +} + /// Converts values to and from [`Payload`]s using different encoding strategies. #[derive(Clone)] +#[non_exhaustive] pub enum PayloadConverter { /// Uses a serde-based converter for encoding/decoding. Serde(Arc), @@ -207,6 +257,7 @@ impl Default for PayloadConverter { /// Errors that can occur during payload conversion. #[derive(Debug)] +#[non_exhaustive] pub enum PayloadConversionError { /// The payload's encoding does not match what the converter expects. WrongEncoding, @@ -354,10 +405,7 @@ impl DecodablePayloads { &self, ) -> Result { self.payload_converter.from_payloads( - &SerializationContext { - data: &self.context, - converter: &self.payload_converter, - }, + &SerializationContext::new(&self.context, &self.payload_converter), self.payloads.clone(), ) } @@ -401,10 +449,7 @@ impl RawValue { RawValue::new(vec![ converter .to_payload( - &SerializationContext { - data: &SerializationContextData::None, - converter, - }, + &SerializationContext::new(&SerializationContextData::None, converter), value, ) .unwrap(), @@ -415,10 +460,7 @@ impl RawValue { pub fn to_value(self, converter: &PayloadConverter) -> T { converter .from_payload( - &SerializationContext { - data: &SerializationContextData::None, - converter, - }, + &SerializationContext::new(&SerializationContextData::None, converter), self.payloads.into_iter().next().unwrap(), ) .unwrap() @@ -495,31 +537,55 @@ impl GenericPayloadConverter for PayloadConverter { context: &SerializationContext<'_>, val: &T, ) -> Result { - // If a single payload is explicitly needed for `()`, then produce a null payload - if std::any::TypeId::of::() == std::any::TypeId::of::<()>() { - return Ok(Payload { - metadata: { - let mut hm = HashMap::new(); - hm.insert("encoding".to_string(), b"binary/null".to_vec()); - hm - }, - data: vec![], - external_payloads: vec![], - }); - } - let mut payloads = self.to_payloads(context, val)?; - if payloads.len() != 1 { - return Err(PayloadConversionError::WrongEncoding); + match self { + PayloadConverter::Serde(pc) => { + if let Some(well_known_type) = WellKnownType::of::() { + Ok(well_known_type.to_payload(val)) + } else { + pc.to_payload(context.data, val.as_serde()?) + } + } + PayloadConverter::UseWrappers => T::to_payload(val, context), + PayloadConverter::Composite(composite) => { + for converter in &composite.converters { + match converter.to_payload(context, val) { + Ok(payload) => return Ok(payload), + Err(PayloadConversionError::WrongEncoding) => continue, + Err(e) => return Err(e), + } + } + Err(PayloadConversionError::WrongEncoding) + } } - Ok(payloads.pop().unwrap()) } fn from_payload( &self, context: &SerializationContext<'_>, - payload: Payload, + mut payload: Payload, ) -> Result { - self.from_payloads(context, vec![payload]) + match self { + PayloadConverter::Serde(pc) => { + if let Some(well_known_type) = WellKnownType::of::() { + payload = match well_known_type.try_from_payload(payload) { + Ok(value) => return Ok(value), + Err(payload) => payload, + }; + } + T::from_serde(pc.as_ref(), context, payload) + } + PayloadConverter::UseWrappers => T::from_payload(context, payload), + PayloadConverter::Composite(composite) => { + for converter in &composite.converters { + match converter.from_payload(context, payload.clone()) { + Ok(value) => return Ok(value), + Err(PayloadConversionError::WrongEncoding) => continue, + Err(e) => return Err(e), + } + } + Err(PayloadConversionError::WrongEncoding) + } + } } fn to_payloads( @@ -529,10 +595,8 @@ impl GenericPayloadConverter for PayloadConverter { ) -> Result, PayloadConversionError> { match self { PayloadConverter::Serde(pc) => { - // Since Rust SDK uses () to denote no input, we must match other SDKs by producing - // no payloads for it. - if std::any::TypeId::of::() == std::any::TypeId::of::<()>() { - Ok(Vec::new()) + if let Some(well_known_type) = WellKnownType::of::() { + Ok(well_known_type.to_payloads(val)) } else { Ok(vec![pc.to_payload(context.data, val.as_serde()?)?]) } @@ -554,23 +618,21 @@ impl GenericPayloadConverter for PayloadConverter { fn from_payloads( &self, context: &SerializationContext<'_>, - payloads: Vec, + mut payloads: Vec, ) -> Result { - // Accept empty payloads (no args) and a single binary/null payload (result from a - // workflow/update with () return type as (). - if std::any::TypeId::of::() == std::any::TypeId::of::<()>() - && is_unit_payloads(&payloads) - { - let boxed: Box = Box::new(()); - return Ok(*boxed.downcast::().unwrap()); - } - match self { PayloadConverter::Serde(pc) => { + if let Some(well_known_type) = WellKnownType::of::() { + payloads = match well_known_type.try_from_payloads(payloads) { + Ok(value) => return Ok(value), + Err(payloads) => payloads, + }; + } if payloads.len() != 1 { return Err(PayloadConversionError::WrongEncoding); } - T::from_serde(pc.as_ref(), context, payloads.into_iter().next().unwrap()) + let payload = payloads.into_iter().next().unwrap(); + T::from_serde(pc.as_ref(), context, payload) } PayloadConverter::UseWrappers => T::from_payloads(context, payloads), PayloadConverter::Composite(composite) => { @@ -587,21 +649,6 @@ impl GenericPayloadConverter for PayloadConverter { } } -fn is_unit_payloads(payloads: &[Payload]) -> bool { - match payloads { - [] => true, - [payload] => { - payload.data.is_empty() - && payload - .metadata - .get("encoding") - .map(|encoding| encoding == b"binary/null") - .unwrap_or(false) - } - _ => false, - } -} - // TODO [rust-sdk-branch]: Potentially allow opt-out / no-serde compile flags impl TemporalSerializable for T where @@ -638,10 +685,16 @@ impl ErasedSerdePayloadConverter for SerdeJsonPayloadConverter { ) -> Result { let as_json = serde_json::to_vec(value) .map_err(|e| PayloadConversionError::EncodingError(e.into()))?; + if as_json.as_slice() == b"null" { + return Ok(binary_null_payload()); + } Ok(Payload { metadata: { let mut hm = HashMap::new(); - hm.insert("encoding".to_string(), b"json/plain".to_vec()); + hm.insert( + ENCODING_PAYLOAD_KEY.to_string(), + JSON_ENCODING_VAL.as_bytes().to_vec(), + ); hm }, data: as_json, @@ -654,12 +707,18 @@ impl ErasedSerdePayloadConverter for SerdeJsonPayloadConverter { _: &SerializationContextData, payload: Payload, ) -> Result>, PayloadConversionError> { - let encoding = payload.metadata.get("encoding").map(|v| v.as_slice()); - if encoding != Some(b"json/plain".as_slice()) { + let encoding = payload + .metadata + .get(ENCODING_PAYLOAD_KEY) + .map(|v| v.as_slice()); + let json_v = if encoding == Some(JSON_ENCODING_VAL.as_bytes()) { + serde_json::from_slice(&payload.data) + .map_err(|e| PayloadConversionError::EncodingError(Box::new(e)))? + } else if encoding == Some(BINARY_NULL_ENCODING_VAL.as_bytes()) { + serde_json::Value::Null + } else { return Err(PayloadConversionError::WrongEncoding); - } - let json_v: serde_json::Value = serde_json::from_slice(&payload.data) - .map_err(|e| PayloadConversionError::EncodingError(Box::new(e)))?; + }; Ok(Box::new(::erase(json_v))) } } @@ -694,7 +753,10 @@ where Ok(Payload { metadata: { let mut hm = HashMap::new(); - hm.insert("encoding".to_string(), b"binary/protobuf".to_vec()); + hm.insert( + ENCODING_PAYLOAD_KEY.to_string(), + PROTOBUF_ENCODING_VAL.as_bytes().to_vec(), + ); hm }, data: as_proto, @@ -713,8 +775,8 @@ where where Self: Sized, { - let encoding = p.metadata.get("encoding").map(|v| v.as_slice()); - if encoding != Some(b"binary/protobuf".as_slice()) { + let encoding = p.metadata.get(ENCODING_PAYLOAD_KEY).map(|v| v.as_slice()); + if encoding != Some(PROTOBUF_ENCODING_VAL.as_bytes()) { return Err(PayloadConversionError::WrongEncoding); } T::decode(p.data.as_slice()) @@ -733,7 +795,7 @@ impl Default for DataConverter { fn default() -> Self { Self::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), DefaultPayloadCodec, ) } @@ -817,28 +879,14 @@ impl_multi_args!(MultiArgs6; 6; 0: A, 1: B, 2: C, 3: D, 4: E, 5: F); #[cfg(test)] mod tests { use super::*; + use crate::data_converters::well_known::BINARY_PLAIN_ENCODING_VAL; + use rstest::rstest; #[test] - fn test_empty_payloads_as_unit_type() { - let converter = PayloadConverter::default(); - let ctx = SerializationContext { - data: &SerializationContextData::Workflow, - converter: &converter, - }; - - let empty_payloads: Vec = vec![]; - let result: Result<(), _> = converter.from_payloads(&ctx, empty_payloads); - - assert!(result.is_ok(), "Empty payloads should deserialize as ()"); - } - - #[test] - fn test_unit_type_roundtrip_serde() { + fn unit_payloads_roundtrip() { let converter = PayloadConverter::serde_json(); - let ctx = SerializationContext { - data: &SerializationContextData::Workflow, - converter: &converter, - }; + let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let ctx = SerializationContext::new(&context_data, &converter); let payloads = converter.to_payloads(&ctx, &()).unwrap(); assert!(payloads.is_empty()); @@ -847,59 +895,119 @@ mod tests { assert_eq!(result, ()); } - #[test] - fn test_unit_composite_roundtrip() { + #[rstest] + #[case::unit((), BINARY_NULL_ENCODING_VAL, b"")] + #[case::none_string(Option::::None, BINARY_NULL_ENCODING_VAL, b"")] + #[case::some_string( + Some("value".to_string()), + JSON_ENCODING_VAL, + br#""value""# + )] + #[case::bytes(vec![0_u8, 1, 2, 255], BINARY_PLAIN_ENCODING_VAL, &[0, 1, 2, 255])] + #[case::some_bytes( + Some(vec![1_u8, 2, 3]), + BINARY_PLAIN_ENCODING_VAL, + &[1, 2, 3] + )] + #[case::none_bytes(Option::>::None, BINARY_NULL_ENCODING_VAL, b"")] + fn value_encodes_as( + #[case] value: T, + #[case] expected_encoding: &str, + #[case] expected_data: &[u8], + ) where + T: TemporalSerializable + std::fmt::Debug + 'static, + { let converter = PayloadConverter::default(); - let ctx = SerializationContext { - data: &SerializationContextData::Workflow, - converter: &converter, - }; - - let payloads = converter.to_payloads(&ctx, &()).unwrap(); - assert!(payloads.is_empty()); + let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let ctx = SerializationContext::new(&context_data, &converter); - let result: () = converter.from_payloads(&ctx, payloads).unwrap(); - assert_eq!(result, ()); - } + let payload = converter.to_payload(&ctx, &value).unwrap(); - #[test] - fn test_unit_to_payload_roundtrip() { + assert_eq!( + payload.metadata.get(ENCODING_PAYLOAD_KEY).unwrap(), + expected_encoding.as_bytes() + ); + assert_eq!(payload.data, expected_data); + } + + #[rstest] + #[case::unit(BINARY_NULL_ENCODING_VAL, b"", ())] + #[case::none_string(BINARY_NULL_ENCODING_VAL, b"", Option::::None)] + #[case::legacy_none_string(JSON_ENCODING_VAL, b"null", Option::::None)] + #[case::bytes(BINARY_PLAIN_ENCODING_VAL, &[0, 1, 2, 255], vec![0_u8, 1, 2, 255])] + #[case::legacy_bytes(JSON_ENCODING_VAL, b"[3,2,1]", vec![3_u8, 2, 1])] + #[case::some_bytes( + BINARY_PLAIN_ENCODING_VAL, + &[1, 2, 3], + Some(vec![1_u8, 2, 3]) + )] + #[case::none_bytes(BINARY_NULL_ENCODING_VAL, b"", Option::>::None)] + #[case::legacy_some_bytes( + JSON_ENCODING_VAL, + b"[3,2,1]", + Some(vec![3_u8, 2, 1]) + )] + #[case::legacy_none_bytes(JSON_ENCODING_VAL, b"null", Option::>::None)] + fn payload_decodes_as(#[case] encoding: &str, #[case] data: &[u8], #[case] expected: T) + where + T: TemporalDeserializable + std::fmt::Debug + PartialEq + 'static, + { let converter = PayloadConverter::default(); - let ctx = SerializationContext { - data: &SerializationContextData::Workflow, - converter: &converter, - }; + let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let ctx = SerializationContext::new(&context_data, &converter); - let mut payloads = vec![converter.to_payload(&ctx, &()).unwrap()]; - assert!(is_unit_payloads(&payloads)); - let result: () = converter - .from_payload(&ctx, payloads.pop().unwrap()) + let actual: T = converter + .from_payload( + &ctx, + Payload { + metadata: HashMap::from([( + ENCODING_PAYLOAD_KEY.to_string(), + encoding.as_bytes().to_vec(), + )]), + data: data.to_vec(), + external_payloads: vec![], + }, + ) .unwrap(); - assert_eq!(result, ()); + assert_eq!(actual, expected); } #[test] - fn test_unit_use_wrappers_returns_wrong_encoding() { + fn use_wrappers_returns_wrong_encoding_for_standard_types() { let converter = PayloadConverter::UseWrappers; - let ctx = SerializationContext { - data: &SerializationContextData::Workflow, - converter: &converter, - }; + let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let ctx = SerializationContext::new(&context_data, &converter); + + let result = converter.to_payload(&ctx, &()); + assert!( + matches!(result, Err(PayloadConversionError::WrongEncoding)), + "{result:?}" + ); let result = converter.to_payloads(&ctx, &()); assert!( matches!(result, Err(PayloadConversionError::WrongEncoding)), "{result:?}" ); + + let result = converter.to_payloads(&ctx, &vec![1_u8, 2, 3]); + assert!( + matches!(result, Err(PayloadConversionError::WrongEncoding)), + "{result:?}" + ); + + let result: Result<(), _> = converter.from_payload(&ctx, binary_null_payload()); + assert!( + matches!(result, Err(PayloadConversionError::WrongEncoding)), + "{result:?}" + ); } #[test] fn multi_args_round_trip() { let converter = PayloadConverter::default(); - let ctx = SerializationContext { - data: &SerializationContextData::Workflow, - converter: &converter, - }; + let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let ctx = SerializationContext::new(&context_data, &converter); let args = MultiArgs2("hello".to_string(), 42i32); let payloads = converter.to_payloads(&ctx, &args).unwrap(); @@ -909,58 +1017,52 @@ mod tests { assert_eq!(result, args); } + #[test] + fn empty_payloads_do_not_decode_as_option() { + let converter = PayloadConverter::default(); + let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let ctx = SerializationContext::new(&context_data, &converter); + + let result: Result, _> = converter.from_payloads(&ctx, vec![]); + assert!(matches!(result, Err(PayloadConversionError::WrongEncoding))); + } + #[test] fn multi_args_from_tuple() { let args: MultiArgs2 = ("hello".to_string(), 42i32).into(); assert_eq!(args, MultiArgs2("hello".to_string(), 42)); } - fn decodable_from_value(value: &T) -> DecodablePayloads { + #[rstest] + #[case::string("hello".to_string())] + #[case::some_string(Some("hello".to_string()))] + #[case::none_string(Option::::None)] + #[case::unit(())] + #[case::strings(vec!["hello".to_string(), "world".to_string()])] + #[case::bytes(vec![1_u8, 2, 3])] + #[case::some_bytes(Some(vec![1_u8, 2, 3]))] + #[case::none_bytes(Option::>::None)] + fn decodable_payloads_roundtrip(#[case] value: T) + where + T: TemporalSerializable + TemporalDeserializable + std::fmt::Debug + PartialEq + 'static, + { let converter = PayloadConverter::default(); let payloads = converter .to_payloads( - &SerializationContext { - data: &SerializationContextData::Workflow, - converter: &converter, - }, - value, + &SerializationContext::new( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + &converter, + ), + &value, ) .unwrap(); - DecodablePayloads::new(payloads, converter, SerializationContextData::Workflow) - } - #[test] - fn decodable_payloads_roundtrip_string() { - let payloads = decodable_from_value(&"hello".to_string()); - - let result: String = payloads.deserialize().unwrap(); - - assert_eq!(result, "hello"); - } - - #[test] - fn decodable_payloads_roundtrip_option_string() { - let payloads = decodable_from_value(&Some("hello".to_string())); - - let result: Option = payloads.deserialize().unwrap(); - - assert_eq!(result, Some("hello".to_string())); - } - - #[test] - fn decodable_payloads_roundtrip_unit() { - let payloads = decodable_from_value(&()); - - let result: () = payloads.deserialize().unwrap(); - - assert_eq!(result, ()); - } - - #[test] - fn decodable_payloads_roundtrip_vec_string() { - let payloads = decodable_from_value(&vec!["hello".to_string(), "world".to_string()]); - - let result: Vec = payloads.deserialize().unwrap(); + let payloads = DecodablePayloads::new( + payloads, + converter, + SerializationContextData::Workflow(WorkflowSerializationContext::new()), + ); - assert_eq!(result, vec!["hello".to_string(), "world".to_string()]); + let result: T = payloads.deserialize().unwrap(); + assert_eq!(result, value); } } diff --git a/crates/common-wasm/src/data_converters/failure_converter.rs b/crates/common-wasm/src/data_converters/failure_converter.rs index 3c17ba3c8..a6e6befde 100644 --- a/crates/common-wasm/src/data_converters/failure_converter.rs +++ b/crates/common-wasm/src/data_converters/failure_converter.rs @@ -9,18 +9,24 @@ //! [`FailureDecodeHint`] implementations adapt that normalized value into the caller-facing error //! type they expect. -use super::{PayloadConversionError, PayloadConverter, SerializationContextData}; +use super::{ + GenericPayloadConverter, PayloadConversionError, PayloadConverter, SerializationContext, + SerializationContextData, +}; use crate::{ error::{ - ActivityExecutionError, ActivityFailureError, ApplicationFailure, CancelledError, - ChildWorkflowExecutionError, ChildWorkflowFailureError, ChildWorkflowStartError, - IncomingError, IncomingNexusHandlerError, IncomingNexusOperationExecutionError, - OutgoingActivityError, OutgoingError, OutgoingWorkflowError, ResetWorkflowError, - ServerError, TerminatedError, TimeoutError, WorkflowSignalError, - WorkflowSignalFailureError, + ActivityExecutionError, ActivityFailureError, ApplicationFailure, + CancelExternalWorkflowError, CancelledError, ChildWorkflowExecutionError, + ChildWorkflowFailureError, ChildWorkflowStartError, IncomingError, + IncomingNexusHandlerError, IncomingNexusOperationExecutionError, OutgoingActivityError, + OutgoingError, OutgoingWorkflowError, ResetWorkflowError, ServerError, TerminatedError, + TimeoutError, WorkflowCancelFailureError, WorkflowSignalError, WorkflowSignalFailureError, }, protos::temporal::api::{ - enums::v1::ApplicationErrorCategory as ProtoApplicationErrorCategory, + enums::v1::{ + ApplicationErrorCategory as ProtoApplicationErrorCategory, + CancelExternalWorkflowExecutionFailedCause, SignalExternalWorkflowExecutionFailedCause, + }, failure::v1::{ ActivityFailureInfo, ApplicationFailureInfo, CanceledFailureInfo, ChildWorkflowExecutionFailureInfo, Failure, failure::FailureInfo, @@ -48,7 +54,35 @@ pub trait FailureConverter { } /// Default failure converter. -pub struct DefaultFailureConverter; +pub struct DefaultFailureConverter { + encode_common_attributes: bool, +} + +impl DefaultFailureConverter { + /// Creates a failure converter, optionally moving failure messages and stack traces into + /// encoded attributes. + pub const fn new(encode_common_attributes: bool) -> Self { + Self { + encode_common_attributes, + } + } +} + +impl Default for DefaultFailureConverter { + fn default() -> Self { + Self::new(false) + } +} + +/// Failure attributes that can be moved into an encoded payload. +#[derive(serde::Deserialize, serde::Serialize)] +#[non_exhaustive] +pub struct CommonAttributes { + /// Failure message. + pub message: String, + /// Failure stack trace. + pub stack_trace: String, +} /// Adapts a normalized incoming failure into a caller-facing error surface. pub trait FailureDecodeHint { @@ -59,13 +93,33 @@ pub trait FailureDecodeHint { fn adapt(self, normalized: IncomingError) -> Self::Output; } +/// No-op decode hint; returns the error unchanged. +#[derive(Debug, Clone, Copy)] +pub struct NoopDecodeHint; + +impl FailureDecodeHint for NoopDecodeHint { + type Output = IncomingError; + + fn adapt(self, normalized: IncomingError) -> Self::Output { + normalized + } +} + /// Decode hint for activity execution results. #[derive(Debug, Clone, Copy)] +#[non_exhaustive] pub struct ActivityExecutionDecodeHint { /// Whether the workflow-side resolution was cancelled rather than failed. pub cancelled: bool, } +impl ActivityExecutionDecodeHint { + /// Creates a decode hint for an activity resolution. + pub fn new(cancelled: bool) -> Self { + Self { cancelled } + } +} + impl FailureDecodeHint for ActivityExecutionDecodeHint { type Output = ActivityExecutionError; @@ -102,7 +156,8 @@ impl FailureDecodeHint for ActivityExecutionDecodeHint { } /// Decode hint for child-workflow start results. -#[derive(Debug, Clone, Copy)] +#[derive(Debug, Clone, Copy, Default)] +#[non_exhaustive] pub struct ChildWorkflowStartDecodeHint; impl FailureDecodeHint for ChildWorkflowStartDecodeHint { @@ -128,7 +183,8 @@ impl FailureDecodeHint for ChildWorkflowStartDecodeHint { } /// Decode hint for child-workflow execution results. -#[derive(Debug, Clone, Copy)] +#[derive(Debug, Clone, Copy, Default)] +#[non_exhaustive] pub struct ChildWorkflowExecutionDecodeHint; impl FailureDecodeHint for ChildWorkflowExecutionDecodeHint { @@ -149,17 +205,62 @@ impl FailureDecodeHint for ChildWorkflowExecutionDecodeHint { } /// Decode hint for workflow signal failures. -#[derive(Debug, Clone, Copy)] -pub struct WorkflowSignalDecodeHint; +#[derive(Debug, Clone, Copy, Default)] +#[non_exhaustive] +pub struct WorkflowSignalDecodeHint { + cause: SignalExternalWorkflowExecutionFailedCause, +} + +impl WorkflowSignalDecodeHint { + /// Creates a decode hint with the server-reported signal failure cause. + pub fn new(cause: SignalExternalWorkflowExecutionFailedCause) -> Self { + Self { cause } + } +} impl FailureDecodeHint for WorkflowSignalDecodeHint { type Output = WorkflowSignalError; fn adapt(self, normalized: IncomingError) -> Self::Output { let failure = normalized.failure().clone(); - WorkflowSignalError::Failed(Box::new(WorkflowSignalFailureError::new( - failure, normalized, - ))) + let error = Box::new(WorkflowSignalFailureError::new(failure, normalized)); + if self.cause + == SignalExternalWorkflowExecutionFailedCause::ExternalWorkflowExecutionNotFound + { + WorkflowSignalError::NotFound(error) + } else { + WorkflowSignalError::Failed(error) + } + } +} + +/// Decode hint for external-workflow cancellation failures. +#[derive(Debug, Clone, Copy, Default)] +#[non_exhaustive] +pub struct CancelExternalWorkflowDecodeHint { + cause: CancelExternalWorkflowExecutionFailedCause, +} + +impl CancelExternalWorkflowDecodeHint { + /// Creates a decode hint with the server-reported cancellation failure cause. + pub fn new(cause: CancelExternalWorkflowExecutionFailedCause) -> Self { + Self { cause } + } +} + +impl FailureDecodeHint for CancelExternalWorkflowDecodeHint { + type Output = CancelExternalWorkflowError; + + fn adapt(self, normalized: IncomingError) -> Self::Output { + let failure = normalized.failure().clone(); + let error = Box::new(WorkflowCancelFailureError::new(failure, normalized)); + if self.cause + == CancelExternalWorkflowExecutionFailedCause::ExternalWorkflowExecutionNotFound + { + CancelExternalWorkflowError::NotFound(error) + } else { + CancelExternalWorkflowError::Failed(error) + } } } @@ -193,13 +294,25 @@ impl FailureConverter for DefaultFailureConverter { OutgoingError::Workflow(OutgoingWorkflowError::WorkflowSignal(signal)) => { signal.encode_failure(payload_converter, context) } + OutgoingError::Workflow(OutgoingWorkflowError::CancelExternalWorkflow(cancel)) => { + cancel.encode_failure(payload_converter, context) + } }; - encoded.unwrap_or_else(|converter_error| { + let mut failure = encoded.unwrap_or_else(|converter_error| { Failure::application_failure( failed_error_conversion_message(&original_error, &converter_error), false, ) - }) + }); + if self.encode_common_attributes + && encode_common_attributes(&mut failure, payload_converter, context).is_err() + { + failure = Failure::application_failure( + "Failed encoding failure attributes".to_owned(), + false, + ); + } + failure } fn to_error( @@ -227,6 +340,7 @@ enum ClassifiedFailure<'a> { ChildWorkflowExecution(&'a ChildWorkflowExecutionError), ChildWorkflowStart(&'a ChildWorkflowStartError), WorkflowSignal(&'a WorkflowSignalError), + CancelExternalWorkflow(&'a CancelExternalWorkflowError), Generic(&'a (dyn std::error::Error + 'static)), } @@ -252,6 +366,8 @@ impl<'a> ClassifiedFailure<'a> { Self::ChildWorkflowStart(child) } else if let Some(child_signal) = err.downcast_ref::() { Self::WorkflowSignal(child_signal) + } else if let Some(cancel_external) = err.downcast_ref::() { + Self::CancelExternalWorkflow(cancel_external) } else { Self::Generic(err) } @@ -299,6 +415,14 @@ impl<'a> ClassifiedFailure<'a> { .unwrap_or_else(|converter_error| { encode_failed_error_conversion(signal, converter_error) }), + Self::CancelExternalWorkflow(cancel) => cancel + .encode_failure( + &PayloadConverter::default(), + &SerializationContextData::None, + ) + .unwrap_or_else(|converter_error| { + encode_failed_error_conversion(cancel, converter_error) + }), Self::Generic(err) => encode_generic_application_failure(err), } } @@ -400,7 +524,20 @@ impl EncodeFailure for WorkflowSignalError { _: &SerializationContextData, ) -> Result { Ok(match self { - Self::Failed(failure) => failure.failure().clone(), + Self::NotFound(failure) | Self::Failed(failure) => failure.failure().clone(), + Self::Serialization(err) => encode_generic_application_failure(err), + }) + } +} + +impl EncodeFailure for CancelExternalWorkflowError { + fn encode_failure( + &self, + _: &PayloadConverter, + _: &SerializationContextData, + ) -> Result { + Ok(match self { + Self::NotFound(error) | Self::Failed(error) => error.failure().clone(), Self::Serialization(err) => encode_generic_application_failure(err), }) } @@ -453,11 +590,39 @@ fn encode_failed_error_conversion( } } +fn encode_common_attributes( + failure: &mut Failure, + payload_converter: &PayloadConverter, + context: &SerializationContextData, +) -> Result<(), PayloadConversionError> { + if let Some(cause) = failure.cause.as_deref_mut() { + encode_common_attributes(cause, payload_converter, context)?; + } + failure.encoded_attributes = Some(payload_converter.to_payload( + &SerializationContext::new(context, payload_converter), + &CommonAttributes { + message: std::mem::take(&mut failure.message), + stack_trace: std::mem::take(&mut failure.stack_trace), + }, + )?); + failure.message = "Encoded failure".to_owned(); + Ok(()) +} + fn decode_failure( - failure: Failure, + mut failure: Failure, payload_converter: &PayloadConverter, context: &SerializationContextData, ) -> IncomingError { + if let Some(encoded_attributes) = failure.encoded_attributes.clone() + && let Ok(attributes) = payload_converter.from_payload::( + &SerializationContext::new(context, payload_converter), + encoded_attributes, + ) + { + failure.message = attributes.message; + failure.stack_trace = attributes.stack_trace; + } let cause = failure .cause .clone() @@ -506,8 +671,11 @@ fn decode_failure( mod tests { use super::*; use crate::{ - data_converters::{GenericPayloadConverter, SerializationContext}, - error::ApplicationErrorCategory, + data_converters::{ + ActivitySerializationContext, GenericPayloadConverter, SerializationContext, + WorkflowSerializationContext, + }, + error::{ApplicationErrorCategory, StartChildWorkflowExecutionFailedCause}, protos::temporal::api::{ common::v1::{Payload, Payloads}, failure::v1::{ @@ -601,17 +769,17 @@ mod tests { } fn convert(err: OutgoingWorkflowError) -> Failure { - DefaultFailureConverter.to_failure( + DefaultFailureConverter::default().to_failure( OutgoingError::Workflow(err), &PayloadConverter::default(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) } fn data_converter() -> crate::data_converters::DataConverter { crate::data_converters::DataConverter::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), crate::data_converters::DefaultPayloadCodec, ) } @@ -673,10 +841,10 @@ mod tests { let converter = PayloadConverter::default(); let details: String = converter .from_payloads( - &SerializationContext { - data: &SerializationContextData::Workflow, - converter: &converter, - }, + &SerializationContext::new( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + &converter, + ), payloads, ) .unwrap(); @@ -685,14 +853,14 @@ mod tests { #[test] fn application_failures_surface_detail_encoding_errors_with_original_message() { - let failure = DefaultFailureConverter.to_failure( + let failure = DefaultFailureConverter::default().to_failure( OutgoingError::Workflow(OutgoingWorkflowError::Application(Box::new( ApplicationFailure::builder(anyhow::anyhow!("app boom")) .details(AlwaysFailsSerialize) .build(), ))), &PayloadConverter::default(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ); assert_eq!( @@ -706,10 +874,10 @@ mod tests { let converter = PayloadConverter::default(); let payloads = converter .to_payloads( - &SerializationContext { - data: &SerializationContextData::Workflow, - converter: &converter, - }, + &SerializationContext::new( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + &converter, + ), &"detail", ) .unwrap(); @@ -724,8 +892,12 @@ mod tests { ..Default::default() }; - let decoded = DefaultFailureConverter - .to_error(failure, &converter, &SerializationContextData::Workflow) + let decoded = DefaultFailureConverter::default() + .to_error( + failure, + &converter, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + ) .unwrap(); let IncomingError::Application(app) = decoded else { @@ -868,11 +1040,11 @@ mod tests { )); assert!(cause.cause.is_none()); - let decoded = DefaultFailureConverter + let decoded = DefaultFailureConverter::default() .to_error( converted.clone(), &PayloadConverter::default(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .unwrap(); @@ -890,13 +1062,86 @@ mod tests { assert!(wrapper.cause().is_none()); } + #[test] + fn failure_converter_encodes_and_decodes_cause_chain() { + let payload_converter = PayloadConverter::default(); + let converter = DefaultFailureConverter::new(true); + let context = SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let failure = Failure { + message: "outer message".to_owned(), + stack_trace: "outer stack trace".to_owned(), + cause: Some(Box::new(Failure { + message: "inner message".to_owned(), + stack_trace: "inner stack trace".to_owned(), + failure_info: Some(FailureInfo::ApplicationFailureInfo( + ApplicationFailureInfo::default(), + )), + ..Default::default() + })), + failure_info: Some(FailureInfo::ActivityFailureInfo( + ActivityFailureInfo::default(), + )), + ..Default::default() + }; + let activity_error = ActivityExecutionError::Failed(ActivityFailureError::new( + failure, + ActivityFailureInfo::default(), + None, + )); + + let failure = converter.to_failure( + OutgoingError::Workflow(OutgoingWorkflowError::ActivityExecution(Box::new( + activity_error, + ))), + &payload_converter, + &context, + ); + + assert_eq!(failure.message, "Encoded failure"); + assert_eq!(failure.cause.as_ref().unwrap().message, "Encoded failure"); + assert!(failure.stack_trace.is_empty()); + assert!(failure.cause.as_ref().unwrap().stack_trace.is_empty()); + let payload_context = SerializationContext::new(&context, &payload_converter); + let outer_attributes: CommonAttributes = payload_converter + .from_payload( + &payload_context, + failure.encoded_attributes.clone().unwrap(), + ) + .unwrap(); + let inner_attributes: CommonAttributes = payload_converter + .from_payload( + &payload_context, + failure + .cause + .as_ref() + .unwrap() + .encoded_attributes + .clone() + .unwrap(), + ) + .unwrap(); + assert_eq!(outer_attributes.message, "outer message"); + assert_eq!(inner_attributes.message, "inner message"); + assert_eq!(outer_attributes.stack_trace, "outer stack trace"); + assert_eq!(inner_attributes.stack_trace, "inner stack trace"); + + let decoded = DefaultFailureConverter::default() + .to_error(failure, &payload_converter, &context) + .unwrap(); + assert_eq!(decoded.failure().message, "outer message"); + assert_eq!(decoded.failure().stack_trace, "outer stack trace"); + let cause = decoded.cause().unwrap().failure(); + assert_eq!(cause.message, "inner message"); + assert_eq!(cause.stack_trace, "inner stack trace"); + } + #[test] fn start_failed_child_workflow_errors_fall_back_to_application_failures() { let failure = convert(OutgoingWorkflowError::ChildWorkflowStart(Box::new( ChildWorkflowStartError::StartFailed { workflow_id: "wf-id".to_owned(), workflow_type: "wf-type".to_owned(), - cause: crate::protos::coresdk::child_workflow::StartChildWorkflowExecutionFailedCause::WorkflowAlreadyExists, + cause: StartChildWorkflowExecutionFailedCause::WorkflowAlreadyExists, }, ))); assert!(matches!( @@ -920,11 +1165,11 @@ mod tests { ..Default::default() }; - let decoded = DefaultFailureConverter + let decoded = DefaultFailureConverter::default() .to_error( failure.clone(), &PayloadConverter::default(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .unwrap(); @@ -953,11 +1198,11 @@ mod tests { ..Default::default() }; - let decoded = DefaultFailureConverter + let decoded = DefaultFailureConverter::default() .to_error( failure.clone(), &PayloadConverter::default(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .unwrap(); @@ -979,11 +1224,11 @@ mod tests { ..Default::default() }; - let decoded = DefaultFailureConverter + let decoded = DefaultFailureConverter::default() .to_error( failure.clone(), &PayloadConverter::default(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .unwrap(); @@ -997,11 +1242,11 @@ mod tests { assert_eq!(reencoded.message, failure.message); assert_eq!(reencoded.cause.as_deref(), failure.cause.as_deref()); - let decoded_reencoded = DefaultFailureConverter + let decoded_reencoded = DefaultFailureConverter::default() .to_error( reencoded, &PayloadConverter::default(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .unwrap(); let IncomingError::Application(roundtripped) = decoded_reencoded else { @@ -1027,11 +1272,11 @@ mod tests { ..Default::default() }; - let decoded = DefaultFailureConverter + let decoded = DefaultFailureConverter::default() .to_error( failure.clone(), &PayloadConverter::default(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .unwrap(); @@ -1098,11 +1343,11 @@ mod tests { ..Default::default() }; - let decoded = DefaultFailureConverter + let decoded = DefaultFailureConverter::default() .to_error( failure.clone(), &PayloadConverter::default(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .unwrap(); @@ -1135,13 +1380,13 @@ mod tests { }; let data_converter = crate::data_converters::DataConverter::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), crate::data_converters::DefaultPayloadCodec, ); let decoded = data_converter .to_error( - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), failure.clone(), ActivityExecutionDecodeHint { cancelled: false }, ) @@ -1192,7 +1437,7 @@ mod tests { ) { let decoded = data_converter() .to_error( - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), failure.clone(), ActivityExecutionDecodeHint { cancelled: true }, ) @@ -1237,11 +1482,11 @@ mod tests { ..Default::default() }; - let decoded = DefaultFailureConverter + let decoded = DefaultFailureConverter::default() .to_error( failure.clone(), &PayloadConverter::default(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .unwrap(); @@ -1273,11 +1518,11 @@ mod tests { ..Default::default() }; - let decoded = DefaultFailureConverter + let decoded = DefaultFailureConverter::default() .to_error( failure.clone(), &PayloadConverter::default(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .unwrap(); @@ -1314,7 +1559,7 @@ mod tests { }; let decoded = data_converter() .to_error( - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), failure.clone(), ChildWorkflowExecutionDecodeHint, ) @@ -1363,7 +1608,7 @@ mod tests { ) { let decoded = data_converter() .to_error( - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), failure.clone(), ChildWorkflowExecutionDecodeHint, ) @@ -1395,7 +1640,7 @@ mod tests { }; let decoded = data_converter() .to_error( - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), failure.clone(), ChildWorkflowStartDecodeHint, ) @@ -1408,6 +1653,58 @@ mod tests { assert!(decoded_failure.cause().is_none()); } + #[test] + fn workflow_signal_decode_hint_recognizes_not_found() { + let failure = Failure { + message: "workflow not found".to_owned(), + ..Default::default() + }; + let decoded = data_converter() + .to_error( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + failure.clone(), + WorkflowSignalDecodeHint::new( + SignalExternalWorkflowExecutionFailedCause::ExternalWorkflowExecutionNotFound, + ), + ) + .unwrap(); + + let WorkflowSignalError::NotFound(decoded_failure) = decoded else { + panic!("expected not-found workflow signal error"); + }; + assert_eq!(decoded_failure.failure(), &failure); + } + + #[test] + fn cancel_external_workflow_decode_hint_recognizes_not_found() { + let failure = Failure { + message: "workflow not found".to_owned(), + cause: Some(Box::new(Failure { + message: "timed out".to_owned(), + failure_info: Some(FailureInfo::TimeoutFailureInfo( + TimeoutFailureInfo::default(), + )), + ..Default::default() + })), + ..Default::default() + }; + let decoded = data_converter() + .to_error( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + failure.clone(), + CancelExternalWorkflowDecodeHint::new( + CancelExternalWorkflowExecutionFailedCause::ExternalWorkflowExecutionNotFound, + ), + ) + .unwrap(); + + let CancelExternalWorkflowError::NotFound(decoded_failure) = decoded else { + panic!("expected not-found external-workflow cancellation error"); + }; + assert_eq!(decoded_failure.failure(), &failure); + assert!(std::error::Error::source(&*decoded_failure).is_some()); + } + #[test] fn child_workflow_signal_decode_hint_preserves_failure_proto() { let failure = Failure { @@ -1423,9 +1720,9 @@ mod tests { }; let decoded = data_converter() .to_error( - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), failure.clone(), - WorkflowSignalDecodeHint, + WorkflowSignalDecodeHint::default(), ) .unwrap(); @@ -1446,10 +1743,10 @@ mod tests { #[test] fn outgoing_cancelled_activity_errors_encode_to_cancelled_failures() { - let failure = DefaultFailureConverter.to_failure( + let failure = DefaultFailureConverter::default().to_failure( OutgoingError::Activity(OutgoingActivityError::Cancelled { details: None }), &PayloadConverter::default(), - &SerializationContextData::Activity, + &SerializationContextData::Activity(ActivitySerializationContext::new()), ); assert_eq!(failure.message, "Activity cancelled"); @@ -1461,19 +1758,19 @@ mod tests { #[test] fn outgoing_cancelled_activity_errors_encode_serializable_details_with_payload_converter() { - let failure = DefaultFailureConverter.to_failure( + let failure = DefaultFailureConverter::default().to_failure( OutgoingError::Activity(OutgoingActivityError::Cancelled { details: Some("detail".to_string().into()), }), &PayloadConverter::default(), - &SerializationContextData::Activity, + &SerializationContextData::Activity(ActivitySerializationContext::new()), ); - let err = DefaultFailureConverter + let err = DefaultFailureConverter::default() .to_error( failure, &PayloadConverter::default(), - &SerializationContextData::Activity, + &SerializationContextData::Activity(ActivitySerializationContext::new()), ) .unwrap(); let cancelled = err.as_cancelled().unwrap(); diff --git a/crates/common-wasm/src/data_converters/well_known.rs b/crates/common-wasm/src/data_converters/well_known.rs new file mode 100644 index 000000000..9420d3c92 --- /dev/null +++ b/crates/common-wasm/src/data_converters/well_known.rs @@ -0,0 +1,123 @@ +use crate::protos::{ENCODING_PAYLOAD_KEY, temporal::api::common::v1::Payload}; +use std::{ + any::{Any, TypeId}, + collections::HashMap, +}; + +pub(super) const BINARY_PLAIN_ENCODING_VAL: &str = "binary/plain"; +pub(super) const BINARY_NULL_ENCODING_VAL: &str = "binary/null"; + +#[derive(Clone, Copy)] +pub(super) enum WellKnownType { + Unit, + Bytes, + OptionalBytes, +} + +impl WellKnownType { + pub(super) fn of() -> Option { + let type_id = TypeId::of::(); + if type_id == TypeId::of::<()>() { + Some(Self::Unit) + } else if type_id == TypeId::of::>() { + Some(Self::Bytes) + } else if type_id == TypeId::of::>>() { + Some(Self::OptionalBytes) + } else { + None + } + } + + pub(super) fn to_payload(self, value: &T) -> Payload { + let value = value as &dyn Any; + match self { + Self::Unit => binary_null_payload(), + Self::Bytes => binary_plain_payload(value.downcast_ref::>().unwrap().clone()), + Self::OptionalBytes => match value.downcast_ref::>>().unwrap() { + Some(bytes) => binary_plain_payload(bytes.clone()), + None => binary_null_payload(), + }, + } + } + + pub(super) fn to_payloads(self, value: &T) -> Vec { + match self { + Self::Unit => Vec::new(), + _ => vec![self.to_payload(value)], + } + } + + pub(super) fn try_from_payload(self, payload: Payload) -> Result { + match self { + Self::Unit if is_binary_null_payload(&payload) => Ok(downcast_well_known(())), + Self::Bytes if is_binary_plain_payload(&payload) => { + Ok(downcast_well_known(payload.data)) + } + Self::OptionalBytes if is_binary_plain_payload(&payload) => { + Ok(downcast_well_known(Some(payload.data))) + } + Self::OptionalBytes if is_binary_null_payload(&payload) => { + Ok(downcast_well_known(None::>)) + } + _ => Err(payload), + } + } + + pub(super) fn try_from_payloads( + self, + mut payloads: Vec, + ) -> Result> { + if matches!(self, Self::Unit) && payloads.is_empty() { + return Ok(downcast_well_known(())); + } + if payloads.len() != 1 { + return Err(payloads); + } + match self.try_from_payload(payloads.pop().unwrap()) { + Ok(value) => Ok(value), + Err(payload) => Err(vec![payload]), + } + } +} + +fn downcast_well_known(value: impl Any) -> T { + let value: Box = Box::new(value); + *value.downcast::().ok().unwrap() +} + +fn binary_plain_payload(data: Vec) -> Payload { + Payload { + metadata: HashMap::from([( + ENCODING_PAYLOAD_KEY.to_string(), + BINARY_PLAIN_ENCODING_VAL.as_bytes().to_vec(), + )]), + data, + external_payloads: vec![], + } +} + +fn is_binary_plain_payload(payload: &Payload) -> bool { + payload + .metadata + .get(ENCODING_PAYLOAD_KEY) + .is_some_and(|encoding| encoding == BINARY_PLAIN_ENCODING_VAL.as_bytes()) +} + +pub(super) fn binary_null_payload() -> Payload { + Payload { + metadata: HashMap::from([( + ENCODING_PAYLOAD_KEY.to_string(), + BINARY_NULL_ENCODING_VAL.as_bytes().to_vec(), + )]), + data: vec![], + external_payloads: vec![], + } +} + +pub(super) fn is_binary_null_payload(payload: &Payload) -> bool { + payload.data.is_empty() + && payload + .metadata + .get(ENCODING_PAYLOAD_KEY) + .is_some_and(|encoding| encoding == BINARY_NULL_ENCODING_VAL.as_bytes()) +} diff --git a/crates/common-wasm/src/error.rs b/crates/common-wasm/src/error.rs index fe47a923e..2155c3bef 100644 --- a/crates/common-wasm/src/error.rs +++ b/crates/common-wasm/src/error.rs @@ -7,20 +7,29 @@ use crate::{ RawValue, SerializationContext, SerializationContextData, TemporalDeserializable, TemporalSerializable, }, - protos::{ - coresdk::child_workflow::StartChildWorkflowExecutionFailedCause, - temporal::api::{ - common::v1::{Payload, Payloads}, - enums::v1::{ - ApplicationErrorCategory as ProtoApplicationErrorCategory, - RetryState as ProtoRetryState, TimeoutType as ProtoTimeoutType, - }, - failure::v1::Failure, + protos::temporal::api::{ + common::v1::{Payload, Payloads}, + enums::v1::{ + ApplicationErrorCategory as ProtoApplicationErrorCategory, + RetryState as ProtoRetryState, TimeoutType as ProtoTimeoutType, }, + failure::v1::Failure, }, }; use std::time::Duration; +/// Why starting a child workflow failed. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +#[non_exhaustive] +pub enum StartChildWorkflowExecutionFailedCause { + /// No cause was specified. + Unspecified, + /// A workflow with the requested ID already exists. + WorkflowAlreadyExists, + /// A cause introduced by a newer server or API version. + Unknown, +} + /// Describes why a retry did or did not occur. #[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] #[non_exhaustive] @@ -123,13 +132,7 @@ where payload_converter: &PayloadConverter, context: &SerializationContextData, ) -> Result, PayloadConversionError> { - payload_converter.to_payloads( - &SerializationContext { - data: context, - converter: payload_converter, - }, - self, - ) + payload_converter.to_payloads(&SerializationContext::new(context, payload_converter), self) } } @@ -370,7 +373,7 @@ impl ApplicationFailure { FailurePayloads::from(DecodablePayloads::new( details.payloads, payload_converter.clone(), - *context, + context.clone(), )) }), failure: Some(failure), @@ -455,6 +458,9 @@ pub enum OutgoingWorkflowError { /// A workflow failure sourced from signaling a workflow. #[error(transparent)] WorkflowSignal(#[from] Box), + /// A workflow failure sourced from requesting cancellation of an external workflow. + #[error(transparent)] + CancelExternalWorkflow(#[from] Box), } impl OutgoingWorkflowError { @@ -468,6 +474,7 @@ impl OutgoingWorkflowError { Self::ChildWorkflowExecution(err) => err.as_cancelled(), Self::ChildWorkflowStart(err) => err.as_cancelled(), Self::WorkflowSignal(err) => err.as_cancelled(), + Self::CancelExternalWorkflow(err) => err.as_cancelled(), } } } @@ -520,8 +527,18 @@ impl From for OutgoingWorkflowError { } } +impl From for OutgoingWorkflowError { + fn from(value: CancelExternalWorkflowError) -> Self { + match value { + CancelExternalWorkflowError::Serialization(err) => Self::PayloadConversion(err), + other => Self::CancelExternalWorkflow(Box::new(other)), + } + } +} + /// A normalized incoming Temporal failure decoded from a protobuf [`Failure`]. #[derive(Debug)] +#[non_exhaustive] pub enum IncomingError { /// A decoded application failure. Application(ApplicationFailure), @@ -735,7 +752,7 @@ impl TimeoutError { cause: cause.map(Box::new), timeout_type: TimeoutType::from_raw(failure_info.timeout_type), last_heartbeat_details: failure_info.last_heartbeat_details.map(|details| { - DecodablePayloads::new(details.payloads, payload_converter.clone(), *context) + DecodablePayloads::new(details.payloads, payload_converter.clone(), context.clone()) }), } } @@ -786,7 +803,7 @@ impl CancelledError { failure, cause: cause.map(Box::new), details: failure_info.details.map(|details| { - DecodablePayloads::new(details.payloads, payload_converter.clone(), *context) + DecodablePayloads::new(details.payloads, payload_converter.clone(), context.clone()) }), } } @@ -982,6 +999,7 @@ incoming_failure_wrapper!( /// Error type for activity execution outcomes. #[derive(Debug, thiserror::Error)] +#[non_exhaustive] pub enum ActivityExecutionError { /// The activity failed with the given failure details. #[error("Activity failed: {}", .0.failure().message)] @@ -1043,6 +1061,7 @@ impl ActivityExecutionError { /// Error returned when starting a child workflow fails. #[derive(Debug, thiserror::Error)] +#[non_exhaustive] pub enum ChildWorkflowStartError { /// The child workflow start was cancelled before the normal execution wrapper path existed. #[error("Child workflow start cancelled: {}", .0.failure().message)] @@ -1096,6 +1115,7 @@ impl ChildWorkflowStartError { /// Error returned when a child workflow execution fails. #[derive(Debug, thiserror::Error)] +#[non_exhaustive] pub enum ChildWorkflowExecutionError { /// The child workflow failed. #[error("Child workflow failed: {}", .0.failure().message)] @@ -1151,7 +1171,11 @@ impl ChildWorkflowExecutionError { /// Error returned when signaling a workflow fails. #[derive(Debug, thiserror::Error)] +#[non_exhaustive] pub enum WorkflowSignalError { + /// The target workflow was not found. + #[error("Workflow not found: {}", .0.failure().message)] + NotFound(#[source] Box), /// The signal delivery failed. #[error("Child workflow signal failed: {}", .0.failure().message)] Failed(#[source] Box), @@ -1160,11 +1184,58 @@ pub enum WorkflowSignalError { Serialization(#[from] PayloadConversionError), } +/// Error returned when requesting cancellation of an external workflow fails. +#[derive(Debug, thiserror::Error)] +pub enum CancelExternalWorkflowError { + /// The target workflow was not found. + #[error("Workflow not found: {}", .0.failure().message)] + NotFound(#[source] Box), + /// The cancellation request failed. + #[error("External workflow cancellation request failed: {}", .0.failure().message)] + Failed(#[source] Box), + /// Failed to deserialize payloads attached to the cancellation failure. + #[error("External workflow cancellation failure conversion failed: {0}")] + Serialization(#[from] PayloadConversionError), +} + +impl CancelExternalWorkflowError { + /// Returns the retained top-level cancellation failure proto, if one exists. + pub fn failure(&self) -> Option<&Failure> { + match self { + Self::NotFound(err) | Self::Failed(err) => Some(err.failure()), + Self::Serialization(_) => None, + } + } + + /// Returns the normalized cause of the cancellation failure, if any. + pub fn cause(&self) -> Option<&IncomingError> { + match self { + Self::NotFound(err) | Self::Failed(err) => err.cause(), + Self::Serialization(_) => None, + } + } + + /// Returns the normalized cancellation failure itself, if one exists. + pub fn reason(&self) -> Option<&IncomingError> { + match self { + Self::NotFound(err) | Self::Failed(err) => Some(err.error()), + Self::Serialization(_) => None, + } + } + + /// If this error was caused by cancellation, returns the associated [`CancelledError`]. + pub fn as_cancelled(&self) -> Option<&CancelledError> { + self.reason()?.as_cancelled() + } +} + impl WorkflowSignalError { /// Returns the retained top-level workflow signal failure proto, if one exists. pub fn failure(&self) -> Option<&Failure> { match self { - WorkflowSignalError::Failed(err) => Some(err.failure()), + WorkflowSignalError::NotFound(err) | WorkflowSignalError::Failed(err) => { + Some(err.failure()) + } WorkflowSignalError::Serialization(_) => None, } } @@ -1172,7 +1243,7 @@ impl WorkflowSignalError { /// Returns the normalized cause of the workflow signal failure, if any. pub fn cause(&self) -> Option<&IncomingError> { match self { - WorkflowSignalError::Failed(err) => err.cause(), + WorkflowSignalError::NotFound(err) | WorkflowSignalError::Failed(err) => err.cause(), WorkflowSignalError::Serialization(_) => None, } } @@ -1180,7 +1251,9 @@ impl WorkflowSignalError { /// Returns the underlying failure reason for wrapper-shaped signal failures. pub fn reason(&self) -> Option<&IncomingError> { match self { - WorkflowSignalError::Failed(err) => Some(err.error()), + WorkflowSignalError::NotFound(err) | WorkflowSignalError::Failed(err) => { + Some(err.error()) + } WorkflowSignalError::Serialization(_) => None, } } @@ -1237,13 +1310,58 @@ impl std::error::Error for WorkflowSignalFailureError { } } +/// A normalized external workflow cancellation failure wrapper. +#[derive(Debug)] +pub struct WorkflowCancelFailureError { + failure: Failure, + error: Box, +} + +impl WorkflowCancelFailureError { + /// Creates an external workflow cancellation failure wrapper. + pub(crate) fn new(failure: Failure, error: IncomingError) -> Self { + Self { + failure, + error: Box::new(error), + } + } + + /// Returns the retained top-level proto failure. + pub fn failure(&self) -> &Failure { + &self.failure + } + + /// Returns the normalized direct cause of the external workflow cancellation failure, if any. + pub fn cause(&self) -> Option<&IncomingError> { + self.error.cause() + } + + /// Returns the direct decoded incoming error represented by the top-level proto failure. + pub fn error(&self) -> &IncomingError { + &self.error + } +} + +impl std::fmt::Display for WorkflowCancelFailureError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + self.failure.fmt(f) + } +} + +impl std::error::Error for WorkflowCancelFailureError { + fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { + self.cause() + .map(|cause| cause as &(dyn std::error::Error + 'static)) + } +} + #[cfg(test)] mod tests { use super::*; use crate::{ data_converters::{ DefaultFailureConverter, FailureConverter, GenericPayloadConverter, PayloadConverter, - SerializationContext, SerializationContextData, + SerializationContext, SerializationContextData, WorkflowSerializationContext, }, protos::temporal::api::{ common::v1::Payload, @@ -1276,11 +1394,11 @@ mod tests { ..Default::default() }; - let decoded = DefaultFailureConverter + let decoded = DefaultFailureConverter::default() .to_error( failure, &PayloadConverter::default(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .unwrap(); let IncomingError::Activity(activity) = decoded else { @@ -1321,7 +1439,7 @@ mod tests { ..Default::default() }], }; - let failure = DefaultFailureConverter.to_failure( + let failure = DefaultFailureConverter::default().to_failure( OutgoingError::Workflow(OutgoingWorkflowError::Application(Box::new( ApplicationFailure::builder(anyhow::anyhow!("oops")) .type_name("MyType".to_owned()) @@ -1351,7 +1469,7 @@ mod tests { data: b"details".to_vec(), ..Default::default() }; - let failure = DefaultFailureConverter.to_failure( + let failure = DefaultFailureConverter::default().to_failure( OutgoingError::Workflow(OutgoingWorkflowError::Application(Box::new( ApplicationFailure::builder(anyhow::anyhow!("oops")) .details(RawValue::new(vec![payload.clone()])) @@ -1369,7 +1487,7 @@ mod tests { #[test] fn builder_accepts_serializable_details() { - let failure = DefaultFailureConverter.to_failure( + let failure = DefaultFailureConverter::default().to_failure( OutgoingError::Workflow(OutgoingWorkflowError::Application(Box::new( ApplicationFailure::builder(anyhow::anyhow!("oops")) .details("details".to_string()) @@ -1386,10 +1504,7 @@ mod tests { let converter = PayloadConverter::default(); let details: String = converter .from_payloads( - &SerializationContext { - data: &SerializationContextData::None, - converter: &converter, - }, + &SerializationContext::new(&SerializationContextData::None, &converter), payloads, ) .unwrap(); @@ -1398,7 +1513,7 @@ mod tests { #[test] fn application_failure_encoding_surfaces_detail_encoding_errors() { - let failure = DefaultFailureConverter.to_failure( + let failure = DefaultFailureConverter::default().to_failure( OutgoingError::Workflow(OutgoingWorkflowError::Application(Box::new( ApplicationFailure::builder(anyhow::anyhow!("oops")) .details(AlwaysFailsSerialize) diff --git a/crates/common-wasm/src/lib.rs b/crates/common-wasm/src/lib.rs index fb3d7b5b7..505b0d3b4 100644 --- a/crates/common-wasm/src/lib.rs +++ b/crates/common-wasm/src/lib.rs @@ -7,6 +7,8 @@ #[macro_use] extern crate tracing; +use std::time::Duration; + mod activity_definition; pub mod data_converters; pub mod error; @@ -16,7 +18,8 @@ mod retry_policy; mod workflow_execution; pub mod protos { //! Protobuf definitions re-exported from `temporalio-protos`. - + //! + //! Because this module re-exports generated types, updating it might include breaking changes. pub use temporalio_protos::*; } pub mod search_attributes; @@ -24,7 +27,7 @@ pub mod worker; mod workflow_definition; pub use activity_definition::{ActivityDefinition, ActivityError, UntypedActivity}; -pub use memo::Memo; +pub use memo::{Memo, MemoValue, MemoValues}; pub use priority::Priority; pub use retry_policy::RetryPolicy; pub use search_attributes::{ @@ -48,3 +51,53 @@ macro_rules! dbg_panic { } #[allow(unused_imports)] pub(crate) use dbg_panic; + +/// Represents Activity schedule-to-close and start-to-close timeouts for the purposes of specifying +/// Activity options. Specifying at least one of them is required, but specifying both is also +/// allowed. Note that this type does not cover all available timeout options for an Activity. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[non_exhaustive] +pub enum ActivityCloseTimeouts { + /// Total time the Activity is allowed to run, including retries. + ScheduleToClose(Duration), + /// Maximum time of a single Activity execution attempt. Note that the Temporal Server doesn't + /// detect Worker process failures directly. It relies on this timeout to detect that an + /// Activity that didn't complete on time. So this timeout should be as short as the longest + /// possible execution of the Activity body. Potentially long running Activities must specify + /// `heartbeat_timeout` in options and heartbeat from the activity periodically for timely + /// failure detection. + StartToClose(Duration), + /// Applies both execution-attempt and overall-completion bounds. + ScheduleAndStartToClose { + /// Total time the Activity is allowed to run, including retries. + schedule_to_close: Duration, + /// Maximum time of a single Activity execution attempt. + start_to_close: Duration, + }, +} + +impl ActivityCloseTimeouts { + /// Returns value of [`Self::ScheduleToClose`] or + /// [`Self::ScheduleAndStartToClose::schedule_to_close`]. + pub fn schedule_to_close(&self) -> Option { + match self { + ActivityCloseTimeouts::ScheduleToClose(schedule_to_close) + | ActivityCloseTimeouts::ScheduleAndStartToClose { + schedule_to_close, .. + } => Some(*schedule_to_close), + _ => None, + } + } + + /// Returns value of [`Self::StartToClose`] or + /// [`Self::ScheduleAndStartToClose::start_to_close`]. + pub fn start_to_close(&self) -> Option { + match self { + ActivityCloseTimeouts::StartToClose(start_to_close) + | ActivityCloseTimeouts::ScheduleAndStartToClose { start_to_close, .. } => { + Some(*start_to_close) + } + _ => None, + } + } +} diff --git a/crates/common-wasm/src/memo.rs b/crates/common-wasm/src/memo.rs index 58f3cd8ae..ec1f7cbf4 100644 --- a/crates/common-wasm/src/memo.rs +++ b/crates/common-wasm/src/memo.rs @@ -1,16 +1,18 @@ use crate::{ data_converters::{ GenericPayloadConverter, PayloadConversionError, PayloadConverter, SerializationContext, - SerializationContextData, TemporalDeserializable, + SerializationContextData, TemporalDeserializable, TemporalSerializable, }, protos::temporal::api::common::v1::{Memo as ProtoMemo, Payload}, }; +use std::{collections::BTreeMap, sync::Arc}; /// A collection of memo payloads that can be deserialized into typed values. #[derive(Clone, Debug)] #[non_exhaustive] pub struct Memo { raw: ProtoMemo, + ordered_keys: BTreeMap, payload_converter: PayloadConverter, context: SerializationContextData, } @@ -18,14 +20,16 @@ pub struct Memo { impl Memo { /// Construct a memo with the payload converter and serialization context associated with its /// source. - #[doc(hidden)] pub fn from_raw( raw: Option, payload_converter: PayloadConverter, context: SerializationContextData, ) -> Self { + let raw = raw.unwrap_or_default(); + let ordered_keys = raw.fields.keys().cloned().map(|key| (key, ())).collect(); Self { - raw: raw.unwrap_or_default(), + raw, + ordered_keys, payload_converter, context, } @@ -41,10 +45,7 @@ impl Memo { }; self.payload_converter .from_payload( - &SerializationContext { - data: &self.context, - converter: &self.payload_converter, - }, + &SerializationContext::new(&self.context, &self.payload_converter), payload.clone(), ) .map(Some) @@ -65,9 +66,9 @@ impl Memo { self.raw.fields.is_empty() } - /// Iterates over memo keys. + /// Iterates over memo keys in lexicographic order. pub fn keys(&self) -> impl Iterator { - self.raw.fields.keys().map(String::as_str) + self.ordered_keys.keys().map(String::as_str) } /// Returns the underlying payload without applying payload conversion. @@ -86,18 +87,95 @@ impl Memo { } } +trait SerializableMemoValue: Send + Sync { + fn to_payload( + &self, + context: &SerializationContext<'_>, + ) -> Result; +} + +impl SerializableMemoValue for T +where + T: TemporalSerializable + Send + Sync + 'static, +{ + fn to_payload( + &self, + context: &SerializationContext<'_>, + ) -> Result { + context.converter.to_payload(context, self) + } +} + +/// A typed value used in a workflow memo update. +#[derive(Clone, derive_more::Debug)] +#[non_exhaustive] +pub struct MemoValue { + #[debug(skip)] + value: Arc, +} + +impl MemoValue { + /// Create a memo value that will be serialized with the workflow's data converter. + pub fn new(value: T) -> Self { + Self { + value: Arc::new(value), + } + } +} + +impl TemporalSerializable for MemoValue { + fn to_payload( + &self, + context: &SerializationContext<'_>, + ) -> Result { + self.value.to_payload(context) + } +} + +/// A complete set of memo values for a new workflow execution. +#[derive(Clone, Debug, Default)] +#[non_exhaustive] +pub struct MemoValues { + values: BTreeMap, +} + +impl MemoValues { + /// Create an empty set of memo values. + pub fn new() -> Self { + Self::default() + } + + /// Add or replace a memo value. + pub fn insert(&mut self, key: impl Into, value: T) -> &mut Self + where + T: TemporalSerializable + Send + Sync + 'static, + { + self.values.insert(key.into(), MemoValue::new(value)); + self + } + + /// Returns the value for `key`, if present. + pub fn get(&self, key: &str) -> Option<&MemoValue> { + self.values.get(key) + } + + /// Iterates over the memo entries in key order. + pub fn iter(&self) -> impl Iterator { + self.values.iter().map(|(key, value)| (key.as_str(), value)) + } +} + #[cfg(test)] mod tests { use super::*; + use crate::data_converters::WorkflowSerializationContext; use std::collections::HashMap; #[test] fn memo_decodes_serialized_values() { let payload_converter = PayloadConverter::default(); - let context = SerializationContext { - data: &SerializationContextData::Workflow, - converter: &payload_converter, - }; + let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let context = SerializationContext::new(&context_data, &payload_converter); let payload = payload_converter.to_payload(&context, &7_u32).unwrap(); let raw = ProtoMemo { fields: HashMap::from([("count".to_owned(), payload.clone())]), @@ -105,7 +183,7 @@ mod tests { let memo = Memo::from_raw( Some(raw.clone()), payload_converter, - SerializationContextData::Workflow, + SerializationContextData::Workflow(WorkflowSerializationContext::new()), ); assert_eq!(memo.get::("count").unwrap(), Some(7)); @@ -117,19 +195,70 @@ mod tests { #[test] fn memo_reports_deserialization_errors() { let payload_converter = PayloadConverter::default(); - let context = SerializationContext { - data: &SerializationContextData::Workflow, - converter: &payload_converter, - }; + let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let context = SerializationContext::new(&context_data, &payload_converter); let payload = payload_converter.to_payload(&context, &7_u32).unwrap(); let memo = Memo::from_raw( Some(ProtoMemo { fields: HashMap::from([("count".to_owned(), payload)]), }), payload_converter, - SerializationContextData::Workflow, + SerializationContextData::Workflow(WorkflowSerializationContext::new()), ); assert!(memo.get::("count").is_err()); } + + #[test] + fn memo_keys_have_replay_stable_order() { + let memo = Memo::from_raw( + Some(ProtoMemo { + fields: HashMap::from([ + ("zebra".to_owned(), Payload::default()), + ("alpha".to_owned(), Payload::default()), + ("middle".to_owned(), Payload::default()), + ]), + }), + PayloadConverter::default(), + SerializationContextData::Workflow(WorkflowSerializationContext::new()), + ); + + assert_eq!( + memo.keys().collect::>(), + vec!["alpha", "middle", "zebra"] + ); + } + + #[test] + fn memo_values_serialize_heterogeneous_values() { + let payload_converter = PayloadConverter::default(); + let mut values = MemoValues::new(); + values + .insert("count", 7_u32) + .insert("label", "hello".to_string()); + + let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let context = SerializationContext::new(&context_data, &payload_converter); + let fields = values + .iter() + .map(|(key, value)| { + ( + key.to_owned(), + payload_converter.to_payload(&context, value).unwrap(), + ) + }) + .collect(); + + let memo = Memo::from_raw( + Some(ProtoMemo { fields }), + payload_converter.clone(), + SerializationContextData::Workflow(WorkflowSerializationContext::new()), + ); + + assert_eq!(memo.get::("count").unwrap(), Some(7)); + assert_eq!( + memo.get::("label").unwrap(), + Some("hello".to_string()) + ); + } } diff --git a/crates/common-wasm/src/priority.rs b/crates/common-wasm/src/priority.rs index 85a382506..c5da0271d 100644 --- a/crates/common-wasm/src/priority.rs +++ b/crates/common-wasm/src/priority.rs @@ -19,7 +19,9 @@ use crate::protos::temporal::api::common; /// The overall semantics of Priority are: /// (more will be added here later) /// 1. First, consider "priority_key": lower number goes first. -#[derive(Debug, Clone, Default, PartialEq)] +#[derive(Debug, Clone, Default, PartialEq, bon::Builder)] +#[builder(on(String, into), state_mod(vis = "pub"))] +#[non_exhaustive] pub struct Priority { /// Priority key is a positive integer from 1 to n, where smaller integers /// correspond to higher priorities (tasks run sooner). In general, tasks in diff --git a/crates/common-wasm/src/search_attributes.rs b/crates/common-wasm/src/search_attributes.rs index 78f1b0388..8d2a3adf2 100644 --- a/crates/common-wasm/src/search_attributes.rs +++ b/crates/common-wasm/src/search_attributes.rs @@ -17,7 +17,10 @@ //! let unset = MY_KW.value_unset(); //! ``` -use std::{collections::HashMap, marker::PhantomData}; +use std::{ + collections::{BTreeMap, HashMap}, + marker::PhantomData, +}; use tracing::warn; @@ -248,10 +251,7 @@ fn encode_json_search_attr( indexed_value_type: IndexedValueType, ) -> Result { let converter = PayloadConverter::serde_json(); - let context = SerializationContext { - data: &SerializationContextData::None, - converter: &converter, - }; + let context = SerializationContext::new(&SerializationContextData::None, &converter); let mut payload = converter.to_payload(&context, value)?; payload.metadata.insert( TYPE_METADATA_KEY.to_string(), @@ -267,10 +267,7 @@ fn decode_json_search_attr( payload: &Payload, ) -> Result { let converter = PayloadConverter::serde_json(); - let context = SerializationContext { - data: &SerializationContextData::None, - converter: &converter, - }; + let context = SerializationContext::new(&SerializationContextData::None, &converter); Ok(converter.from_payload(&context, payload.clone())?) } @@ -571,7 +568,7 @@ impl SearchAttributeUpdate { /// [`SearchAttributeKey`]. #[derive(Debug, Clone, Default, PartialEq)] pub struct SearchAttributes { - fields: HashMap, + fields: BTreeMap, } impl SearchAttributes { @@ -579,7 +576,7 @@ impl SearchAttributes { /// /// Updates with `None` payloads remove any existing entry for that key. pub fn new(updates: impl IntoIterator) -> Self { - let mut fields = HashMap::new(); + let mut fields = BTreeMap::new(); for update in updates { match update.payload { Some(payload) => { @@ -653,7 +650,7 @@ impl SearchAttributes { self.fields.len() } - /// Returns an iterator over the attribute names in this collection. + /// Returns an iterator over the attribute names in lexicographic order. pub fn keys(&self) -> impl Iterator { self.fields.keys().map(|s| s.as_str()) } @@ -668,7 +665,7 @@ impl SearchAttributes { /// Convert to the proto wire representation. pub fn to_proto(&self) -> ProtoSearchAttributes { ProtoSearchAttributes { - indexed_fields: self.fields.clone(), + indexed_fields: self.fields.clone().into_iter().collect(), } } @@ -676,14 +673,14 @@ impl SearchAttributes { /// cloning. pub fn into_proto(self) -> ProtoSearchAttributes { ProtoSearchAttributes { - indexed_fields: self.fields, + indexed_fields: self.fields.into_iter().collect(), } } /// Construct from the proto wire representation by cloning the inner map. pub fn from_proto(attrs: &ProtoSearchAttributes) -> Self { Self { - fields: attrs.indexed_fields.clone(), + fields: attrs.indexed_fields.clone().into_iter().collect(), } } } @@ -692,7 +689,7 @@ impl From for SearchAttributes { /// Construct from an owned proto, moving the inner map without cloning. fn from(attrs: ProtoSearchAttributes) -> Self { Self { - fields: attrs.indexed_fields, + fields: attrs.indexed_fields.into_iter().collect(), } } } @@ -1126,11 +1123,19 @@ mod tests { } #[test] - fn keys_returns_attribute_names() { - let attrs = SearchAttributes::new([BOOL_KEY.value_set(true), INT_KEY.value_set(42)]); - let mut keys: Vec<&str> = attrs.keys().collect(); - keys.sort(); - assert_eq!(keys, vec!["my_bool", "my_int"]); + fn search_attribute_keys_have_replay_stable_order() { + let attrs = SearchAttributes::from(ProtoSearchAttributes { + indexed_fields: HashMap::from([ + ("zebra".to_owned(), Payload::default()), + ("alpha".to_owned(), Payload::default()), + ("middle".to_owned(), Payload::default()), + ]), + }); + + assert_eq!( + attrs.keys().collect::>(), + vec!["alpha", "middle", "zebra"] + ); } #[test] diff --git a/crates/common-wasm/src/worker.rs b/crates/common-wasm/src/worker.rs index bc133e1d4..c02f79088 100644 --- a/crates/common-wasm/src/worker.rs +++ b/crates/common-wasm/src/worker.rs @@ -4,7 +4,9 @@ use crate::protos::{coresdk, temporal}; use std::str::FromStr; /// Identifies a specific version of a worker deployment. -#[derive(Clone, Debug, Eq, PartialEq, Hash)] +#[derive(Clone, Debug, Eq, PartialEq, Hash, bon::Builder)] +#[builder(on(String, into), state_mod(vis = "pub"))] +#[non_exhaustive] pub struct WorkerDeploymentVersion { /// Name of the deployment pub deployment_name: String, diff --git a/crates/common/Cargo.toml b/crates/common/Cargo.toml index d09874739..85dcf5b9e 100644 --- a/crates/common/Cargo.toml +++ b/crates/common/Cargo.toml @@ -1,7 +1,8 @@ [package] name = "temporalio-common" -version = "0.6.0" +version = "1.0.0" edition = "2024" +rust-version = "1.88.0" authors = ["Temporal Technologies Inc. "] license-file = { workspace = true } description = "Common functionality for the Temporal SDK Core, Client, and Rust SDK" @@ -36,6 +37,7 @@ envconfig = ["dep:toml", "dep:dirs"] serde_serialize = ["temporalio-common-wasm/serde_serialize", "temporalio-protos/serde_serialize"] core-telemetry-bridge = ["dep:ringbuf", "dep:futures-channel"] core-based-sdk = ["core-telemetry-bridge", "prometheus", "envconfig"] +vendored-protox = ["temporalio-protos/vendored-protox"] [dependencies] anyhow = "1.0" @@ -106,12 +108,12 @@ uuid = { version = "1.18", default-features = false, features = ["v4"] } [dependencies.temporalio-protos] path = "../protos" -version = "0.6" +version = "~0.9.0" features = ["grpc-clients"] [dependencies.temporalio-common-wasm] path = "../common-wasm" -version = "0.6" +version = "~1.0.0" [build-dependencies] prost = { workspace = true } diff --git a/crates/common/README.md b/crates/common/README.md new file mode 100644 index 000000000..f4fe4cee6 --- /dev/null +++ b/crates/common/README.md @@ -0,0 +1,8 @@ +# `temporalio-common` + +[![crates.io](https://img.shields.io/crates/v/temporalio-common.svg)](https://crates.io/crates/temporalio-common) +[![docs.rs](https://docs.rs/temporalio-common/badge.svg)](https://docs.rs/temporalio-common) + +Part of [Temporal](https://temporal.io)'s [Rust SDK](https://github.com/temporalio/sdk-rust). + +Shared types and functionality used by Temporal SDK Core, the Rust Client, and the Rust SDK. diff --git a/crates/common/build.rs b/crates/common/build.rs index ecf894587..f8f1429f7 100644 --- a/crates/common/build.rs +++ b/crates/common/build.rs @@ -261,7 +261,11 @@ impl PayloadVisitorGenerator { } } - // Process oneofs + // Process oneofs in index order. This ordering reaches the generated file, and + // iterating the map directly emits a different field order on every build. + let mut oneof_fields: Vec<(i32, Vec<&FieldDescriptorProto>)> = + oneof_fields.into_iter().collect(); + oneof_fields.sort_by_key(|(oneof_index, _)| *oneof_index); for (oneof_index, oneof_field_list) in oneof_fields { let oneof_desc = &msg.oneof_decl[oneof_index as usize]; let oneof_name = oneof_desc.name.as_deref().unwrap_or(""); @@ -338,8 +342,12 @@ impl PayloadVisitorGenerator { let mut output = String::new(); output.push_str("// Generated from descriptors.bin - DO NOT EDIT\n\n"); - // Generate impls for each payload-containing type - for name in self.model.payload_containing.iter() { + // Generate impls for each payload-containing type, in name order. This output is + // compiled, so iterating the set directly emits the same impls in a different + // order on every build. + let mut payload_containing: Vec<&String> = self.model.payload_containing.iter().collect(); + payload_containing.sort(); + for name in payload_containing { if name == "temporal.api.common.v1.Payload" || name == "temporal.api.common.v1.Payloads" { continue; @@ -749,6 +757,7 @@ const BLOB_FIELDS: &[&str] = &[ "temporal.api.command.v1.SignalExternalWorkflowExecutionCommandAttributes.input", "temporal.api.command.v1.StartChildWorkflowExecutionCommandAttributes.input", "temporal.api.command.v1.UpsertWorkflowSearchAttributesCommandAttributes.search_attributes", // indexed_fields data-sum + "temporal.api.common.v1.Callback.NexusHandler.source_context", "temporal.api.protocol.v1.Message.body", // whole Any body; see EXTRA_WHOLE_MESSAGE_LEAVES "temporal.api.query.v1.WorkflowQuery.query_args", "temporal.api.workflow.v1.NewWorkflowExecutionInfo.input", diff --git a/crates/common/src/envconfig.rs b/crates/common/src/envconfig.rs index 1116d017e..9501b8c91 100644 --- a/crates/common/src/envconfig.rs +++ b/crates/common/src/envconfig.rs @@ -101,6 +101,7 @@ impl From for ConfigError { /// A source for configuration or a TLS certificate/key, from a path or raw data. #[derive(Debug, Clone, PartialEq)] +#[non_exhaustive] pub enum DataSource { /// A filesystem path to the data. Path(String), @@ -109,14 +110,18 @@ pub enum DataSource { } /// ClientConfig represents a client config file. -#[derive(Debug, Clone, PartialEq, Default)] +#[derive(Debug, Clone, PartialEq, Default, bon::Builder)] +#[non_exhaustive] pub struct ClientConfig { /// Profiles, keyed by profile name + #[builder(default)] pub profiles: HashMap, } /// ClientConfigProfile is profile-level configuration for a client. -#[derive(Debug, Clone, PartialEq, Default)] +#[derive(Debug, Clone, PartialEq, Default, bon::Builder)] +#[builder(on(String, into))] +#[non_exhaustive] pub struct ClientConfigProfile { /// Client address pub address: Option, @@ -137,11 +142,14 @@ pub struct ClientConfigProfile { /// Client gRPC metadata (aka headers). When loading from TOML and env var, or writing to TOML, the keys are /// lowercased and underscores are replaced with hyphens. This is used for deduplicating/overriding too, so manually /// set values that are not normalized may not get overridden when applying environment variables. + #[builder(default)] pub grpc_meta: HashMap, } /// ClientConfigTLS is TLS configuration for a client. -#[derive(Debug, Clone, PartialEq, Default)] +#[derive(Debug, Clone, PartialEq, Default, bon::Builder)] +#[builder(on(String, into))] +#[non_exhaustive] pub struct ClientConfigTLS { /// If Some(true), TLS is explicitly disabled. If Some(false), TLS is explicitly enabled. /// If None, TLS behavior depends on other factors (API key presence, etc.) @@ -160,11 +168,14 @@ pub struct ClientConfigTLS { pub server_name: Option, /// True if host verification should be skipped + #[builder(default)] pub disable_host_verification: bool, } /// Codec configuration for a client -#[derive(Debug, Clone, PartialEq, Default)] +#[derive(Debug, Clone, PartialEq, Default, bon::Builder)] +#[builder(on(String, into))] +#[non_exhaustive] pub struct ClientConfigCodec { /// Remote endpoint for the codec pub endpoint: Option, diff --git a/crates/common/src/lib.rs b/crates/common/src/lib.rs index 36a384a28..897487b76 100644 --- a/crates/common/src/lib.rs +++ b/crates/common/src/lib.rs @@ -18,10 +18,10 @@ pub mod protos; pub mod telemetry; pub mod worker; pub use temporalio_common_wasm::{ - ActivityDefinition, ActivityError, HasWorkflowDefinition, Memo, Priority, QueryDefinition, - RetryPolicy, SignalDefinition, UntypedActivity, UntypedWorkflow, UpdateDefinition, - WorkerDeploymentVersion, WorkflowDefinition, WorkflowExecution, data_converters, error, - search_attributes, + ActivityCloseTimeouts, ActivityDefinition, ActivityError, HasWorkflowDefinition, Memo, + MemoValue, MemoValues, Priority, QueryDefinition, RetryPolicy, SignalDefinition, + UntypedActivity, UntypedWorkflow, UpdateDefinition, WorkerDeploymentVersion, + WorkflowDefinition, WorkflowExecution, data_converters, error, search_attributes, }; macro_rules! dbg_panic { diff --git a/crates/common/src/payload_visitor.rs b/crates/common/src/payload_visitor.rs index 317d01c50..f9b1845a3 100644 --- a/crates/common/src/payload_visitor.rs +++ b/crates/common/src/payload_visitor.rs @@ -208,30 +208,33 @@ include!(concat!(env!("OUT_DIR"), "/payload_visitor_impl.rs")); #[cfg(test)] mod tests { use super::*; - use crate::protos::{ - coresdk::{ - activity_result::{ - ActivityResolution, Success, activity_resolution::Status as ActivityStatus, - }, - workflow_activation::{ - InitializeWorkflow, ResolveActivity, WorkflowActivation, WorkflowActivationJob, - workflow_activation_job::Variant, - }, - workflow_commands::{ - ContinueAsNewWorkflowExecution, ScheduleActivity, StartChildWorkflowExecution, - UpsertWorkflowSearchAttributes, WorkflowCommand, - workflow_command::Variant as CmdVariant, + use crate::{ + data_converters::WorkflowSerializationContext, + protos::{ + coresdk::{ + activity_result::{ + ActivityResolution, Success, activity_resolution::Status as ActivityStatus, + }, + workflow_activation::{ + InitializeWorkflow, ResolveActivity, WorkflowActivation, WorkflowActivationJob, + workflow_activation_job::Variant, + }, + workflow_commands::{ + ContinueAsNewWorkflowExecution, ScheduleActivity, StartChildWorkflowExecution, + UpsertWorkflowSearchAttributes, WorkflowCommand, + workflow_command::Variant as CmdVariant, + }, + workflow_completion::{ + WorkflowActivationCompletion, workflow_activation_completion::Status, + }, }, - workflow_completion::{ - WorkflowActivationCompletion, workflow_activation_completion::Status, + temporal::api::{ + common::v1::{Memo, SearchAttributes}, + failure::v1::failure::FailureInfo, + workflow::v1::WorkflowExecutionInfo, + workflowservice::v1::DescribeWorkflowExecutionResponse, }, }, - temporal::api::{ - common::v1::{Memo, SearchAttributes}, - failure::v1::failure::FailureInfo, - workflow::v1::WorkflowExecutionInfo, - workflowservice::v1::DescribeWorkflowExecutionResponse, - }, }; use futures::FutureExt; use std::{ @@ -427,7 +430,7 @@ mod tests { }, ..Default::default() })), - user_metadata: None, + ..Default::default() }], ..Default::default() }, @@ -438,7 +441,7 @@ mod tests { encode_payloads( &mut completion, &MarkingCodec, - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .await .unwrap(); @@ -482,7 +485,7 @@ mod tests { decode_payloads( &mut activation, &MarkingCodec, - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .await .unwrap(); @@ -521,7 +524,7 @@ mod tests { decode_payloads( &mut activation, &MarkingCodec, - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .await .unwrap(); @@ -566,7 +569,7 @@ mod tests { }), }, )), - user_metadata: None, + ..Default::default() }, // ContinueAsNewWorkflowExecution command WorkflowCommand { @@ -586,7 +589,7 @@ mod tests { ..Default::default() }, )), - user_metadata: None, + ..Default::default() }, // StartChildWorkflowExecution command WorkflowCommand { @@ -608,7 +611,7 @@ mod tests { ..Default::default() }, )), - user_metadata: None, + ..Default::default() }, ], ..Default::default() @@ -620,7 +623,7 @@ mod tests { encode_payloads( &mut completion, &MarkingCodec, - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .await .unwrap(); @@ -698,7 +701,7 @@ mod tests { decode_payloads( &mut response, &MarkingCodec, - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .await .unwrap(); @@ -728,7 +731,7 @@ mod tests { encode_payloads( &mut payload, &MarkingCodec, - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .await .unwrap(); @@ -743,7 +746,7 @@ mod tests { decode_payloads( &mut payload, &MarkingCodec, - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .await .unwrap(); @@ -760,7 +763,7 @@ mod tests { encode_payloads( &mut payloads, &MarkingCodec, - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .await .unwrap(); @@ -787,9 +790,13 @@ mod tests { ..Default::default() }; - let err = decode_payloads(&mut activation, &codec, &SerializationContextData::Workflow) - .await - .unwrap_err(); + let err = decode_payloads( + &mut activation, + &codec, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + ) + .await + .unwrap_err(); assert_eq!(err.to_string(), "Encoding error: visitor decode failed"); assert_eq!(codec.decode_calls.load(Ordering::SeqCst), 1); @@ -797,7 +804,7 @@ mod tests { #[tokio::test] async fn test_encode_failure_encodes_application_failure_details() { - let mut failure = DefaultFailureConverter.to_failure( + let mut failure = DefaultFailureConverter::default().to_failure( OutgoingError::Workflow(OutgoingWorkflowError::Application(Box::new( ApplicationFailure::builder(anyhow::anyhow!("app boom")) .details(crate::data_converters::RawValue::new(vec![make_payload( @@ -806,13 +813,13 @@ mod tests { .build(), ))), &PayloadConverter::default(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ); encode_payloads( &mut failure, &MarkingCodec, - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .await .unwrap(); diff --git a/crates/common/src/telemetry/prometheus_meter.rs b/crates/common/src/telemetry/prometheus_meter.rs index 82f8a7725..7df30f288 100644 --- a/crates/common/src/telemetry/prometheus_meter.rs +++ b/crates/common/src/telemetry/prometheus_meter.rs @@ -483,6 +483,7 @@ impl MetricAttributable> for PromHistogramF64 { #[derive(Debug)] pub struct CorePrometheusMeter { registry: Registry, + counters_total_suffix: bool, use_seconds_for_durations: bool, unit_suffix: bool, bucket_overrides: crate::telemetry::HistogramBucketOverrides, @@ -491,12 +492,14 @@ pub struct CorePrometheusMeter { impl CorePrometheusMeter { pub(super) fn new( registry: Registry, + counters_total_suffix: bool, use_seconds_for_durations: bool, unit_suffix: bool, bucket_overrides: crate::telemetry::HistogramBucketOverrides, ) -> Self { Self { registry, + counters_total_suffix, use_seconds_for_durations, unit_suffix, bucket_overrides, @@ -558,7 +561,10 @@ impl CoreMeter for CorePrometheusMeter { } fn counter(&self, params: MetricParameters) -> Counter { - let metric_name = params.name.to_string(); + let mut metric_name = params.name.to_string(); + if self.counters_total_suffix { + metric_name.push_str("_total"); + } Counter::new(Arc::new(PromMetric::::new( metric_name, params.description.to_string(), @@ -676,6 +682,7 @@ mod tests { metrics::{MetricKeyValue, NewAttributes, WORKFLOW_E2E_LATENCY_HISTOGRAM_NAME}, }; use prometheus::{Encoder, TextEncoder}; + use rstest::rstest; #[test] fn test_prometheus_meter_dynamic_labels() { @@ -684,6 +691,7 @@ mod tests { registry.clone(), false, false, + false, HistogramBucketOverrides::default(), ); @@ -717,6 +725,34 @@ mod tests { ); } + #[rstest] + #[case(false, "test_counter")] + #[case(true, "test_counter_total")] + fn counter_total_suffix_option_controls_counter_names( + #[case] counters_total_suffix: bool, + #[case] expected_name: &str, + ) { + let registry = Registry::new(HashMap::new()); + let meter = CorePrometheusMeter::new( + registry.clone(), + counters_total_suffix, + false, + false, + HistogramBucketOverrides::default(), + ); + let counter = meter.counter(MetricParameters { + name: "test_counter".into(), + description: "A test counter metric".into(), + unit: "".into(), + }); + + counter.adds(1); + + let output = output_string(®istry); + let expected_sample = format!("{expected_name} 1"); + assert!(output.lines().any(|line| line == expected_sample)); + } + #[test] fn test_extend_attributes() { let registry = Registry::new(HashMap::new()); @@ -724,6 +760,7 @@ mod tests { registry.clone(), false, false, + false, HistogramBucketOverrides::default(), ); @@ -763,6 +800,7 @@ mod tests { registry.clone(), false, false, + false, HistogramBucketOverrides::default(), ); @@ -787,6 +825,7 @@ mod tests { let registry_s = Registry::new(HashMap::new()); let meter_s = CorePrometheusMeter::new( registry_s.clone(), + false, true, false, HistogramBucketOverrides::default(), @@ -817,6 +856,7 @@ mod tests { registry.clone(), false, false, + false, HistogramBucketOverrides::default(), ); let counter = meter.counter(MetricParameters { @@ -841,6 +881,7 @@ mod tests { registry.clone(), false, false, + false, HistogramBucketOverrides::default(), ); @@ -875,6 +916,7 @@ mod tests { registry.clone(), false, false, + false, HistogramBucketOverrides::default(), ); let dashes = meter.counter(MetricParameters { @@ -896,6 +938,7 @@ mod tests { registry.clone(), false, false, + false, HistogramBucketOverrides::default(), ); let dashes = meter.counter(MetricParameters { diff --git a/crates/common/src/telemetry/prometheus_server.rs b/crates/common/src/telemetry/prometheus_server.rs index dd95d3e9e..bb54dd50e 100644 --- a/crates/common/src/telemetry/prometheus_server.rs +++ b/crates/common/src/telemetry/prometheus_server.rs @@ -29,6 +29,7 @@ pub fn start_prometheus_metric_exporter( let meter = Arc::new( crate::telemetry::prometheus_meter::CorePrometheusMeter::new( srv.registry().clone(), + opts.counters_total_suffix, opts.use_seconds_for_durations, opts.unit_suffix, opts.histogram_bucket_overrides, diff --git a/crates/common/src/worker.rs b/crates/common/src/worker.rs index 9852053c0..ff9bf26e4 100644 --- a/crates/common/src/worker.rs +++ b/crates/common/src/worker.rs @@ -135,10 +135,12 @@ pub struct WorkerDeploymentOptions { impl WorkerDeploymentOptions { /// Create deployment options from just a build ID, without opting into worker versioning. pub fn from_build_id(build_id: String) -> Self { - Self::new(WorkerDeploymentVersion { - deployment_name: "".to_owned(), - build_id, - }) + Self::new( + WorkerDeploymentVersion::builder() + .deployment_name("") + .build_id(build_id) + .build(), + ) .build() } } diff --git a/crates/macros/Cargo.toml b/crates/macros/Cargo.toml index a4ec61de9..009542072 100644 --- a/crates/macros/Cargo.toml +++ b/crates/macros/Cargo.toml @@ -1,7 +1,8 @@ [package] name = "temporalio-macros" -version = "0.6.0" +version = "1.0.0" edition = "2024" +rust-version = "1.88.0" authors = ["Temporal Technologies Inc. "] license-file = { workspace = true } description = "Procmacros used in Temporal Core & Rust SDKs" @@ -20,7 +21,7 @@ quote = "1.0" [dev-dependencies] # This is to enable doctests for the macros -temporalio-common = { path = "../common" } +temporalio-common = { path = "../common", version = "~1.0.0" } derive_more = { workspace = true } [package.metadata.workspaces] diff --git a/crates/macros/README.md b/crates/macros/README.md new file mode 100644 index 000000000..7ad6f207f --- /dev/null +++ b/crates/macros/README.md @@ -0,0 +1,8 @@ +# `temporalio-macros` + +[![crates.io](https://img.shields.io/crates/v/temporalio-macros.svg)](https://crates.io/crates/temporalio-macros) +[![docs.rs](https://docs.rs/temporalio-macros/badge.svg)](https://docs.rs/temporalio-macros) + +Part of [Temporal](https://temporal.io)'s [Rust SDK](https://github.com/temporalio/sdk-rust). + +Procedural macros for defining Temporal Workflows and Activities in Rust. diff --git a/crates/macros/src/activities_definitions.rs b/crates/macros/src/activities_definitions.rs index e91595443..b3277280f 100644 --- a/crates/macros/src/activities_definitions.rs +++ b/crates/macros/src/activities_definitions.rs @@ -366,8 +366,8 @@ impl ActivitiesDefinition { }) .collect(); - // Run methods and `ExecutableActivity`/`ActivityImplementer`/`HasOnlyStaticMethods` - // impls only make sense for real activities; definitions skip them entirely. + // Run methods and `ExecutableActivity`/`ActivityImplementer` impls only make sense for + // real activities; definitions skip them entirely. let run_impls: Vec<_> = if is_definitions { Vec::new() } else { @@ -396,14 +396,6 @@ impl ActivitiesDefinition { self.generate_activity_implementer_impl(impl_type, &module_ident) }; - let has_only_static = if !is_definitions && self.activities.iter().all(|a| a.is_static) { - quote! { - impl ::temporalio_sdk::activities::HasOnlyStaticMethods for #impl_type {} - } - } else { - quote! {} - }; - // Generate impl block with consts let const_impl = quote! { impl #impl_type { @@ -427,8 +419,6 @@ impl ActivitiesDefinition { #(#activity_impls)* #implementer_impl - - #has_only_static }; output.into() @@ -608,6 +598,7 @@ impl ActivitiesDefinition { let prefixed_method = format_ident!("__{}", activity.method.sig.ident); let has_input = !activity.input_types.is_empty(); + let is_instance = !activity.is_static; let receiver_pattern = if activity.is_static { quote! { _receiver } @@ -666,6 +657,7 @@ impl ActivitiesDefinition { impl ::temporalio_sdk::activities::ExecutableActivity for #module_ident::#struct_ident { type Implementer = #impl_type; + const REQUIRES_INSTANCE: bool = #is_instance; fn definition() -> Self { #module_ident::#struct_ident diff --git a/crates/macros/src/cloud_test_exclusion.rs b/crates/macros/src/cloud_test_exclusion.rs new file mode 100644 index 000000000..cd4455b34 --- /dev/null +++ b/crates/macros/src/cloud_test_exclusion.rs @@ -0,0 +1,240 @@ +use proc_macro2::TokenStream; +use quote::quote; +use syn::{ + Attribute, ItemFn, LitStr, Path, PathArguments, Token, + parse::{Parse, ParseStream}, + parse_quote, +}; + +pub(crate) fn expand(attr: TokenStream, item: TokenStream) -> syn::Result { + let exclusion = syn::parse2::(attr)?; + let mut test_fn = syn::parse2::(item)?; + let reason = &exclusion.reason; + let reason_type_check = quote!(const _: crate::CloudTestExclusionReason = #reason;); + let test_attribute_index = test_fn + .attrs + .iter() + .position(is_test_executor_attribute) + .ok_or_else(|| { + syn::Error::new_spanned( + &test_fn.sig.ident, + "cloud_test_exclusion can only be applied to a test function", + ) + })?; + + if test_fn + .attrs + .iter() + .any(|attribute| attribute.path().is_ident("ignore")) + { + return Ok(quote! { + #reason_type_check + #test_fn + }); + } + + let ignore_reason = exclusion.ignore_reason(); + // Rstest treats attributes before a case as case-specific, but copies attributes next to the + // test executor to every generated case. + test_fn.attrs.insert( + test_attribute_index, + parse_quote!(#[cfg_attr(feature = "cloud-test-mode", ignore = #ignore_reason)]), + ); + Ok(quote! { + #reason_type_check + #test_fn + }) +} + +struct CloudTestExclusion { + reason: Path, + note: Option, +} + +impl CloudTestExclusion { + fn ignore_reason(&self) -> LitStr { + let reason = self.reason.segments.last().unwrap().ident.to_string(); + let reason = match self.note.as_ref() { + Some(note) => format!("{reason}: {}", note.value()), + None => reason, + }; + LitStr::new(&reason, self.reason.segments.last().unwrap().ident.span()) + } +} + +impl Parse for CloudTestExclusion { + fn parse(input: ParseStream<'_>) -> syn::Result { + let reason = input.parse::()?; + if reason.leading_colon.is_some() + || reason.segments.len() != 3 + || reason.segments[0].ident != "crate" + || reason.segments[1].ident != "CloudTestExclusionReason" + || reason + .segments + .iter() + .any(|segment| !matches!(segment.arguments, PathArguments::None)) + { + return Err(syn::Error::new_spanned( + &reason, + "expected crate::CloudTestExclusionReason::", + )); + } + let note = if input.peek(Token![,]) { + input.parse::()?; + if input.is_empty() { + None + } else { + let note = input.parse::()?; + if note.value().trim().is_empty() { + return Err(syn::Error::new_spanned( + note, + "exclusion note cannot be empty", + )); + } + if input.peek(Token![,]) { + input.parse::()?; + } + Some(note) + } + } else { + None + }; + if !input.is_empty() { + return Err(input.error("unexpected cloud_test_exclusion argument")); + } + Ok(Self { reason, note }) + } +} + +fn is_test_executor_attribute(attribute: &Attribute) -> bool { + let path = attribute.path(); + path.is_ident("test") + || (path.segments.len() == 2 + && path.segments.first().unwrap().ident == "tokio" + && path.segments.last().unwrap().ident == "test") +} + +#[cfg(test)] +mod tests { + use super::*; + use quote::quote; + + #[test] + fn adds_note_to_ignore_reason() { + let expanded = expand( + quote!( + crate::CloudTestExclusionReason::DoesNotUseServer, + "Uses synthetic workflow history." + ), + quote! { + #[tokio::test] + async fn example() {} + }, + ) + .unwrap() + .to_string(); + + assert!(expanded.contains("cloud-test-mode")); + assert!(expanded.contains("DoesNotUseServer: Uses synthetic workflow history.")); + assert!(expanded.contains("const _ : crate :: CloudTestExclusionReason")); + } + + #[test] + fn accepts_reason_without_note() { + let expanded = expand( + quote!(crate::CloudTestExclusionReason::DoesNotUseServer), + quote! { + #[test] + fn example() {} + }, + ) + .unwrap() + .to_string(); + + assert!(expanded.contains("ignore = \"DoesNotUseServer\"")); + } + + #[test] + fn puts_ignore_after_rstest_cases() { + let expanded = expand( + quote!( + crate::CloudTestExclusionReason::RequiresCloudProvisioning, + "Requires a configured search attribute." + ), + quote! { + #[rstest::rstest] + #[case(true)] + #[case(false)] + #[tokio::test] + async fn example(#[case] value: bool) {} + }, + ) + .unwrap(); + let expanded = syn::parse2::(expanded).unwrap(); + let syn::Item::Fn(expanded) = &expanded.items[1] else { + panic!("second expanded item must remain the annotated test") + }; + + assert_eq!( + expanded.attrs[0].path().segments.last().unwrap().ident, + "rstest" + ); + assert!(expanded.attrs[1].path().is_ident("case")); + assert!(expanded.attrs[2].path().is_ident("case")); + assert!(expanded.attrs[3].path().is_ident("cfg_attr")); + assert!(is_test_executor_attribute(&expanded.attrs[4])); + } + + #[test] + fn preserves_permanently_ignored_test() { + let expanded = expand( + quote!( + crate::CloudTestExclusionReason::RequiresLocalServer, + "Runs only against the local test server." + ), + quote! { + #[ignore = "Manual test"] + #[test] + fn example() {} + }, + ) + .unwrap() + .to_string(); + + assert!(!expanded.contains("cfg_attr")); + assert!(expanded.contains("Manual test")); + assert!(expanded.contains("const _ : crate :: CloudTestExclusionReason")); + } + + #[test] + fn rejects_noncanonical_reason_path() { + let error = expand( + quote!(DoesNotUseServer), + quote! { + #[test] + fn example() {} + }, + ) + .unwrap_err(); + + assert!( + error + .to_string() + .contains("expected crate::CloudTestExclusionReason::") + ); + } + + #[test] + fn rejects_empty_note() { + let error = expand( + quote!(crate::CloudTestExclusionReason::NeedsCloudAdaptation, " "), + quote! { + #[test] + fn example() {} + }, + ) + .unwrap_err(); + + assert!(error.to_string().contains("exclusion note cannot be empty")); + } +} diff --git a/crates/macros/src/lib.rs b/crates/macros/src/lib.rs index 4c880e6e7..c426da368 100644 --- a/crates/macros/src/lib.rs +++ b/crates/macros/src/lib.rs @@ -3,6 +3,7 @@ use proc_macro2::TokenStream as TokenStream2; use syn::{parse::Parser, parse_macro_input}; mod activities_definitions; +mod cloud_test_exclusion; mod fsm_impl; mod macro_utils; mod workflow_definitions; @@ -35,6 +36,19 @@ pub fn activity_definitions(_attr: TokenStream, item: TokenStream) -> TokenStrea def.codegen() } +/// Marks an SDK integration-test function that should not run against Temporal Cloud. +/// +/// The first argument is a `crate::CloudTestExclusionReason` variant. An optional second argument +/// records information not already conveyed by the variant. +#[doc(hidden)] +#[proc_macro_attribute] +pub fn cloud_test_exclusion(attr: TokenStream, item: TokenStream) -> TokenStream { + match cloud_test_exclusion::expand(attr.into(), item.into()) { + Ok(output) => output.into(), + Err(error) => error.into_compile_error().into(), + } +} + /// Marks a struct as a workflow definition. /// /// By default, the struct name is used as the workflow type name. To specify a custom workflow diff --git a/crates/macros/src/workflow_definitions.rs b/crates/macros/src/workflow_definitions.rs index f73fc63c6..108f4e546 100644 --- a/crates/macros/src/workflow_definitions.rs +++ b/crates/macros/src/workflow_definitions.rs @@ -75,10 +75,13 @@ fn generate_decode_arm( ) -> TokenStream2 { quote! { #handler_name => { - let ctx = ::temporalio_workflow::common::data_converters::SerializationContext { - data: &::temporalio_workflow::common::data_converters::SerializationContextData::Workflow, + let context_data = ::temporalio_workflow::common::data_converters::SerializationContextData::Workflow( + ::temporalio_workflow::common::data_converters::WorkflowSerializationContext::new() + ); + let ctx = ::temporalio_workflow::common::data_converters::SerializationContext::new( + &context_data, converter, - }; + ); let input: #input_type = <::temporalio_workflow::common::data_converters::PayloadConverter as ::temporalio_workflow::common::data_converters::GenericPayloadConverter>::from_payloads( converter, &ctx, @@ -771,8 +774,8 @@ impl WorkflowMethodsDefinition { fn handle( mut ctx: ::temporalio_workflow::WorkflowContext, input: <#module_ident::#struct_ident as ::temporalio_workflow::common::SignalDefinition>::Input, - ) -> ::temporalio_workflow::__private::futures_util::future::LocalBoxFuture<'static, ()> { - ::temporalio_workflow::__private::futures_util::FutureExt::boxed_local( + ) -> ::temporalio_workflow::__private::LocalBoxFuture<'static, ()> { + ::temporalio_workflow::__private::FutureExt::boxed_local( async move { #method_call.await } ) } @@ -904,14 +907,14 @@ impl WorkflowMethodsDefinition { }; let handle_body = if update.is_fallible { quote! { - ::temporalio_workflow::__private::futures_util::FutureExt::boxed_local( + ::temporalio_workflow::__private::FutureExt::boxed_local( async move { #method_call.await } ) } } else { quote! { - ::temporalio_workflow::__private::futures_util::FutureExt::boxed_local( - ::temporalio_workflow::__private::futures_util::FutureExt::map( + ::temporalio_workflow::__private::FutureExt::boxed_local( + ::temporalio_workflow::__private::FutureExt::map( async move { #method_call.await }, Ok, ) @@ -923,7 +926,7 @@ impl WorkflowMethodsDefinition { fn handle( mut ctx: ::temporalio_workflow::WorkflowContext, input: <#module_ident::#struct_ident as ::temporalio_workflow::common::UpdateDefinition>::Input, - ) -> ::temporalio_workflow::__private::futures_util::future::LocalBoxFuture<'static, Result<<#module_ident::#struct_ident as ::temporalio_workflow::common::UpdateDefinition>::Output, Box>> { + ) -> ::temporalio_workflow::__private::LocalBoxFuture<'static, Result<<#module_ident::#struct_ident as ::temporalio_workflow::common::UpdateDefinition>::Output, Box>> { #handle_body } @@ -1012,7 +1015,7 @@ impl WorkflowMethodsDefinition { }; let run_impl_body = quote! { - ::temporalio_workflow::__private::futures_util::FutureExt::boxed_local(async move { + ::temporalio_workflow::__private::FutureExt::boxed_local(async move { let result = #run_call; match result { Ok(value) => Ok( @@ -1103,7 +1106,7 @@ impl WorkflowMethodsDefinition { _ctx: ::temporalio_workflow::WorkflowContext, name: &str, _input: ::std::boxed::Box, - ) -> ::temporalio_workflow::__private::futures_util::future::LocalBoxFuture<'static, Result<(), ::temporalio_workflow::workflows::WorkflowError>> { + ) -> ::temporalio_workflow::__private::LocalBoxFuture<'static, Result<(), ::temporalio_workflow::workflows::WorkflowError>> { unreachable!("typed signal dispatch called for unknown signal handler '{name}'") } } @@ -1113,7 +1116,7 @@ impl WorkflowMethodsDefinition { ctx: ::temporalio_workflow::WorkflowContext, name: &str, input: ::std::boxed::Box, - ) -> ::temporalio_workflow::__private::futures_util::future::LocalBoxFuture<'static, Result<(), ::temporalio_workflow::workflows::WorkflowError>> { + ) -> ::temporalio_workflow::__private::LocalBoxFuture<'static, Result<(), ::temporalio_workflow::workflows::WorkflowError>> { match name { #(#dispatch_signal_arms)* _ => unreachable!("typed signal dispatch called for unknown signal handler '{name}'"), @@ -1177,7 +1180,7 @@ impl WorkflowMethodsDefinition { let handler_name = &info.handler_name; let has_validator = u.validator.is_some(); quote! { - ::temporalio_workflow::runtime::types::UpdateDefinitionDescriptor { + ::temporalio_workflow::__private::macros::UpdateDefinitionDescriptor { name: (#handler_name).to_string(), has_validator: #has_validator, } @@ -1235,7 +1238,7 @@ impl WorkflowMethodsDefinition { _ctx: ::temporalio_workflow::WorkflowContext, name: &str, _input: ::std::boxed::Box, - ) -> ::temporalio_workflow::__private::futures_util::future::LocalBoxFuture<'static, Result<::std::boxed::Box, ::temporalio_workflow::workflows::WorkflowError>> { + ) -> ::temporalio_workflow::__private::LocalBoxFuture<'static, Result<::std::boxed::Box, ::temporalio_workflow::workflows::WorkflowError>> { unreachable!("typed update dispatch called for unknown update handler '{name}'") } @@ -1254,7 +1257,7 @@ impl WorkflowMethodsDefinition { ctx: ::temporalio_workflow::WorkflowContext, name: &str, input: ::std::boxed::Box, - ) -> ::temporalio_workflow::__private::futures_util::future::LocalBoxFuture<'static, Result<::std::boxed::Box, ::temporalio_workflow::workflows::WorkflowError>> { + ) -> ::temporalio_workflow::__private::LocalBoxFuture<'static, Result<::std::boxed::Box, ::temporalio_workflow::workflows::WorkflowError>> { match name { #(#dispatch_update_arms)* _ => unreachable!("typed update dispatch called for unknown update handler '{name}'"), @@ -1374,7 +1377,7 @@ impl WorkflowMethodsDefinition { }; quote! { - impl ::temporalio_workflow::runtime::entry::WorkflowImplementation for #impl_type { + impl ::temporalio_workflow::__private::macros::WorkflowImplementation for #impl_type { type Run = #module_ident::#run_struct_ident; const HAS_INIT: bool = #has_init; @@ -1384,8 +1387,8 @@ impl WorkflowMethodsDefinition { <#impl_type>::name() } - fn definition() -> ::temporalio_workflow::runtime::types::WorkflowDefinitionDescriptor { - ::temporalio_workflow::runtime::types::WorkflowDefinitionDescriptor { + fn definition() -> ::temporalio_workflow::__private::macros::WorkflowDefinitionDescriptor { + ::temporalio_workflow::__private::macros::WorkflowDefinitionDescriptor { workflow_type: Self::name().to_string(), has_init: #has_init, init_takes_input: #init_has_input, @@ -1405,7 +1408,7 @@ impl WorkflowMethodsDefinition { fn run( mut ctx: ::temporalio_workflow::WorkflowContext, input: ::std::option::Option<::Input>, - ) -> ::temporalio_workflow::__private::futures_util::future::LocalBoxFuture<'static, Result<::std::boxed::Box, ::temporalio_workflow::WorkflowTermination>> { + ) -> ::temporalio_workflow::__private::LocalBoxFuture<'static, Result<::std::boxed::Box, ::temporalio_workflow::WorkflowTermination>> { #run_impl_body } diff --git a/crates/protos/Cargo.toml b/crates/protos/Cargo.toml index 8e8485c82..d71f372b4 100644 --- a/crates/protos/Cargo.toml +++ b/crates/protos/Cargo.toml @@ -1,12 +1,13 @@ [package] name = "temporalio-protos" -version = "0.6.0" +version = "0.9.0" edition = "2024" +rust-version = "1.88.0" authors = ["Temporal Technologies Inc. "] license-file = { workspace = true } description = "Compiled protobuf definitions for the Temporal Rust SDK" homepage = "https://temporal.io/" -repository = "https://github.com/temporalio/sdk-core" +repository = "https://github.com/temporalio/sdk-rust" keywords = ["temporal", "protobuf"] categories = ["development-tools"] exclude = ["protos/*/.github/*"] @@ -16,10 +17,11 @@ links = "temporalio_protos" default = [] serde_serialize = [] grpc-clients = ["tonic/channel"] +vendored-protox = ["dep:protox", "prost-types/vendored-protox"] [dependencies] anyhow = "1.0" -base64 = "0.22" +base64 = "0.23" derive_more = { workspace = true } http = "1" prost = { workspace = true } @@ -33,9 +35,9 @@ pbjson = { workspace = true } [build-dependencies] prost = { workspace = true } -prost-types = "0.14" tonic-prost-build = { workspace = true } pbjson-build = { workspace = true } +protox = { version = "0.9.1", optional = true } [lints] workspace = true diff --git a/crates/protos/README.md b/crates/protos/README.md new file mode 100644 index 000000000..f580d6cfa --- /dev/null +++ b/crates/protos/README.md @@ -0,0 +1,18 @@ +# `temporalio-protos` + +[![crates.io](https://img.shields.io/crates/v/temporalio-protos.svg)](https://crates.io/crates/temporalio-protos) +[![docs.rs](https://docs.rs/temporalio-protos/badge.svg)](https://docs.rs/temporalio-protos) + +Part of [Temporal](https://temporal.io)'s [Rust SDK](https://github.com/temporalio/sdk-rust). + +Compiled protobuf definitions for Temporal APIs and SDK Core protocols. + +This crate remains on a `0.x` version because generated protobuf messages are not marked +`#[non_exhaustive]`. Adding a protobuf field can therefore be a source-breaking change for code +that constructs a message with a struct literal, so minor releases may contain breaking changes. + +Most Rust SDK users should use [`temporalio-client`](https://crates.io/crates/temporalio-client) or +[`temporalio-sdk`](https://crates.io/crates/temporalio-sdk) instead. + +Enable the `vendored-protox` feature to compile the protobuf definitions with the pure-Rust +`protox` implementation instead of requiring an installed `protoc` binary. diff --git a/crates/protos/build.rs b/crates/protos/build.rs index 5ca388c9f..d829d8d76 100644 --- a/crates/protos/build.rs +++ b/crates/protos/build.rs @@ -1,4 +1,7 @@ use std::{env, path::PathBuf}; + +#[cfg(feature = "vendored-protox")] +use prost::Message; use tonic_prost_build::Config; static ALWAYS_SERDE: &str = "#[cfg_attr(not(feature = \"serde_serialize\"), \ @@ -28,6 +31,7 @@ const SERDE_DERIVE_PREFIXES: &[&str] = &[ ".temporal.api.history", ".temporal.api.namespace", ".temporal.api.nexus", + ".temporal.api.nexusoperation", ".temporal.api.nexusservices", ".temporal.api.notification", ".temporal.api.operatorservice", @@ -52,6 +56,26 @@ fn main() -> Result<(), Box> { let out = PathBuf::from(env::var("OUT_DIR").unwrap()); let descriptor_file = out.join("descriptors.bin"); println!("cargo:descriptor_path={}", descriptor_file.display()); + let protos = &[ + "./protos/local/temporal/sdk/core/core_interface.proto", + "./protos/api_upstream/temporal/api/sdk/v1/workflow_metadata.proto", + "./protos/api_upstream/temporal/api/workflowservice/v1/service.proto", + "./protos/api_upstream/temporal/api/nexusservices/workerservice/v1/request_response.proto", + "./protos/api_upstream/temporal/api/operatorservice/v1/service.proto", + "./protos/api_upstream/temporal/api/errordetails/v1/message.proto", + "./protos/api_cloud_upstream/temporal/api/cloud/cloudservice/v1/service.proto", + "./protos/testsrv_upstream/temporal/api/testservice/v1/service.proto", + "./protos/grpc/health/v1/health.proto", + "./protos/google/rpc/status.proto", + ]; + let includes = &[ + "./protos/api_upstream", + "./protos/api_cloud_upstream", + "./protos/local", + "./protos/testsrv_upstream", + "./protos/grpc", + "./protos", + ]; let mut builder = tonic_prost_build::configure() // Workflow guests need message structs, while the native common crate enables this // feature to preserve the generated clients it re-exports today. @@ -147,6 +171,13 @@ fn main() -> Result<(), Box> { builder = builder.type_attribute(*prefix, SERDE_ATTR); } + #[cfg(feature = "vendored-protox")] + { + let descriptors = protox::compile(protos, includes)?; + std::fs::write(&descriptor_file, descriptors.encode_to_vec())?; + builder = builder.skip_protoc_run(); + } + builder .file_descriptor_set_path(&descriptor_file) .compile_with_config( @@ -155,26 +186,8 @@ fn main() -> Result<(), Box> { c.enable_type_names(); c }, - &[ - "./protos/local/temporal/sdk/core/core_interface.proto", - "./protos/api_upstream/temporal/api/sdk/v1/workflow_metadata.proto", - "./protos/api_upstream/temporal/api/workflowservice/v1/service.proto", - "./protos/api_upstream/temporal/api/nexusservices/workerservice/v1/request_response.proto", - "./protos/api_upstream/temporal/api/operatorservice/v1/service.proto", - "./protos/api_upstream/temporal/api/errordetails/v1/message.proto", - "./protos/api_cloud_upstream/temporal/api/cloud/cloudservice/v1/service.proto", - "./protos/testsrv_upstream/temporal/api/testservice/v1/service.proto", - "./protos/grpc/health/v1/health.proto", - "./protos/google/rpc/status.proto", - ], - &[ - "./protos/api_upstream", - "./protos/api_cloud_upstream", - "./protos/local", - "./protos/testsrv_upstream", - "./protos/grpc", - "./protos", - ], + protos, + includes, )?; // TODO [rust-sdk-branch]: support normal JSON and proto JSON serialization diff --git a/crates/protos/protos/api_upstream/.github/CODEOWNERS b/crates/protos/protos/api_upstream/.github/CODEOWNERS index 07a91fabc..1b11a8dad 100644 --- a/crates/protos/protos/api_upstream/.github/CODEOWNERS +++ b/crates/protos/protos/api_upstream/.github/CODEOWNERS @@ -2,4 +2,5 @@ # https://docs.github.com/en/github/creating-cloning-and-archiving-repositories/about-code-owners#codeowners-syntax * @temporalio/server @temporalio/sdk +*.wit @temporalio/sdk api/temporal/api/sdk/* @temporalio/sdk diff --git a/crates/protos/protos/api_upstream/.github/workflows/ci.yml b/crates/protos/protos/api_upstream/.github/workflows/ci.yml index 79b216936..53cde483c 100644 --- a/crates/protos/protos/api_upstream/.github/workflows/ci.yml +++ b/crates/protos/protos/api_upstream/.github/workflows/ci.yml @@ -8,7 +8,7 @@ jobs: name: ci runs-on: ubuntu-latest steps: - - uses: actions/checkout@34e114876b0b11c390a56381ad16ebd13914f8d5 # v4.3.1 + - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 - uses: actions/setup-go@7b8cf10d4e4a01d4992d18a89f4d7dc5a3e6d6f4 # v4.3.0 with: go-version: '^1.21' diff --git a/crates/protos/protos/api_upstream/.github/workflows/create-release.yml b/crates/protos/protos/api_upstream/.github/workflows/create-release.yml index c9c480964..749493f7e 100644 --- a/crates/protos/protos/api_upstream/.github/workflows/create-release.yml +++ b/crates/protos/protos/api_upstream/.github/workflows/create-release.yml @@ -20,6 +20,22 @@ on: description: An ID used by external tools to identify workflow runs(can be left empty when running manually) default: "none" type: string + api_ref: + description: "api commit or ref to release; defaults to `branch`" + default: "" + type: string + api_go_ref: + description: "api-go commit or ref to release; defaults to `branch`" + default: "" + type: string + auto_publish: + description: "Publish the release instead of leaving a draft. A draft creates no tag, so automated callers need this." + default: false + type: boolean + +permissions: + contents: read + jobs: dispatch: runs-on: ubuntu-latest @@ -33,19 +49,24 @@ jobs: api_commit_sha: ${{ steps.pin_commits.outputs.api_commit_sha }} api_go_commit_sha: ${{ steps.pin_commits.outputs.api_go_commit_sha }} steps: + # api_ref/api_go_ref let a caller release specific commits rather than + # the head of a branch. They must name a corresponding pair: the "Pin + # commits sha" step below asserts that the api-go commit's proto/api + # submodule points at the api commit. Left empty, both fall back to + # `branch`. - name: Checkout api - uses: actions/checkout@34e114876b0b11c390a56381ad16ebd13914f8d5 # v4.3.1 + uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 with: - ref: ${{ github.event.inputs.branch }} + ref: ${{ inputs.api_ref || inputs.branch }} fetch-depth: 0 fetch-tags: true path: api - name: Checkout api-go - uses: actions/checkout@34e114876b0b11c390a56381ad16ebd13914f8d5 # v4.3.1 + uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 with: repository: temporalio/api-go - ref: ${{ github.event.inputs.branch }} + ref: ${{ inputs.api_go_ref || inputs.branch }} submodules: true path: api-go @@ -114,7 +135,7 @@ jobs: owner: ${{ github.repository_owner }} - name: Checkout - uses: actions/checkout@34e114876b0b11c390a56381ad16ebd13914f8d5 # v4.3.1 + uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 with: ref: ${{ needs.prepare-inputs.outputs.api_commit_sha }} token: ${{ steps.generate_token.outputs.token }} @@ -141,3 +162,33 @@ jobs: api_commit_sha: ${{ needs.prepare-inputs.outputs.api_commit_sha }} base_tag: ${{ github.event.inputs.base_tag }} secrets: inherit + + # Publishing this repo's release fires `release: published`, which + # trigger-api-go-publish-release.yml turns into api-go's publish-release.yml + # -- the step that actually creates the api-go tag. + publish-release: + name: "Publish release" + needs: + - create-release + - release-api-go + if: | + !cancelled() && + (inputs.auto_publish == true || inputs.auto_publish == 'true') && + needs.create-release.result == 'success' && + needs.release-api-go.result == 'success' + runs-on: ubuntu-latest + + steps: + - name: Generate token + id: generate_token + uses: actions/create-github-app-token@d72941d797fd3113feb6b93fd0dec494b13a2547 # v1.12.0 + with: + app-id: ${{ secrets.TEMPORAL_CICD_APP_ID }} + private-key: ${{ secrets.TEMPORAL_CICD_PRIVATE_KEY }} + owner: ${{ github.repository_owner }} + + - name: Publish release + env: + GH_TOKEN: ${{ steps.generate_token.outputs.token }} + TAG: ${{ github.event.inputs.tag }} + run: gh release edit "$TAG" --draft=false -R "$GITHUB_REPOSITORY" --latest diff --git a/crates/protos/protos/api_upstream/.github/workflows/push-to-buf.yml b/crates/protos/protos/api_upstream/.github/workflows/push-to-buf.yml index f2cd0a675..e99912c80 100644 --- a/crates/protos/protos/api_upstream/.github/workflows/push-to-buf.yml +++ b/crates/protos/protos/api_upstream/.github/workflows/push-to-buf.yml @@ -13,7 +13,7 @@ jobs: runs-on: ubuntu-latest steps: - name: Checkout repo - uses: actions/checkout@34e114876b0b11c390a56381ad16ebd13914f8d5 # v4.3.1 + uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 - uses: bufbuild/buf-action@fd21066df7214747548607aaa45548ba2b9bc1ff # v1.4.0 with: version: 1.49.0 diff --git a/crates/protos/protos/api_upstream/Makefile b/crates/protos/protos/api_upstream/Makefile index d43b05677..f6fcf899d 100644 --- a/crates/protos/protos/api_upstream/Makefile +++ b/crates/protos/protos/api_upstream/Makefile @@ -97,7 +97,7 @@ api-linter-install: buf-install: printf $(COLOR) "Install/update buf..." - go install github.com/bufbuild/buf/cmd/buf@v1.27.0 + go install github.com/bufbuild/buf/cmd/buf@v1.49.0 ##### Sync external proto dependencies ##### sync-nexus-annotations: @@ -112,12 +112,12 @@ api-linter: $(STAMPDIR): mkdir $@ -$(STAMPDIR)/buf-mod-prune: $(STAMPDIR) buf.yaml +$(STAMPDIR)/buf-dep-prune: $(STAMPDIR) buf.yaml printf $(COLOR) "Pruning buf module" - buf mod prune + buf dep prune touch $@ -buf-lint: $(STAMPDIR)/buf-mod-prune +buf-lint: $(STAMPDIR)/buf-dep-prune printf $(COLOR) "Run buf linter..." (cd $(PROTO_ROOT) && buf lint) diff --git a/crates/protos/protos/api_upstream/buf.lock b/crates/protos/protos/api_upstream/buf.lock index f43352bf2..e0559987a 100644 --- a/crates/protos/protos/api_upstream/buf.lock +++ b/crates/protos/protos/api_upstream/buf.lock @@ -1,18 +1,9 @@ # Generated by buf. DO NOT EDIT. -version: v1 +version: v2 deps: - - remote: buf.build - owner: googleapis - repository: googleapis + - name: buf.build/googleapis/googleapis commit: 004180b77378443887d3b55cabc00384 - digest: shake256:d26c7c2fd95f0873761af33ca4a0c0d92c8577122b6feb74eb3b0a57ebe47a98ab24a209a0e91945ac4c77204e9da0c2de0020b2cedc27bdbcdea6c431eec69b - - remote: buf.build - owner: grpc-ecosystem - repository: grpc-gateway - commit: 6467306b4f624747aaf6266762ee7a1c - digest: shake256:833d648b99b9d2c18b6882ef41aaeb113e76fc38de20dda810c588d133846e6593b4da71b388bcd921b1c7ab41c7acf8f106663d7301ae9e82ceab22cf64b1b7 - - remote: buf.build - owner: temporalio - repository: nexus-annotations + digest: b5:e8f475fe3330f31f5fd86ac689093bcd274e19611a09db91f41d637cb9197881ce89882b94d13a58738e53c91c6e4bae7dc1feba85f590164c975a89e25115dc + - name: buf.build/temporalio/nexus-annotations commit: 599b78404fbe4e78b833d527a1d0da40 - digest: shake256:1f41ef11ccbf31d7318b0fe1915550ba6567c99dc94694d60b117fc1ffc756290ba9766c58b403986f079e2b861b42538e5f8cf0495f744cd390d223b81854ca + digest: b5:feb0298a2e7e60058a5dee533e166e152bd0c3b9f776170946fba80737722022a39a65b132af1028d150d2c5bc52990694f60fbcf89fbb9693f7a6b7803d9203 diff --git a/crates/protos/protos/api_upstream/buf.yaml b/crates/protos/protos/api_upstream/buf.yaml index 2f2fa5389..d00ec17a7 100644 --- a/crates/protos/protos/api_upstream/buf.yaml +++ b/crates/protos/protos/api_upstream/buf.yaml @@ -1,24 +1,23 @@ -version: v1 -name: buf.build/temporalio/api +version: v2 +modules: + - path: . + name: buf.build/temporalio/api + excludes: + # Vendored for api-linter (can't read the BSR); excluded so buf sees them once. + - google + - nexusannotations deps: - - buf.build/grpc-ecosystem/grpc-gateway - buf.build/googleapis/googleapis - buf.build/temporalio/nexus-annotations -build: - excludes: - # Buf won't accept a local dependency on the google protos but we need them - # to run api-linter, so just tell buf it ignore it - - google - # Same for nexusannotations - local copy for api-linter, BSR dep for buf - - nexusannotations -breaking: +lint: use: - - WIRE_JSON + - STANDARD ignore: + - cmd - google -lint: + disallow_comment_ignores: true +breaking: use: - - DEFAULT + - WIRE_JSON ignore: - google - - cmd diff --git a/crates/protos/protos/api_upstream/nexus/deps/nexus-temporal-types/model.wit b/crates/protos/protos/api_upstream/nexus/deps/nexus-temporal-types/model.wit index e435bebaa..c661ee510 100644 --- a/crates/protos/protos/api_upstream/nexus/deps/nexus-temporal-types/model.wit +++ b/crates/protos/protos/api_upstream/nexus/deps/nexus-temporal-types/model.wit @@ -9,15 +9,28 @@ interface model { /// python="typing.Any" /// typescript="common.Payload" /// dotnet="object?" + /// dotnet-from="ProtoExtensions.FromPayload" /// dotnet-to="ProtoExtensions.ToPayload" /// typescript-import="@temporalio/common" type payload = placeholder; /// @nexus.proto "temporal.api.common.v1.Payloads" /// typescript-import="@temporalio/proto" - /// @nexus.type dotnet="IReadOnlyCollection" dotnet-to="ProtoExtensions.ToPayloads" + /// @nexus.type dotnet="IReadOnlyCollection" dotnet-from="ProtoExtensions.FromPayloads" dotnet-to="ProtoExtensions.ToPayloads" type payloads = list; + /// Temporal failure represented by the target SDK's native exception/error type. + /// The SDK failure converter owns the recursive cause and failure-info structure. + /// @nexus.proto "temporal.api.failure.v1.Failure" typescript-import="@temporalio/proto" + /// @nexus.type + /// python="BaseException" + /// typescript="Error" + /// go="error" + /// dotnet="System.Exception" + /// dotnet-from="ProtoExtensions.FromFailureProto" + /// dotnet-to="ProtoExtensions.ToFailureProto" + type failure = placeholder; + /// Callable result annotation for workflow functions. /// @nexus.type /// python="collections.abc.Awaitable[WorkflowResult]" @@ -50,6 +63,7 @@ interface model { /// python="str" /// typescript="string" /// dotnet="string" + /// dotnet-from="ProtoExtensions.FromWorkflowTypeProto" /// dotnet-to="ProtoExtensions.ToWorkflowTypeProto" type workflow-type = placeholder; @@ -72,9 +86,9 @@ interface model { /// typescript-name-extractor="signalFunctionName" /// dotnet-name-extractor="TemporalFunctionNames.SignalName" /// dotnet-call-extractor="TemporalFunctionNames.ExtractCall" - /// typescript-value-type="workflow.SignalDefinition" - /// typescript-args-type="Value extends workflow.SignalDefinition ? Args : never" - /// typescript-import="@temporalio/workflow" + /// typescript-value-type="common.SignalDefinition" + /// typescript-args-type="Value extends common.SignalDefinition ? Args : never" + /// typescript-import="@temporalio/common" /// @nexus.add-rpc-compatible-with "string" type signal-function = placeholder; @@ -83,6 +97,7 @@ interface model { /// python="temporalio.common.RetryPolicy" /// typescript="common.RetryPolicy" /// dotnet="Temporalio.Common.RetryPolicy" + /// dotnet-from="ProtoExtensions.FromRetryPolicyProto" /// typescript-import="@temporalio/common" type retry-policy = placeholder; @@ -91,18 +106,31 @@ interface model { /// python="str" /// typescript="string" /// dotnet="string" + /// dotnet-from="ProtoExtensions.FromTaskQueueProto" /// dotnet-to="ProtoExtensions.ToTaskQueueProto" type task-queue = placeholder; /// @nexus.proto "temporal.api.common.v1.Memo" typescript-import="@temporalio/proto" - /// @nexus.type python="collections.abc.Mapping[str, typing.Any]" typescript="Record" dotnet="IReadOnlyDictionary" + /// @nexus.type python="collections.abc.Mapping[str, typing.Any]" typescript="Record" dotnet="IReadOnlyDictionary" dotnet-from="ProtoExtensions.FromMemoProto" type memo = placeholder; + /// @nexus.proto "temporal.api.common.v1.Header" typescript-import="@temporalio/proto" + /// @nexus.type + /// python="collections.abc.Mapping[str, typing.Any]" + /// typescript="Record" + /// go="map[string]any" + /// dotnet="IReadOnlyDictionary" + /// dotnet-from="ProtoExtensions.FromHeaderProto" + /// dotnet-to="ProtoExtensions.ToHeaderProto" + /// typescript-import="@temporalio/common" + type header = placeholder; + /// @nexus.proto "temporal.api.common.v1.SearchAttributes" typescript-import="@temporalio/proto" /// @nexus.type /// python="temporalio.common.TypedSearchAttributes" /// typescript="common.TypedSearchAttributes" /// dotnet="Temporalio.Common.SearchAttributeCollection" + /// dotnet-from="ProtoExtensions.FromSearchAttributesProto" /// typescript-import="@temporalio/common" type search-attributes = placeholder; @@ -111,6 +139,7 @@ interface model { /// python="temporalio.common.Priority" /// typescript="common.Priority" /// dotnet="Temporalio.Common.Priority" + /// dotnet-from="ProtoExtensions.FromPriorityProto" /// typescript-import="@temporalio/common" type priority = placeholder; @@ -119,6 +148,7 @@ interface model { /// python="temporalio.common.VersioningOverride" /// typescript="common.VersioningOverride" /// dotnet="Temporalio.Common.VersioningOverride" + /// dotnet-from="ProtoExtensions.FromVersioningOverrideProto" /// typescript-import="@temporalio/common" type versioning-override = placeholder; @@ -127,6 +157,7 @@ interface model { /// python="datetime.timedelta" /// typescript="common.Duration" /// dotnet="System.TimeSpan" + /// dotnet-from="ProtoExtensions.FromDurationProto" /// typescript-import="@temporalio/common" type duration = placeholder; @@ -146,8 +177,10 @@ interface model { /// typescript-import="@temporalio/common" type workflow-id-conflict-policy = placeholder; + /// @nexus.doc "Static metadata for a workflow execution." /// @nexus.proto "temporal.api.sdk.v1.UserMetadata" typescript-import="@temporalio/proto" /// @nexus.flatten-in-api + /// @nexus.experimental record user-metadata { /// @nexus.doc "Single-line fixed summary for the workflow execution that may appear in UI and CLI. This can be in single-line Temporal Markdown format." /// @nexus.proto-field "summary" diff --git a/crates/protos/protos/api_upstream/nexus/workflow-service.wit b/crates/protos/protos/api_upstream/nexus/workflow-service.wit index dd87646d0..caf3842a3 100644 --- a/crates/protos/protos/api_upstream/nexus/workflow-service.wit +++ b/crates/protos/protos/api_upstream/nexus/workflow-service.wit @@ -13,6 +13,7 @@ world system { interface workflow-service { use nexus:temporal-types/model@1.0.0.{ duration, + header, memo, payloads, placeholder, @@ -36,9 +37,10 @@ interface workflow-service { /// python="Workflow type name or callable identifying the workflow to start." /// typescript="Workflow type name or workflow function identifying the workflow to start." /// dotnet="Workflow type name or workflow expression identifying the workflow to start." + /// go="Workflow function identifying the workflow to start." /// @nexus.proto-field "workflow_type" workflow: workflow-function, - /// @nexus.doc "Unique identifier for the workflow execution." + /// @nexus.doc "Unique identifier for the workflow execution. Must be nonempty." /// @nexus.proto-field "workflow_id" id: string, /// @nexus.doc "Task queue to run the workflow on." @@ -49,24 +51,27 @@ interface workflow-service { /// dotnet="Signal name or signal expression to send with the start request." /// @nexus.proto-field "signal_name" signal: signal-function, - /// @nexus.doc "Total workflow execution timeout, including retries and continue-as-new." + /// @nexus.doc "Total workflow execution timeout, including retries and continue-as-new. Defaults to unlimited." /// @nexus.proto-field "workflow_execution_timeout" + /// @nexus.name go="WorkflowExecutionTimeout" execution-timeout: option, - /// @nexus.doc "Timeout of a single workflow run." + /// @nexus.doc "Timeout of a single workflow run. Defaults to the workflow execution timeout." /// @nexus.proto-field "workflow_run_timeout" + /// @nexus.name go="WorkflowRunTimeout" run-timeout: option, - /// @nexus.doc "Timeout of a single workflow task." + /// @nexus.doc "Timeout of a single workflow task. Defaults to 10 seconds." /// @nexus.proto-field "workflow_task_timeout" + /// @nexus.name go="WorkflowTaskTimeout" task-timeout: option, /// @nexus.omit identity: placeholder, - /// @nexus.doc "Request ID used to deduplicate workflow start requests." - request-id: option, + /// @nexus.omit + request-id: placeholder, /// @nexus.doc "Behavior when a closed workflow with the same ID exists. Default is allow-duplicate." /// @nexus.proto-field "workflow_id_reuse_policy" /// @nexus.default "allow-duplicate" id-reuse-policy: workflow-id-reuse-policy, - /// @nexus.doc "Behavior when a workflow is currently running with the same ID. Set to use-existing for idempotent deduplication on workflow ID. Cannot be set if id-reuse-policy is terminate-if-running." + /// @nexus.doc "Behavior when a workflow is currently running with the same ID. Set to use-existing for idempotent deduplication on workflow ID. Cannot be set if id-reuse-policy is terminate-if-running. Defaults to use-existing." /// @nexus.proto-field "workflow_id_conflict_policy" id-conflict-policy: option, /// @nexus.doc "Retry policy for the workflow." @@ -84,23 +89,30 @@ interface workflow-service { /// @nexus.doc "Amount of time to wait before starting the workflow. This does not work with cron-schedule." /// @nexus.proto-field "workflow_start_delay" start-delay: option, + /// @nexus.doc "Static metadata for the workflow execution." user-metadata: option, - /// @nexus.source python="workflow_namespace" typescript="workflowNamespace" dotnet="TemporalWorkflowContext.WorkflowNamespace" + /// @nexus.source python="workflow_namespace()" typescript="workflowNamespace()" go="workflow.GetInfo(ctx).Namespace" dotnet="TemporalWorkflowContext.WorkflowNamespace()" + /// @nexus.doc "Namespace of the workflow execution." namespace: string, /// @nexus.omit control: placeholder, - /// @nexus.omit - header: placeholder, + /// @nexus.api-omit + /// @nexus.doc "Headers for the request." + /// @nexus.proto-field "header" + headers: option
, /// @nexus.omit links: placeholder, /// @nexus.omit time-skipping-config: placeholder, } + /// @nexus.doc "Result of signaling a workflow and starting it if needed." /// @nexus.experimental /// @nexus.proto "temporal.api.workflowservice.v1.SignalWithStartWorkflowExecutionResponse" typescript-import="@temporalio/proto" record signal-with-start-workflow-response { + /// @nexus.doc "Run ID of the started workflow." run-id: option, + /// @nexus.doc "Whether the workflow was started." started: option, /// @nexus.omit signal-link: placeholder, @@ -114,12 +126,14 @@ interface workflow-service { /// @nexus.output-transform /// python-type="temporalio.workflow.ExternalWorkflowHandle[WorkflowResult]" /// python="temporalio.workflow.get_external_workflow_handle(request.id, run_id=result.run_id)" - /// typescript-type="workflow.ExternalWorkflowHandle" + /// typescript-type="ExternalWorkflowHandle" + /// typescript-type-import="../../../../workflow-handle" /// typescript="workflow.getExternalWorkflowHandle(request.id, result.runId ?? undefined)" /// typescript-import="@temporalio/workflow" /// dotnet-type="Temporalio.Workflows.ExternalWorkflowHandle" /// dotnet="Temporalio.Workflows.Workflow.GetExternalWorkflowHandle(request.Id, result.RunId)" /// @nexus.operation name="SignalWithStartWorkflowExecution" + /// @nexus.serialization-context python="signal_with_start_workflow_serialization_context" typescript="signalWithStartWorkflowSerializationContext" dotnet="WorkflowServiceSerializationContexts.SignalWithStartWorkflow" /// @nexus.experimental signal-with-start-workflow: func( request: signal-with-start-workflow-request, diff --git a/crates/protos/protos/api_upstream/openapi/openapiv2.json b/crates/protos/protos/api_upstream/openapi/openapiv2.json index 1437120e5..0b643ceb7 100644 --- a/crates/protos/protos/api_upstream/openapi/openapiv2.json +++ b/crates/protos/protos/api_upstream/openapi/openapiv2.json @@ -1451,7 +1451,7 @@ }, "/api/v1/namespaces/{namespace}/current-deployment/{deployment.seriesName}": { "post": { - "summary": "Sets a deployment as the current deployment for its deployment series. Can optionally update\nthe metadata of the deployment as well.\nExperimental. This API might significantly change or be removed in a future release.\nDeprecated. Replaced by `SetWorkerDeploymentCurrentVersion`.", + "summary": "Sets a deployment as the current deployment for its deployment series. Can optionally update\nthe metadata of the deployment as well.\nDeprecated. Replaced by `SetWorkerDeploymentCurrentVersion`.", "operationId": "SetCurrentDeployment2", "responses": { "200": { @@ -1497,7 +1497,7 @@ }, "/api/v1/namespaces/{namespace}/current-deployment/{seriesName}": { "get": { - "summary": "Returns the current deployment (and its info) for a given deployment series.\nExperimental. This API might significantly change or be removed in a future release.\nDeprecated. Replaced by `current_version` returned by `DescribeWorkerDeployment`.", + "summary": "Returns the current deployment (and its info) for a given deployment series.\nDeprecated. Replaced by `current_version` returned by `DescribeWorkerDeployment`.", "operationId": "GetCurrentDeployment2", "responses": { "200": { @@ -1534,7 +1534,7 @@ }, "/api/v1/namespaces/{namespace}/deployments": { "get": { - "summary": "Lists worker deployments in the namespace. Optionally can filter based on deployment series\nname.\nExperimental. This API might significantly change or be removed in a future release.\nDeprecated. Replaced with `ListWorkerDeployments`.", + "summary": "Lists worker deployments in the namespace. Optionally can filter based on deployment series\nname.\nDeprecated. Replaced with `ListWorkerDeployments`.", "operationId": "ListDeployments2", "responses": { "200": { @@ -1586,7 +1586,7 @@ }, "/api/v1/namespaces/{namespace}/deployments/{deployment.seriesName}/{deployment.buildId}": { "get": { - "summary": "Describes a worker deployment.\nExperimental. This API might significantly change or be removed in a future release.\nDeprecated. Replaced with `DescribeWorkerDeploymentVersion`.", + "summary": "Describes a worker deployment.\nDeprecated. Replaced with `DescribeWorkerDeploymentVersion`.", "operationId": "DescribeDeployment2", "responses": { "200": { @@ -1631,7 +1631,7 @@ }, "/api/v1/namespaces/{namespace}/deployments/{deployment.seriesName}/{deployment.buildId}/reachability": { "get": { - "summary": "Returns the reachability level of a worker deployment to help users decide when it is time\nto decommission a deployment. Reachability level is calculated based on the deployment's\n`status` and existing workflows that depend on the given deployment for their execution.\nCalculating reachability is relatively expensive. Therefore, server might return a recently\ncached value. In such a case, the `last_update_time` will inform you about the actual\nreachability calculation time.\nExperimental. This API might significantly change or be removed in a future release.\nDeprecated. Replaced with `DrainageInfo` returned by `DescribeWorkerDeploymentVersion`.", + "summary": "Returns the reachability level of a worker deployment to help users decide when it is time\nto decommission a deployment. Reachability level is calculated based on the deployment's\n`status` and existing workflows that depend on the given deployment for their execution.\nCalculating reachability is relatively expensive. Therefore, server might return a recently\ncached value. In such a case, the `last_update_time` will inform you about the actual\nreachability calculation time.\nDeprecated. Replaced with `DrainageInfo` returned by `DescribeWorkerDeploymentVersion`.", "operationId": "GetDeploymentReachability2", "responses": { "200": { @@ -2849,7 +2849,7 @@ }, "/api/v1/namespaces/{namespace}/worker-deployment-versions/{deploymentVersion.deploymentName}/{deploymentVersion.buildId}": { "get": { - "summary": "Describes a worker deployment version.\nExperimental. This API might significantly change or be removed in a future release.", + "summary": "Describes a worker deployment version.", "operationId": "DescribeWorkerDeploymentVersion2", "responses": { "200": { @@ -2906,7 +2906,7 @@ ] }, "delete": { - "summary": "Used for manual deletion of Versions. User can delete a Version only when all the\nfollowing conditions are met:\n - It is not the Current or Ramping Version of its Deployment.\n - It has no active pollers (none of the task queues in the Version have pollers)\n - It is not draining (see WorkerDeploymentVersionInfo.drainage_info). This condition\n can be skipped by passing `skip-drainage=true`.\nExperimental. This API might significantly change or be removed in a future release.", + "summary": "Used for manual deletion of Versions. User can delete a Version only when all the\nfollowing conditions are met:\n - It is not the Current or Ramping Version of its Deployment.\n - It has no active pollers (none of the task queues in the Version have pollers)\n - It is not draining (see WorkerDeploymentVersionInfo.drainage_info). This condition\n can be skipped by passing `skip-drainage=true`.", "operationId": "DeleteWorkerDeploymentVersion2", "responses": { "200": { @@ -3025,7 +3025,7 @@ }, "/api/v1/namespaces/{namespace}/worker-deployment-versions/{deploymentVersion.deploymentName}/{deploymentVersion.buildId}/update-metadata": { "post": { - "summary": "Updates the user-given metadata attached to a Worker Deployment Version.\nExperimental. This API might significantly change or be removed in a future release.", + "summary": "Updates the user-given metadata attached to a Worker Deployment Version.", "operationId": "UpdateWorkerDeploymentVersionMetadata2", "responses": { "200": { @@ -3131,7 +3131,7 @@ }, "/api/v1/namespaces/{namespace}/worker-deployments": { "get": { - "summary": "Lists all Worker Deployments that are tracked in the Namespace.\nExperimental. This API might significantly change or be removed in a future release.", + "summary": "Lists all Worker Deployments that are tracked in the Namespace.", "operationId": "ListWorkerDeployments2", "responses": { "200": { @@ -3176,7 +3176,7 @@ }, "/api/v1/namespaces/{namespace}/worker-deployments/{deploymentName}": { "get": { - "summary": "Describes a Worker Deployment.\nExperimental. This API might significantly change or be removed in a future release.", + "summary": "Describes a Worker Deployment.", "operationId": "DescribeWorkerDeployment2", "responses": { "200": { @@ -3211,7 +3211,7 @@ ] }, "delete": { - "summary": "Deletes records of (an old) Deployment. A deployment can only be deleted if\nit has no Version in it.\nExperimental. This API might significantly change or be removed in a future release.", + "summary": "Deletes records of (an old) Deployment. A deployment can only be deleted if\nit has no Version in it.", "operationId": "DeleteWorkerDeployment2", "responses": { "200": { @@ -3300,7 +3300,7 @@ }, "/api/v1/namespaces/{namespace}/worker-deployments/{deploymentName}/set-current-version": { "post": { - "summary": "Set/unset the Current Version of a Worker Deployment. Automatically unsets the Ramping\nVersion if it is the Version being set as Current.\nExperimental. This API might significantly change or be removed in a future release.", + "summary": "Set/unset the Current Version of a Worker Deployment. Automatically unsets the Ramping\nVersion if it is the Version being set as Current.", "operationId": "SetWorkerDeploymentCurrentVersion2", "responses": { "200": { @@ -3390,7 +3390,7 @@ }, "/api/v1/namespaces/{namespace}/worker-deployments/{deploymentName}/set-ramping-version": { "post": { - "summary": "Set/unset the Ramping Version of a Worker Deployment and its ramp percentage. Can be used for\ngradual ramp to unversioned workers too.\nExperimental. This API might significantly change or be removed in a future release.", + "summary": "Set/unset the Ramping Version of a Worker Deployment and its ramp percentage. Can be used for\ngradual ramp to unversioned workers too.", "operationId": "SetWorkerDeploymentRampingVersion2", "responses": { "200": { @@ -7104,7 +7104,7 @@ }, "/namespaces/{namespace}/current-deployment/{deployment.seriesName}": { "post": { - "summary": "Sets a deployment as the current deployment for its deployment series. Can optionally update\nthe metadata of the deployment as well.\nExperimental. This API might significantly change or be removed in a future release.\nDeprecated. Replaced by `SetWorkerDeploymentCurrentVersion`.", + "summary": "Sets a deployment as the current deployment for its deployment series. Can optionally update\nthe metadata of the deployment as well.\nDeprecated. Replaced by `SetWorkerDeploymentCurrentVersion`.", "operationId": "SetCurrentDeployment", "responses": { "200": { @@ -7150,7 +7150,7 @@ }, "/namespaces/{namespace}/current-deployment/{seriesName}": { "get": { - "summary": "Returns the current deployment (and its info) for a given deployment series.\nExperimental. This API might significantly change or be removed in a future release.\nDeprecated. Replaced by `current_version` returned by `DescribeWorkerDeployment`.", + "summary": "Returns the current deployment (and its info) for a given deployment series.\nDeprecated. Replaced by `current_version` returned by `DescribeWorkerDeployment`.", "operationId": "GetCurrentDeployment", "responses": { "200": { @@ -7187,7 +7187,7 @@ }, "/namespaces/{namespace}/deployments": { "get": { - "summary": "Lists worker deployments in the namespace. Optionally can filter based on deployment series\nname.\nExperimental. This API might significantly change or be removed in a future release.\nDeprecated. Replaced with `ListWorkerDeployments`.", + "summary": "Lists worker deployments in the namespace. Optionally can filter based on deployment series\nname.\nDeprecated. Replaced with `ListWorkerDeployments`.", "operationId": "ListDeployments", "responses": { "200": { @@ -7239,7 +7239,7 @@ }, "/namespaces/{namespace}/deployments/{deployment.seriesName}/{deployment.buildId}": { "get": { - "summary": "Describes a worker deployment.\nExperimental. This API might significantly change or be removed in a future release.\nDeprecated. Replaced with `DescribeWorkerDeploymentVersion`.", + "summary": "Describes a worker deployment.\nDeprecated. Replaced with `DescribeWorkerDeploymentVersion`.", "operationId": "DescribeDeployment", "responses": { "200": { @@ -7284,7 +7284,7 @@ }, "/namespaces/{namespace}/deployments/{deployment.seriesName}/{deployment.buildId}/reachability": { "get": { - "summary": "Returns the reachability level of a worker deployment to help users decide when it is time\nto decommission a deployment. Reachability level is calculated based on the deployment's\n`status` and existing workflows that depend on the given deployment for their execution.\nCalculating reachability is relatively expensive. Therefore, server might return a recently\ncached value. In such a case, the `last_update_time` will inform you about the actual\nreachability calculation time.\nExperimental. This API might significantly change or be removed in a future release.\nDeprecated. Replaced with `DrainageInfo` returned by `DescribeWorkerDeploymentVersion`.", + "summary": "Returns the reachability level of a worker deployment to help users decide when it is time\nto decommission a deployment. Reachability level is calculated based on the deployment's\n`status` and existing workflows that depend on the given deployment for their execution.\nCalculating reachability is relatively expensive. Therefore, server might return a recently\ncached value. In such a case, the `last_update_time` will inform you about the actual\nreachability calculation time.\nDeprecated. Replaced with `DrainageInfo` returned by `DescribeWorkerDeploymentVersion`.", "operationId": "GetDeploymentReachability", "responses": { "200": { @@ -8432,7 +8432,7 @@ }, "/namespaces/{namespace}/worker-deployment-versions/{deploymentVersion.deploymentName}/{deploymentVersion.buildId}": { "get": { - "summary": "Describes a worker deployment version.\nExperimental. This API might significantly change or be removed in a future release.", + "summary": "Describes a worker deployment version.", "operationId": "DescribeWorkerDeploymentVersion", "responses": { "200": { @@ -8489,7 +8489,7 @@ ] }, "delete": { - "summary": "Used for manual deletion of Versions. User can delete a Version only when all the\nfollowing conditions are met:\n - It is not the Current or Ramping Version of its Deployment.\n - It has no active pollers (none of the task queues in the Version have pollers)\n - It is not draining (see WorkerDeploymentVersionInfo.drainage_info). This condition\n can be skipped by passing `skip-drainage=true`.\nExperimental. This API might significantly change or be removed in a future release.", + "summary": "Used for manual deletion of Versions. User can delete a Version only when all the\nfollowing conditions are met:\n - It is not the Current or Ramping Version of its Deployment.\n - It has no active pollers (none of the task queues in the Version have pollers)\n - It is not draining (see WorkerDeploymentVersionInfo.drainage_info). This condition\n can be skipped by passing `skip-drainage=true`.", "operationId": "DeleteWorkerDeploymentVersion", "responses": { "200": { @@ -8608,7 +8608,7 @@ }, "/namespaces/{namespace}/worker-deployment-versions/{deploymentVersion.deploymentName}/{deploymentVersion.buildId}/update-metadata": { "post": { - "summary": "Updates the user-given metadata attached to a Worker Deployment Version.\nExperimental. This API might significantly change or be removed in a future release.", + "summary": "Updates the user-given metadata attached to a Worker Deployment Version.", "operationId": "UpdateWorkerDeploymentVersionMetadata", "responses": { "200": { @@ -8714,7 +8714,7 @@ }, "/namespaces/{namespace}/worker-deployments": { "get": { - "summary": "Lists all Worker Deployments that are tracked in the Namespace.\nExperimental. This API might significantly change or be removed in a future release.", + "summary": "Lists all Worker Deployments that are tracked in the Namespace.", "operationId": "ListWorkerDeployments", "responses": { "200": { @@ -8759,7 +8759,7 @@ }, "/namespaces/{namespace}/worker-deployments/{deploymentName}": { "get": { - "summary": "Describes a Worker Deployment.\nExperimental. This API might significantly change or be removed in a future release.", + "summary": "Describes a Worker Deployment.", "operationId": "DescribeWorkerDeployment", "responses": { "200": { @@ -8794,7 +8794,7 @@ ] }, "delete": { - "summary": "Deletes records of (an old) Deployment. A deployment can only be deleted if\nit has no Version in it.\nExperimental. This API might significantly change or be removed in a future release.", + "summary": "Deletes records of (an old) Deployment. A deployment can only be deleted if\nit has no Version in it.", "operationId": "DeleteWorkerDeployment", "responses": { "200": { @@ -8883,7 +8883,7 @@ }, "/namespaces/{namespace}/worker-deployments/{deploymentName}/set-current-version": { "post": { - "summary": "Set/unset the Current Version of a Worker Deployment. Automatically unsets the Ramping\nVersion if it is the Version being set as Current.\nExperimental. This API might significantly change or be removed in a future release.", + "summary": "Set/unset the Current Version of a Worker Deployment. Automatically unsets the Ramping\nVersion if it is the Version being set as Current.", "operationId": "SetWorkerDeploymentCurrentVersion", "responses": { "200": { @@ -8973,7 +8973,7 @@ }, "/namespaces/{namespace}/worker-deployments/{deploymentName}/set-ramping-version": { "post": { - "summary": "Set/unset the Ramping Version of a Worker Deployment and its ramp percentage. Can be used for\ngradual ramp to unversioned workers too.\nExperimental. This API might significantly change or be removed in a future release.", + "summary": "Set/unset the Ramping Version of a Worker Deployment and its ramp percentage. Can be used for\ngradual ramp to unversioned workers too.", "operationId": "SetWorkerDeploymentRampingVersion", "responses": { "200": { @@ -10815,6 +10815,10 @@ "type": "object", "description": "Trigger for when the activity is closed." }, + "CallbackInfoOperationCompleted": { + "type": "object", + "description": "Trigger for when the Nexus operation is completed, covering both success cases as\nwell as any type of failure." + }, "CallbackInfoUpdateWorkflowExecutionCompleted": { "type": "object", "properties": { @@ -10853,7 +10857,32 @@ }, "description": "Header to attach to callback request." } - } + }, + "title": "Nexus callbacks are used to delivery Nexus operation completions, as defined in the Nexus RPC spec: \nhttps://github.com/nexus-rpc/api/blob/main/SPEC.md#callback-urls" + }, + "CallbackNexusHandler": { + "type": "object", + "properties": { + "taskQueueName": { + "type": "string", + "description": "Nexus task queue the Temporal worker is listening on.\n\nNOTE: This is not a temporal.api.taskqueue.v1.TaskQueue to avoid a circular dependency." + }, + "service": { + "type": "string", + "description": "Target Nexus service, e.g. \"HTTPAdapter\"." + }, + "operation": { + "type": "string", + "description": "Target operation, e.g. \"DeliverAsWebhook\"." + }, + "sourceContext": { + "$ref": "#/definitions/v1Payload", + "description": "There are restrictions on the maxium payload size a single callback can carry, as well as the\ntotal sum of all source context payloads attached to an execution. See dynamic configuration:\n\"callback.nexusHandler.sourceContext.maxSize\", \"callback.nexusHandler.sourceContext.aggregateMaxSize\".", + "title": "Arbitrary user-supplied data from the source operation's callsite. (As applicable, not all operations\nsupport attaching context data.)" + } + }, + "description": "The targeted Nexus service must be registered within the same namespace as the source operation\nthe callback is attached to. (While Nexus allows for cross-namespace operations, NexusHandler callbacks\nare strictly caller-side.)\n\nNexusHandler callbacks are only supported for certain types of operations, e.g. standalone Nexus operations.\nAttempting to attach a Worker callback for an unsupported operation will result in an INVALID_ARGUMENT\nerror from the server.", + "title": "NexusHandler callbacks are requests to invoke a specific shape of Nexus operation on a Temporal worker.\nThe specified Nexus operation must have the following:\n- Input: temporal.api.notificationservice.v1.OnCompleteRequest\n- Output: temporal.api.notificationservice.v1.OnCompleteResponse" }, "ComputeStatusProviderValidationStatus": { "type": "object", @@ -11544,7 +11573,7 @@ "type": "array", "items": { "type": "object", - "$ref": "#/definitions/v1Callback" + "$ref": "#/definitions/commonV1Callback" }, "description": "Completion callbacks attached to the running workflow update." } @@ -12225,6 +12254,10 @@ "deploymentOptions": { "$ref": "#/definitions/v1WorkerDeploymentOptions", "description": "Worker deployment options that user has set in the worker." + }, + "cause": { + "$ref": "#/definitions/v1ActivityTaskFailedCause", + "description": "Why did the task fail? When unset, the failure is treated as an unspecified activity failure." } } }, @@ -12250,6 +12283,10 @@ "resourceId": { "type": "string", "description": "Resource ID for routing. Contains \"workflow:workflow_id\" or \"activity:activity_id\" for standalone activities." + }, + "cause": { + "$ref": "#/definitions/v1ActivityTaskFailedCause", + "description": "Why did the activity task fail? Optional; when unset the failure is treated as a normal\nactivity failure. See the type's doc for more." } } }, @@ -12578,7 +12615,7 @@ "type": "array", "items": { "type": "object", - "$ref": "#/definitions/v1Callback" + "$ref": "#/definitions/commonV1Callback" }, "description": "Callbacks to be called by the server when this activity reaches a terminal state.\nCallback addresses must be whitelisted in the server's dynamic configuration." }, @@ -12717,6 +12754,10 @@ "$ref": "#/definitions/v1NexusOperationIdConflictPolicy", "description": "Defines how to resolve an operation id conflict with a *running* operation.\nThe default policy is NEXUS_OPERATION_ID_CONFLICT_POLICY_FAIL." }, + "onConflictOptions": { + "$ref": "#/definitions/apiNexusoperationV1OnConflictOptions", + "description": "Defines actions to be done to the existing running standalone Nexus when the conflict policy\nNEXUS_OPERATION_ID_CONFLICT_POLICY_USE_EXISTING is used. If not set or set to a empty object\n(all options with default value), it will not modify the running operation." + }, "searchAttributes": { "$ref": "#/definitions/v1SearchAttributes", "description": "Search attributes for indexing." @@ -12731,6 +12772,22 @@ "userMetadata": { "$ref": "#/definitions/v1UserMetadata", "description": "Metadata for use by user interfaces to display the fixed as-of-start summary and details of the operation." + }, + "completionCallbacks": { + "type": "array", + "items": { + "type": "object", + "$ref": "#/definitions/commonV1Callback" + }, + "description": "Completion callbacks to be invoked once the Nexus operation reaches a terminal state." + }, + "links": { + "type": "array", + "items": { + "type": "object", + "$ref": "#/definitions/apiCommonV1Link" + }, + "description": "Links to be associated with the Nexus operation. Callbacks may also have associated links;\nlinks already included with a callback should not be duplicated here." } } }, @@ -12811,7 +12868,7 @@ "type": "array", "items": { "type": "object", - "$ref": "#/definitions/v1Callback" + "$ref": "#/definitions/commonV1Callback" }, "description": "Callbacks to be called by the server when this workflow reaches a terminal state.\nIf the workflow continues-as-new, these callbacks will be carried over to the new execution.\nCallback addresses must be whitelisted in the server's dynamic configuration." }, @@ -13355,7 +13412,7 @@ "type": "array", "items": { "type": "object", - "$ref": "#/definitions/v1Callback" + "$ref": "#/definitions/commonV1Callback" }, "description": "Callbacks to be called by the server when this update reaches a terminal state." }, @@ -13455,7 +13512,7 @@ "type": "object", "properties": { "callback": { - "$ref": "#/definitions/v1Callback", + "$ref": "#/definitions/commonV1Callback", "description": "Information on how this callback should be invoked (e.g. its URL and type)." }, "registrationTime": { @@ -13489,6 +13546,10 @@ "blockedReason": { "type": "string", "description": "If the state is BLOCKED, blocked reason provides additional information." + }, + "requestId": { + "type": "string", + "description": "Server-generated request ID used as an idempotency token when invoking callbacks.\nIt has no relation to caller-side request_id sent in operations like StartNexusOperationExecutionRequest." } }, "description": "Common callback information. Specific CallbackInfo messages should embed this and may include additional fields." @@ -13511,11 +13572,51 @@ }, "description": "When starting an execution with a conflict policy that uses an existing execution and there is already an existing\nrunning execution, OnConflictOptions defines actions to be taken on the existing running execution." }, + "apiNexusoperationV1CallbackInfo": { + "type": "object", + "properties": { + "trigger": { + "$ref": "#/definitions/apiNexusoperationV1CallbackInfoTrigger", + "description": "Trigger for this callback." + }, + "info": { + "$ref": "#/definitions/apiCallbackV1CallbackInfo", + "description": "Common callback info." + } + }, + "description": "CallbackInfo contains the state of a callback attached to a standalone Nexus operation." + }, + "apiNexusoperationV1CallbackInfoTrigger": { + "type": "object", + "properties": { + "operationCompleted": { + "$ref": "#/definitions/CallbackInfoOperationCompleted" + } + } + }, + "apiNexusoperationV1OnConflictOptions": { + "type": "object", + "properties": { + "attachRequestId": { + "type": "boolean", + "description": "Attaches the request ID to the running operation." + }, + "attachCompletionCallbacks": { + "type": "boolean", + "description": "Attaches the completion callbacks to the running operation." + }, + "attachLinks": { + "type": "boolean", + "description": "Attaches any new links to the running operation." + } + }, + "description": "When StartNexusOperationExecutionRequest uses the conflict policy NEXUS_OPERATION_ID_CONFLICT_POLICY_USE_EXISTING\nand there is already an existing, running standalone Nexus operation, OnConflictOptions defines actions to be\ntaken." + }, "apiWorkflowV1CallbackInfo": { "type": "object", "properties": { "callback": { - "$ref": "#/definitions/v1Callback", + "$ref": "#/definitions/commonV1Callback", "description": "Information on how this callback should be invoked (e.g. its URL and type)." }, "trigger": { @@ -13552,6 +13653,10 @@ "blockedReason": { "type": "string", "description": "If the state is BLOCKED, blocked reason provides additional information." + }, + "requestId": { + "type": "string", + "description": "Server-generated request ID used as an idempotency token when invoking callbacks.\nIt has no relation to caller-side request_id sent in operations like StartWorkflowExecutionRequest." } }, "description": "CallbackInfo contains the state of an attached workflow callback." @@ -13585,6 +13690,29 @@ }, "description": "When StartWorkflowExecution uses the conflict policy WORKFLOW_ID_CONFLICT_POLICY_USE_EXISTING and\nthere is already an existing running workflow, OnConflictOptions defines actions to be taken on\nthe existing running workflow. In this case, it will create a WorkflowExecutionOptionsUpdatedEvent\nhistory event in the running workflow with the changes requested in this object." }, + "commonV1Callback": { + "type": "object", + "properties": { + "nexus": { + "$ref": "#/definitions/CallbackNexus" + }, + "internal": { + "$ref": "#/definitions/CallbackInternal" + }, + "nexusHandler": { + "$ref": "#/definitions/CallbackNexusHandler" + }, + "links": { + "type": "array", + "items": { + "type": "object", + "$ref": "#/definitions/v1Link" + }, + "description": "Links associated with the callback. It can be used to link to underlying resources of the\ncallback." + } + }, + "description": "Callback to attach to various events in the system, e.g. workflow run completion." + }, "protobufAny": { "type": "object", "properties": { @@ -14046,6 +14174,17 @@ } } }, + "v1ActivityTaskFailedCause": { + "type": "string", + "enum": [ + "ACTIVITY_TASK_FAILED_CAUSE_UNSPECIFIED", + "ACTIVITY_TASK_FAILED_CAUSE_PAYLOADS_TOO_LARGE", + "ACTIVITY_TASK_FAILED_CAUSE_EXTERNAL_STORAGE_FAILURE", + "ACTIVITY_TASK_FAILED_CAUSE_ACTIVITY_WORKER_UNHANDLED_FAILURE" + ], + "default": "ACTIVITY_TASK_FAILED_CAUSE_UNSPECIFIED", + "description": "Activity tasks can fail for various reasons. Note that some of these reasons can only originate\nfrom the server, and some of them can only originate from the SDK/worker.\n\n - ACTIVITY_TASK_FAILED_CAUSE_PAYLOADS_TOO_LARGE: A payload-bearing field on a request the worker sent for this activity task exceeded the\nper-field size limit configured on the server for the namespace.\nCheck the activity task failure message for more information.\n - ACTIVITY_TASK_FAILED_CAUSE_EXTERNAL_STORAGE_FAILURE: The worker failed to offload a payload to, or retrieve one from, external storage while\nprocessing this activity task.\nCheck the activity task failure message for more information.\n - ACTIVITY_TASK_FAILED_CAUSE_ACTIVITY_WORKER_UNHANDLED_FAILURE: The default cause for an activity task failure reported by a worker; a more specific cause\ntakes precedence whenever the condition is recognized.\nCheck the activity task failure message for more information." + }, "v1ActivityTaskFailedEventAttributes": { "type": "object", "properties": { @@ -14073,6 +14212,10 @@ "workerVersion": { "$ref": "#/definitions/v1WorkerVersionStamp", "title": "Version info of the worker who processed this workflow task.\nDeprecated. This field should be cleaned up when versioning-2 API is removed. [cleanup-experimental-wv]" + }, + "cause": { + "$ref": "#/definitions/v1ActivityTaskFailedCause", + "description": "Why did the task fail? When unset, the failure is treated as an unspecified activity failure." } } }, @@ -14663,26 +14806,6 @@ }, "description": "CalendarSpec describes an event specification relative to the calendar,\nsimilar to a traditional cron specification, but with labeled fields. Each\nfield can be one of:\n *: matches always\n x: matches when the field equals x\n x/y : matches when the field equals x+n*y where n is an integer\n x-z: matches when the field is between x and z inclusive\n w,x,y,...: matches when the field is one of the listed values\nEach x, y, z, ... is either a decimal integer, or a month or day of week name\nor abbreviation (in the appropriate fields).\nA timestamp matches if all fields match.\nNote that fields have different default values, for convenience.\nNote that the special case that some cron implementations have for treating\nday_of_month and day_of_week as \"or\" instead of \"and\" when both are set is\nnot implemented.\nday_of_week can accept 0 or 7 as Sunday\nCalendarSpec gets compiled into StructuredCalendarSpec, which is what will be\nreturned if you describe the schedule." }, - "v1Callback": { - "type": "object", - "properties": { - "nexus": { - "$ref": "#/definitions/CallbackNexus" - }, - "internal": { - "$ref": "#/definitions/CallbackInternal" - }, - "links": { - "type": "array", - "items": { - "type": "object", - "$ref": "#/definitions/v1Link" - }, - "description": "Links associated with the callback. It can be used to link to underlying resources of the\ncallback." - } - }, - "description": "Callback to attach to various events in the system, e.g. workflow run completion." - }, "v1CallbackState": { "type": "string", "enum": [ @@ -14695,7 +14818,7 @@ "CALLBACK_STATE_BLOCKED" ], "default": "CALLBACK_STATE_UNSPECIFIED", - "description": "State of a callback.\n\n - CALLBACK_STATE_UNSPECIFIED: Default value, unspecified state.\n - CALLBACK_STATE_STANDBY: Callback is standing by, waiting to be triggered.\n - CALLBACK_STATE_SCHEDULED: Callback is in the queue waiting to be executed or is currently executing.\n - CALLBACK_STATE_BACKING_OFF: Callback has failed with a retryable error and is backing off before the next attempt.\n - CALLBACK_STATE_FAILED: Callback has failed.\n - CALLBACK_STATE_SUCCEEDED: Callback has succeeded.\n - CALLBACK_STATE_BLOCKED: Callback is blocked (eg: by circuit breaker)." + "description": "State of a callback.\n\n - CALLBACK_STATE_UNSPECIFIED: Default value, unspecified state.\n - CALLBACK_STATE_STANDBY: Callback is standing by, waiting to be triggered.\n - CALLBACK_STATE_SCHEDULED: Callback is in the queue waiting to be executed or is currently executing.\n - CALLBACK_STATE_BACKING_OFF: Callback has failed with a retryable error and is backing off before the next attempt.\n - CALLBACK_STATE_FAILED: Callback has failed.\n - CALLBACK_STATE_SUCCEEDED: Callback has succeeded.\n - CALLBACK_STATE_BLOCKED: Callback is blocked, e.g. by circuit breaker." }, "v1CancelExternalWorkflowExecutionFailedCause": { "type": "string", @@ -14929,6 +15052,10 @@ "properties": { "clusterName": { "type": "string" + }, + "replicationRampDuration": { + "type": "string", + "description": "Ramp duration when this cluster is added as passive by UpdateNamespace; unset or non-positive disables gradual connect.\nThis field is not persisted and is omitted from namespace responses." } } }, @@ -15620,6 +15747,14 @@ "type": "string", "format": "byte", "description": "Token for follow-on long-poll requests. Absent only if the operation is complete." + }, + "completionCallbacks": { + "type": "array", + "items": { + "type": "object", + "$ref": "#/definitions/apiNexusoperationV1CallbackInfo" + }, + "description": "Completion callbacks to be invoked once the Nexus operation reaches a terminal state.\nThey will remain in the CALLBACK_STATE_STANDBY state until the Nexus operation is finished." } } }, @@ -16000,17 +16135,18 @@ "type": "string" } }, - "description": "Identifies a specific execution within a namespace. This is used for standalone activities\nexecutions in batch jobs currently." + "description": "Identifies a specific execution within a namespace." }, "v1ExecutionType": { "type": "string", "enum": [ "EXECUTION_TYPE_UNSPECIFIED", "EXECUTION_TYPE_WORKFLOW", - "EXECUTION_TYPE_ACTIVITY" + "EXECUTION_TYPE_ACTIVITY", + "EXECUTION_TYPE_NEXUS_OPERATION" ], "default": "EXECUTION_TYPE_UNSPECIFIED", - "description": " - EXECUTION_TYPE_WORKFLOW: A workflow execution archetype.\n - EXECUTION_TYPE_ACTIVITY: An activity execution archetype. This is reserved for standalone activities." + "description": " - EXECUTION_TYPE_WORKFLOW: A workflow execution archetype.\n - EXECUTION_TYPE_ACTIVITY: An activity execution archetype. This is reserved for standalone activities.\n - EXECUTION_TYPE_NEXUS_OPERATION: A Nexus operation execution archetype. This is reserved for standalone Nexus operations." }, "v1ExternalWorkflowExecutionCancelRequestedEventAttributes": { "type": "object", @@ -16742,10 +16878,36 @@ }, "workflow": { "$ref": "#/definitions/LinkWorkflow" + }, + "callback": { + "$ref": "#/definitions/v1LinkCallback" } }, "description": "Link can be associated with history events. It might contain information about an external entity\nrelated to the history event. For example, workflow A makes a Nexus call that starts workflow B:\nin this case, a history event in workflow A could contain a Link to the workflow started event in\nworkflow B, and vice-versa." }, + "v1LinkCallback": { + "type": "object", + "properties": { + "namespace": { + "type": "string" + }, + "execution": { + "$ref": "#/definitions/v1Execution" + }, + "componentPath": { + "type": "array", + "items": { + "type": "string" + }, + "title": "In most cases, the Execution is sufficient to identify the callback's source. But the callback could have\nbeen attached some child component of that execution. e.g. a workflow update. The component path describes\nthe unique component as applicable, typically ending with a unique ID. e.g. [\"Update\", $workflowUpdateId ]" + }, + "requestId": { + "type": "string", + "description": "Server-generate request ID sent when the callback was dispatched." + } + }, + "description": "A link to a worker callback attached to an execution. An execution (e.g. standalone Nexus operation) can have\nmultiple callbacks attached, and will be differentiated by the request_id used when the callback is invoked." + }, "v1ListActivityExecutionsResponse": { "type": "object", "properties": { @@ -18746,7 +18908,7 @@ "type": "array", "items": { "type": "object", - "$ref": "#/definitions/v1Callback" + "$ref": "#/definitions/commonV1Callback" }, "description": "Callbacks to be called by the server when this update reaches a terminal state." }, @@ -21651,7 +21813,7 @@ "type": "array", "items": { "type": "object", - "$ref": "#/definitions/v1Callback" + "$ref": "#/definitions/commonV1Callback" }, "description": "Completion callbacks attached to the running workflow execution." }, @@ -21867,7 +22029,7 @@ "type": "array", "items": { "type": "object", - "$ref": "#/definitions/v1Callback" + "$ref": "#/definitions/commonV1Callback" }, "description": "Completion callbacks attached when this workflow was started." }, diff --git a/crates/protos/protos/api_upstream/openapi/openapiv3.yaml b/crates/protos/protos/api_upstream/openapi/openapiv3.yaml index 007d77168..2aa7c45bf 100644 --- a/crates/protos/protos/api_upstream/openapi/openapiv3.yaml +++ b/crates/protos/protos/api_upstream/openapi/openapiv3.yaml @@ -1324,7 +1324,6 @@ paths: description: |- Sets a deployment as the current deployment for its deployment series. Can optionally update the metadata of the deployment as well. - Experimental. This API might significantly change or be removed in a future release. Deprecated. Replaced by `SetWorkerDeploymentCurrentVersion`. operationId: SetCurrentDeployment parameters: @@ -1363,7 +1362,6 @@ paths: - WorkflowService description: |- Returns the current deployment (and its info) for a given deployment series. - Experimental. This API might significantly change or be removed in a future release. Deprecated. Replaced by `current_version` returned by `DescribeWorkerDeployment`. operationId: GetCurrentDeployment parameters: @@ -1397,7 +1395,6 @@ paths: description: |- Lists worker deployments in the namespace. Optionally can filter based on deployment series name. - Experimental. This API might significantly change or be removed in a future release. Deprecated. Replaced with `ListWorkerDeployments`. operationId: ListDeployments parameters: @@ -1440,7 +1437,6 @@ paths: - WorkflowService description: |- Describes a worker deployment. - Experimental. This API might significantly change or be removed in a future release. Deprecated. Replaced with `DescribeWorkerDeploymentVersion`. operationId: DescribeDeployment parameters: @@ -1500,7 +1496,6 @@ paths: Calculating reachability is relatively expensive. Therefore, server might return a recently cached value. In such a case, the `last_update_time` will inform you about the actual reachability calculation time. - Experimental. This API might significantly change or be removed in a future release. Deprecated. Replaced with `DrainageInfo` returned by `DescribeWorkerDeploymentVersion`. operationId: GetDeploymentReachability parameters: @@ -2570,9 +2565,7 @@ paths: get: tags: - WorkflowService - description: |- - Describes a worker deployment version. - Experimental. This API might significantly change or be removed in a future release. + description: Describes a worker deployment version. operationId: DescribeWorkerDeploymentVersion parameters: - name: namespace @@ -2637,7 +2630,6 @@ paths: - It has no active pollers (none of the task queues in the Version have pollers) - It is not draining (see WorkerDeploymentVersionInfo.drainage_info). This condition can be skipped by passing `skip-drainage=true`. - Experimental. This API might significantly change or be removed in a future release. operationId: DeleteWorkerDeploymentVersion parameters: - name: namespace @@ -2746,9 +2738,7 @@ paths: : post: tags: - WorkflowService - description: |- - Updates the user-given metadata attached to a Worker Deployment Version. - Experimental. This API might significantly change or be removed in a future release. + description: Updates the user-given metadata attached to a Worker Deployment Version. operationId: UpdateWorkerDeploymentVersionMetadata parameters: - name: namespace @@ -2832,9 +2822,7 @@ paths: get: tags: - WorkflowService - description: |- - Lists all Worker Deployments that are tracked in the Namespace. - Experimental. This API might significantly change or be removed in a future release. + description: Lists all Worker Deployments that are tracked in the Namespace. operationId: ListWorkerDeployments parameters: - name: namespace @@ -2869,9 +2857,7 @@ paths: get: tags: - WorkflowService - description: |- - Describes a Worker Deployment. - Experimental. This API might significantly change or be removed in a future release. + description: Describes a Worker Deployment. operationId: DescribeWorkerDeployment parameters: - name: namespace @@ -2945,7 +2931,6 @@ paths: description: |- Deletes records of (an old) Deployment. A deployment can only be deleted if it has no Version in it. - Experimental. This API might significantly change or be removed in a future release. operationId: DeleteWorkerDeployment parameters: - name: namespace @@ -2983,7 +2968,6 @@ paths: description: |- Set/unset the Current Version of a Worker Deployment. Automatically unsets the Ramping Version if it is the Version being set as Current. - Experimental. This API might significantly change or be removed in a future release. operationId: SetWorkerDeploymentCurrentVersion parameters: - name: namespace @@ -3060,7 +3044,6 @@ paths: description: |- Set/unset the Ramping Version of a Worker Deployment and its ramp percentage. Can be used for gradual ramp to unversioned workers too. - Experimental. This API might significantly change or be removed in a future release. operationId: SetWorkerDeploymentRampingVersion parameters: - name: namespace @@ -6422,7 +6405,6 @@ paths: description: |- Sets a deployment as the current deployment for its deployment series. Can optionally update the metadata of the deployment as well. - Experimental. This API might significantly change or be removed in a future release. Deprecated. Replaced by `SetWorkerDeploymentCurrentVersion`. operationId: SetCurrentDeployment parameters: @@ -6461,7 +6443,6 @@ paths: - WorkflowService description: |- Returns the current deployment (and its info) for a given deployment series. - Experimental. This API might significantly change or be removed in a future release. Deprecated. Replaced by `current_version` returned by `DescribeWorkerDeployment`. operationId: GetCurrentDeployment parameters: @@ -6495,7 +6476,6 @@ paths: description: |- Lists worker deployments in the namespace. Optionally can filter based on deployment series name. - Experimental. This API might significantly change or be removed in a future release. Deprecated. Replaced with `ListWorkerDeployments`. operationId: ListDeployments parameters: @@ -6538,7 +6518,6 @@ paths: - WorkflowService description: |- Describes a worker deployment. - Experimental. This API might significantly change or be removed in a future release. Deprecated. Replaced with `DescribeWorkerDeploymentVersion`. operationId: DescribeDeployment parameters: @@ -6598,7 +6577,6 @@ paths: Calculating reachability is relatively expensive. Therefore, server might return a recently cached value. In such a case, the `last_update_time` will inform you about the actual reachability calculation time. - Experimental. This API might significantly change or be removed in a future release. Deprecated. Replaced with `DrainageInfo` returned by `DescribeWorkerDeploymentVersion`. operationId: GetDeploymentReachability parameters: @@ -7610,9 +7588,7 @@ paths: get: tags: - WorkflowService - description: |- - Describes a worker deployment version. - Experimental. This API might significantly change or be removed in a future release. + description: Describes a worker deployment version. operationId: DescribeWorkerDeploymentVersion parameters: - name: namespace @@ -7677,7 +7653,6 @@ paths: - It has no active pollers (none of the task queues in the Version have pollers) - It is not draining (see WorkerDeploymentVersionInfo.drainage_info). This condition can be skipped by passing `skip-drainage=true`. - Experimental. This API might significantly change or be removed in a future release. operationId: DeleteWorkerDeploymentVersion parameters: - name: namespace @@ -7786,9 +7761,7 @@ paths: : post: tags: - WorkflowService - description: |- - Updates the user-given metadata attached to a Worker Deployment Version. - Experimental. This API might significantly change or be removed in a future release. + description: Updates the user-given metadata attached to a Worker Deployment Version. operationId: UpdateWorkerDeploymentVersionMetadata parameters: - name: namespace @@ -7872,9 +7845,7 @@ paths: get: tags: - WorkflowService - description: |- - Lists all Worker Deployments that are tracked in the Namespace. - Experimental. This API might significantly change or be removed in a future release. + description: Lists all Worker Deployments that are tracked in the Namespace. operationId: ListWorkerDeployments parameters: - name: namespace @@ -7909,9 +7880,7 @@ paths: get: tags: - WorkflowService - description: |- - Describes a Worker Deployment. - Experimental. This API might significantly change or be removed in a future release. + description: Describes a Worker Deployment. operationId: DescribeWorkerDeployment parameters: - name: namespace @@ -7985,7 +7954,6 @@ paths: description: |- Deletes records of (an old) Deployment. A deployment can only be deleted if it has no Version in it. - Experimental. This API might significantly change or be removed in a future release. operationId: DeleteWorkerDeployment parameters: - name: namespace @@ -8023,7 +7991,6 @@ paths: description: |- Set/unset the Current Version of a Worker Deployment. Automatically unsets the Ramping Version if it is the Version being set as Current. - Experimental. This API might significantly change or be removed in a future release. operationId: SetWorkerDeploymentCurrentVersion parameters: - name: namespace @@ -8100,7 +8067,6 @@ paths: description: |- Set/unset the Ramping Version of a Worker Deployment and its ramp percentage. Can be used for gradual ramp to unversioned workers too. - Experimental. This API might significantly change or be removed in a future release. operationId: SetWorkerDeploymentRampingVersion parameters: - name: namespace @@ -10227,6 +10193,15 @@ components: description: |- Version info of the worker who processed this workflow task. Deprecated. This field should be cleaned up when versioning-2 API is removed. [cleanup-experimental-wv] + cause: + enum: + - ACTIVITY_TASK_FAILED_CAUSE_UNSPECIFIED + - ACTIVITY_TASK_FAILED_CAUSE_PAYLOADS_TOO_LARGE + - ACTIVITY_TASK_FAILED_CAUSE_EXTERNAL_STORAGE_FAILURE + - ACTIVITY_TASK_FAILED_CAUSE_ACTIVITY_WORKER_UNHANDLED_FAILURE + type: string + description: Why did the task fail? When unset, the failure is treated as an unspecified activity failure. + format: enum ActivityTaskScheduledEventAttributes: type: object properties: @@ -10851,6 +10826,8 @@ components: $ref: '#/components/schemas/Callback_Nexus' internal: $ref: '#/components/schemas/Callback_Internal' + nexusHandler: + $ref: '#/components/schemas/Callback_NexusHandler' links: type: array items: @@ -10903,6 +10880,11 @@ components: blockedReason: type: string description: If the state is BLOCKED, blocked reason provides additional information. + requestId: + type: string + description: |- + Server-generated request ID used as an idempotency token when invoking callbacks. + It has no relation to caller-side request_id sent in operations like StartNexusOperationExecutionRequest. description: Common callback information. Specific CallbackInfo messages should embed this and may include additional fields. Callback_Internal: type: object @@ -10927,6 +10909,45 @@ components: additionalProperties: type: string description: Header to attach to callback request. + description: "Nexus callbacks are used to delivery Nexus operation completions, as defined in the Nexus RPC spec: \n https://github.com/nexus-rpc/api/blob/main/SPEC.md#callback-urls" + Callback_NexusHandler: + type: object + properties: + taskQueueName: + type: string + description: |- + Nexus task queue the Temporal worker is listening on. + + NOTE: This is not a temporal.api.taskqueue.v1.TaskQueue to avoid a circular dependency. + service: + type: string + description: Target Nexus service, e.g. "HTTPAdapter". + operation: + type: string + description: Target operation, e.g. "DeliverAsWebhook". + sourceContext: + allOf: + - $ref: '#/components/schemas/Payload' + description: |- + Arbitrary user-supplied data from the source operation's callsite. (As applicable, not all operations + support attaching context data.) + + There are restrictions on the maxium payload size a single callback can carry, as well as the + total sum of all source context payloads attached to an execution. See dynamic configuration: + "callback.nexusHandler.sourceContext.maxSize", "callback.nexusHandler.sourceContext.aggregateMaxSize". + description: |- + NexusHandler callbacks are requests to invoke a specific shape of Nexus operation on a Temporal worker. + The specified Nexus operation must have the following: + - Input: temporal.api.notificationservice.v1.OnCompleteRequest + - Output: temporal.api.notificationservice.v1.OnCompleteResponse + + The targeted Nexus service must be registered within the same namespace as the source operation + the callback is attached to. (While Nexus allows for cross-namespace operations, NexusHandler callbacks + are strictly caller-side.) + + NexusHandler callbacks are only supported for certain types of operations, e.g. standalone Nexus operations. + Attempting to attach a Worker callback for an unsupported operation will result in an INVALID_ARGUMENT + error from the server. CanceledFailureInfo: type: object properties: @@ -11114,6 +11135,12 @@ components: properties: clusterName: type: string + replicationRampDuration: + pattern: ^-?(?:0|[1-9][0-9]{0,11})(?:\.[0-9]{1,9})?s$ + type: string + description: |- + Ramp duration when this cluster is added as passive by UpdateNamespace; unset or non-positive disables gradual connect. + This field is not persisted and is omitted from namespace responses. CompatibleBuildIdRedirectRule: type: object properties: @@ -11911,6 +11938,13 @@ components: type: string description: Token for follow-on long-poll requests. Absent only if the operation is complete. format: bytes + completionCallbacks: + type: array + items: + $ref: '#/components/schemas/CallbackInfo' + description: |- + Completion callbacks to be invoked once the Nexus operation reaches a terminal state. + They will remain in the CALLBACK_STATE_STANDBY state until the Nexus operation is finished. DescribeScheduleResponse: type: object properties: @@ -12380,15 +12414,14 @@ components: - EXECUTION_TYPE_UNSPECIFIED - EXECUTION_TYPE_WORKFLOW - EXECUTION_TYPE_ACTIVITY + - EXECUTION_TYPE_NEXUS_OPERATION type: string format: enum businessId: type: string runId: type: string - description: |- - Identifies a specific execution within a namespace. This is used for standalone activities - executions in batch jobs currently. + description: Identifies a specific execution within a namespace. ExternalWorkflowExecutionCancelRequestedEventAttributes: type: object properties: @@ -13095,6 +13128,8 @@ components: $ref: '#/components/schemas/Link_NexusOperation' workflow: $ref: '#/components/schemas/Link_Workflow' + callback: + $ref: '#/components/schemas/Link_Callback' description: |- Link can be associated with history events. It might contain information about an external entity related to the history event. For example, workflow A makes a Nexus call that starts workflow B: @@ -13119,6 +13154,27 @@ components: A link to a built-in batch job. Batch jobs can be used to perform operations on a set of workflows (e.g. terminate, signal, cancel, etc). This link can be put on workflow history events generated by actions taken by a batch job. + Link_Callback: + type: object + properties: + namespace: + type: string + execution: + $ref: '#/components/schemas/Execution' + componentPath: + type: array + items: + type: string + description: |- + In most cases, the Execution is sufficient to identify the callback's source. But the callback could have + been attached some child component of that execution. e.g. a workflow update. The component path describes + the unique component as applicable, typically ending with a unique ID. e.g. ["Update", $workflowUpdateId ] + requestId: + type: string + description: Server-generate request ID sent when the callback was dispatched. + description: |- + A link to a worker callback attached to an execution. An execution (e.g. standalone Nexus operation) can have + multiple callbacks attached, and will be differentiated by the request_id used when the callback is invoked. Link_NexusOperation: type: object properties: @@ -15917,6 +15973,17 @@ components: resourceId: type: string description: Resource ID for routing. Contains "workflow:workflow_id" or "activity:activity_id" for standalone activities. + cause: + enum: + - ACTIVITY_TASK_FAILED_CAUSE_UNSPECIFIED + - ACTIVITY_TASK_FAILED_CAUSE_PAYLOADS_TOO_LARGE + - ACTIVITY_TASK_FAILED_CAUSE_EXTERNAL_STORAGE_FAILURE + - ACTIVITY_TASK_FAILED_CAUSE_ACTIVITY_WORKER_UNHANDLED_FAILURE + type: string + description: |- + Why did the activity task fail? Optional; when unset the failure is treated as a normal + activity failure. See the type's doc for more. + format: enum RespondActivityTaskFailedByIdResponse: type: object properties: @@ -15969,6 +16036,15 @@ components: allOf: - $ref: '#/components/schemas/WorkerDeploymentOptions' description: Worker deployment options that user has set in the worker. + cause: + enum: + - ACTIVITY_TASK_FAILED_CAUSE_UNSPECIFIED + - ACTIVITY_TASK_FAILED_CAUSE_PAYLOADS_TOO_LARGE + - ACTIVITY_TASK_FAILED_CAUSE_EXTERNAL_STORAGE_FAILURE + - ACTIVITY_TASK_FAILED_CAUSE_ACTIVITY_WORKER_UNHANDLED_FAILURE + type: string + description: Why did the task fail? When unset, the failure is treated as an unspecified activity failure. + format: enum RespondActivityTaskFailedResponse: type: object properties: @@ -17335,6 +17411,13 @@ components: Defines how to resolve an operation id conflict with a *running* operation. The default policy is NEXUS_OPERATION_ID_CONFLICT_POLICY_FAIL. format: enum + onConflictOptions: + allOf: + - $ref: '#/components/schemas/OnConflictOptions' + description: |- + Defines actions to be done to the existing running standalone Nexus when the conflict policy + NEXUS_OPERATION_ID_CONFLICT_POLICY_USE_EXISTING is used. If not set or set to a empty object + (all options with default value), it will not modify the running operation. searchAttributes: allOf: - $ref: '#/components/schemas/SearchAttributes' @@ -17354,6 +17437,18 @@ components: allOf: - $ref: '#/components/schemas/UserMetadata' description: Metadata for use by user interfaces to display the fixed as-of-start summary and details of the operation. + completionCallbacks: + type: array + items: + $ref: '#/components/schemas/Callback' + description: Completion callbacks to be invoked once the Nexus operation reaches a terminal state. + links: + type: array + items: + $ref: '#/components/schemas/Link' + description: |- + Links to be associated with the Nexus operation. Callbacks may also have associated links; + links already included with a callback should not be duplicated here. StartNexusOperationExecutionResponse: type: object properties: diff --git a/crates/protos/protos/api_upstream/temporal/api/callback/v1/message.proto b/crates/protos/protos/api_upstream/temporal/api/callback/v1/message.proto index f881a4eef..0b14414ad 100644 --- a/crates/protos/protos/api_upstream/temporal/api/callback/v1/message.proto +++ b/crates/protos/protos/api_upstream/temporal/api/callback/v1/message.proto @@ -34,4 +34,8 @@ message CallbackInfo { google.protobuf.Timestamp next_attempt_schedule_time = 7; // If the state is BLOCKED, blocked reason provides additional information. string blocked_reason = 8; -} \ No newline at end of file + + // Server-generated request ID used as an idempotency token when invoking callbacks. + // It has no relation to caller-side request_id sent in operations like StartNexusOperationExecutionRequest. + string request_id = 9; +} diff --git a/crates/protos/protos/api_upstream/temporal/api/common/v1/message.proto b/crates/protos/protos/api_upstream/temporal/api/common/v1/message.proto index 3e2fc0e1c..908c5c80b 100644 --- a/crates/protos/protos/api_upstream/temporal/api/common/v1/message.proto +++ b/crates/protos/protos/api_upstream/temporal/api/common/v1/message.proto @@ -68,8 +68,7 @@ message WorkflowExecution { string run_id = 2; } -// Identifies a specific execution within a namespace. This is used for standalone activities -// executions in batch jobs currently. +// Identifies a specific execution within a namespace. message Execution { temporal.api.enums.v1.ExecutionType type = 1; string business_id = 2; @@ -185,6 +184,8 @@ message ResetOptions { // Callback to attach to various events in the system, e.g. workflow run completion. message Callback { + // Nexus callbacks are used to delivery Nexus operation completions, as defined in the Nexus RPC spec: + // https://github.com/nexus-rpc/api/blob/main/SPEC.md#callback-urls message Nexus { // Callback URL. string url = 1; @@ -201,10 +202,43 @@ message Callback { bytes data = 1; } + // NexusHandler callbacks are requests to invoke a specific shape of Nexus operation on a Temporal worker. + // The specified Nexus operation must have the following: + // - Input: temporal.api.notificationservice.v1.OnCompleteRequest + // - Output: temporal.api.notificationservice.v1.OnCompleteResponse + // + // The targeted Nexus service must be registered within the same namespace as the source operation + // the callback is attached to. (While Nexus allows for cross-namespace operations, NexusHandler callbacks + // are strictly caller-side.) + // + // NexusHandler callbacks are only supported for certain types of operations, e.g. standalone Nexus operations. + // Attempting to attach a Worker callback for an unsupported operation will result in an INVALID_ARGUMENT + // error from the server. + message NexusHandler { + // Nexus task queue the Temporal worker is listening on. + // + // NOTE: This is not a temporal.api.taskqueue.v1.TaskQueue to avoid a circular dependency. + string task_queue_name = 1; + + // Target Nexus service, e.g. "HTTPAdapter". + string service = 2; + // Target operation, e.g. "DeliverAsWebhook". + string operation = 3; + + // Arbitrary user-supplied data from the source operation's callsite. (As applicable, not all operations + // support attaching context data.) + // + // There are restrictions on the maxium payload size a single callback can carry, as well as the + // total sum of all source context payloads attached to an execution. See dynamic configuration: + // "callback.nexusHandler.sourceContext.maxSize", "callback.nexusHandler.sourceContext.aggregateMaxSize". + temporal.api.common.v1.Payload source_context = 4; + } + reserved 1; // For a generic callback mechanism to be added later. oneof variant { Nexus nexus = 2; Internal internal = 3; + NexusHandler nexus_handler = 4; } // Links associated with the callback. It can be used to link to underlying resources of the @@ -273,12 +307,27 @@ message Link { string reason = 4; } + // A link to a worker callback attached to an execution. An execution (e.g. standalone Nexus operation) can have + // multiple callbacks attached, and will be differentiated by the request_id used when the callback is invoked. + message Callback { + string namespace = 1; + Execution execution = 2; + // In most cases, the Execution is sufficient to identify the callback's source. But the callback could have + // been attached some child component of that execution. e.g. a workflow update. The component path describes + // the unique component as applicable, typically ending with a unique ID. e.g. ["Update", $workflowUpdateId ] + repeated string component_path = 3; + + // Server-generate request ID sent when the callback was dispatched. + string request_id = 4; + } + oneof variant { WorkflowEvent workflow_event = 1; BatchJob batch_job = 2; Activity activity = 3; NexusOperation nexus_operation = 4; Workflow workflow = 5; + Callback callback = 6; } } diff --git a/crates/protos/protos/api_upstream/temporal/api/enums/v1/common.proto b/crates/protos/protos/api_upstream/temporal/api/enums/v1/common.proto index cdc387173..ef8381cda 100644 --- a/crates/protos/protos/api_upstream/temporal/api/enums/v1/common.proto +++ b/crates/protos/protos/api_upstream/temporal/api/enums/v1/common.proto @@ -47,7 +47,7 @@ enum CallbackState { CALLBACK_STATE_FAILED = 4; // Callback has succeeded. CALLBACK_STATE_SUCCEEDED = 5; - // Callback is blocked (eg: by circuit breaker). + // Callback is blocked, e.g. by circuit breaker. CALLBACK_STATE_BLOCKED = 6; } @@ -113,4 +113,6 @@ enum ExecutionType { EXECUTION_TYPE_WORKFLOW = 1; // An activity execution archetype. This is reserved for standalone activities. EXECUTION_TYPE_ACTIVITY = 2; -} \ No newline at end of file + // A Nexus operation execution archetype. This is reserved for standalone Nexus operations. + EXECUTION_TYPE_NEXUS_OPERATION = 3; +} diff --git a/crates/protos/protos/api_upstream/temporal/api/enums/v1/failed_cause.proto b/crates/protos/protos/api_upstream/temporal/api/enums/v1/failed_cause.proto index 65987419a..9ff4c9a53 100644 --- a/crates/protos/protos/api_upstream/temporal/api/enums/v1/failed_cause.proto +++ b/crates/protos/protos/api_upstream/temporal/api/enums/v1/failed_cause.proto @@ -105,6 +105,24 @@ enum WorkflowTaskFailedCause { WORKFLOW_TASK_FAILED_CAUSE_BAD_UNSUBSCRIBE_NOTIFICATION_CHANNEL_ATTRIBUTES = 45; } +// Activity tasks can fail for various reasons. Note that some of these reasons can only originate +// from the server, and some of them can only originate from the SDK/worker. +enum ActivityTaskFailedCause { + ACTIVITY_TASK_FAILED_CAUSE_UNSPECIFIED = 0; + // A payload-bearing field on a request the worker sent for this activity task exceeded the + // per-field size limit configured on the server for the namespace. + // Check the activity task failure message for more information. + ACTIVITY_TASK_FAILED_CAUSE_PAYLOADS_TOO_LARGE = 1; + // The worker failed to offload a payload to, or retrieve one from, external storage while + // processing this activity task. + // Check the activity task failure message for more information. + ACTIVITY_TASK_FAILED_CAUSE_EXTERNAL_STORAGE_FAILURE = 2; + // The default cause for an activity task failure reported by a worker; a more specific cause + // takes precedence whenever the condition is recognized. + // Check the activity task failure message for more information. + ACTIVITY_TASK_FAILED_CAUSE_ACTIVITY_WORKER_UNHANDLED_FAILURE = 3; +} + enum StartChildWorkflowExecutionFailedCause { START_CHILD_WORKFLOW_EXECUTION_FAILED_CAUSE_UNSPECIFIED = 0; START_CHILD_WORKFLOW_EXECUTION_FAILED_CAUSE_WORKFLOW_ALREADY_EXISTS = 1; @@ -148,12 +166,14 @@ enum ResourceExhaustedCause { RESOURCE_EXHAUSTED_CAUSE_OPS_LIMIT = 9; // Limits related to Worker Deployments are reached. RESOURCE_EXHAUSTED_CAUSE_WORKER_DEPLOYMENT_LIMITS = 10; + // Namespace exceeds bandwidth limit. + RESOURCE_EXHAUSTED_CAUSE_BANDWIDTH_LIMIT = 11; } enum ResourceExhaustedScope { RESOURCE_EXHAUSTED_SCOPE_UNSPECIFIED = 0; - // Exhausted resource is a system-level resource. - RESOURCE_EXHAUSTED_SCOPE_NAMESPACE = 1; // Exhausted resource is a namespace-level resource. + RESOURCE_EXHAUSTED_SCOPE_NAMESPACE = 1; + // Exhausted resource is a system-level resource. RESOURCE_EXHAUSTED_SCOPE_SYSTEM = 2; } diff --git a/crates/protos/protos/api_upstream/temporal/api/history/v1/message.proto b/crates/protos/protos/api_upstream/temporal/api/history/v1/message.proto index bbb4ff244..31636d401 100644 --- a/crates/protos/protos/api_upstream/temporal/api/history/v1/message.proto +++ b/crates/protos/protos/api_upstream/temporal/api/history/v1/message.proto @@ -530,6 +530,8 @@ message ActivityTaskFailedEventAttributes { // Version info of the worker who processed this workflow task. // Deprecated. This field should be cleaned up when versioning-2 API is removed. [cleanup-experimental-wv] temporal.api.common.v1.WorkerVersionStamp worker_version = 6 [deprecated = true]; + // Why did the task fail? When unset, the failure is treated as an unspecified activity failure. + temporal.api.enums.v1.ActivityTaskFailedCause cause = 7; } message ActivityTaskTimedOutEventAttributes { diff --git a/crates/protos/protos/api_upstream/temporal/api/nexusoperation/v1/message.proto b/crates/protos/protos/api_upstream/temporal/api/nexusoperation/v1/message.proto new file mode 100644 index 000000000..f247013b0 --- /dev/null +++ b/crates/protos/protos/api_upstream/temporal/api/nexusoperation/v1/message.proto @@ -0,0 +1,42 @@ +syntax = "proto3"; + +package temporal.api.nexusoperation.v1; + +option go_package = "go.temporal.io/api/nexusoperation/v1;nexusoperation"; +option java_package = "io.temporal.api.nexusoperation.v1"; +option java_multiple_files = true; +option java_outer_classname = "MessageProto"; +option ruby_package = "Temporalio::Api::NexusOperation::V1"; +option csharp_namespace = "Temporalio.Api.NexusOperation.V1"; + +import "temporal/api/callback/v1/message.proto"; + +// CallbackInfo contains the state of a callback attached to a standalone Nexus operation. +message CallbackInfo { + // Trigger for when the Nexus operation is completed, covering both success cases as + // well as any type of failure. + message OperationCompleted {} + + message Trigger { + oneof variant { + OperationCompleted operation_completed = 1; + } + } + + // Trigger for this callback. + Trigger trigger = 1; + // Common callback info. + temporal.api.callback.v1.CallbackInfo info = 2; +} + +// When StartNexusOperationExecutionRequest uses the conflict policy NEXUS_OPERATION_ID_CONFLICT_POLICY_USE_EXISTING +// and there is already an existing, running standalone Nexus operation, OnConflictOptions defines actions to be +// taken. +message OnConflictOptions { + // Attaches the request ID to the running operation. + bool attach_request_id = 1; + // Attaches the completion callbacks to the running operation. + bool attach_completion_callbacks = 2; + // Attaches any new links to the running operation. + bool attach_links = 3; +} diff --git a/crates/protos/protos/api_upstream/temporal/api/notificationservice/v1/request_response.proto b/crates/protos/protos/api_upstream/temporal/api/notificationservice/v1/request_response.proto new file mode 100644 index 000000000..9a7ac529f --- /dev/null +++ b/crates/protos/protos/api_upstream/temporal/api/notificationservice/v1/request_response.proto @@ -0,0 +1,37 @@ +syntax = "proto3"; + +package temporal.api.notificationservice.v1; + +option go_package = "go.temporal.io/api/notificationservice/v1;notificationservice"; +option java_package = "io.temporal.api.notificationservice.v1"; +option java_multiple_files = true; +option java_outer_classname = "RequestResponseProto"; +option ruby_package = "Temporalio::Api::NotificationService::V1"; +option csharp_namespace = "Temporalio.Api.NotificationService.V1"; + +import "temporal/api/common/v1/message.proto"; +import "temporal/api/failure/v1/message.proto"; + +// OnCompleteRequest is the request type to the NotificationService's OnComplete operation, +// allowing for defining completion handlers for arbitrary operations. +// +// Information about the source operation will be available in the form of a commonpb.Link, +// which will be available separately from this OnCompleteRequest. e.g. a link to the source +// standalone Nexus operation would be found in the nexuspb.StartOperationRequest parameter +// sent to the worker callback. (In addition to this OnCompleteRequest.) +message OnCompleteRequest { + + // The result of the source operation. + oneof result { + // The operation was successful, and resulted in the given payload. + temporal.api.common.v1.Payload success = 1; + // The operation failed. Includes timeout, cancellation, and application errors. + temporal.api.failure.v1.Failure failure = 2; + } + + // User-supplied data which was added to the source invocation. (As applicable.) + temporal.api.common.v1.Payload source_context = 3; +} + +// OnCompleteResponse is the return type of the OnComplete operation. +message OnCompleteResponse {} diff --git a/crates/protos/protos/api_upstream/temporal/api/replication/v1/message.proto b/crates/protos/protos/api_upstream/temporal/api/replication/v1/message.proto index 0c2f614eb..241e93066 100644 --- a/crates/protos/protos/api_upstream/temporal/api/replication/v1/message.proto +++ b/crates/protos/protos/api_upstream/temporal/api/replication/v1/message.proto @@ -9,12 +9,16 @@ option java_outer_classname = "MessageProto"; option ruby_package = "Temporalio::Api::Replication::V1"; option csharp_namespace = "Temporalio.Api.Replication.V1"; +import "google/protobuf/duration.proto"; import "google/protobuf/timestamp.proto"; import "temporal/api/enums/v1/namespace.proto"; message ClusterReplicationConfig { string cluster_name = 1; + // Ramp duration when this cluster is added as passive by UpdateNamespace; unset or non-positive disables gradual connect. + // This field is not persisted and is omitted from namespace responses. + google.protobuf.Duration replication_ramp_duration = 2; } message NamespaceReplicationConfig { diff --git a/crates/protos/protos/api_upstream/temporal/api/workflow/v1/message.proto b/crates/protos/protos/api_upstream/temporal/api/workflow/v1/message.proto index 4fea9a63e..1cb3e5eba 100644 --- a/crates/protos/protos/api_upstream/temporal/api/workflow/v1/message.proto +++ b/crates/protos/protos/api_upstream/temporal/api/workflow/v1/message.proto @@ -497,6 +497,10 @@ message CallbackInfo { // If the state is BLOCKED, blocked reason provides additional information. string blocked_reason = 9; + + // Server-generated request ID used as an idempotency token when invoking callbacks. + // It has no relation to caller-side request_id sent in operations like StartWorkflowExecutionRequest. + string request_id = 10; } // PendingNexusOperationInfo contains the state of a pending Nexus operation. diff --git a/crates/protos/protos/api_upstream/temporal/api/workflowservice/v1/request_response.proto b/crates/protos/protos/api_upstream/temporal/api/workflowservice/v1/request_response.proto index fb70ab180..d69732a71 100644 --- a/crates/protos/protos/api_upstream/temporal/api/workflowservice/v1/request_response.proto +++ b/crates/protos/protos/api_upstream/temporal/api/workflowservice/v1/request_response.proto @@ -47,6 +47,7 @@ import "temporal/api/batch/v1/message.proto"; import "temporal/api/sdk/v1/task_complete_metadata.proto"; import "temporal/api/sdk/v1/user_metadata.proto"; import "temporal/api/nexus/v1/message.proto"; +import "temporal/api/nexusoperation/v1/message.proto"; import "temporal/api/worker/v1/message.proto"; import "google/protobuf/duration.proto"; @@ -748,6 +749,8 @@ message RespondActivityTaskFailedRequest { temporal.api.deployment.v1.Deployment deployment = 7 [deprecated = true]; // Worker deployment options that user has set in the worker. temporal.api.deployment.v1.WorkerDeploymentOptions deployment_options = 8; + // Why did the task fail? When unset, the failure is treated as an unspecified activity failure. + temporal.api.enums.v1.ActivityTaskFailedCause cause = 10; } message RespondActivityTaskFailedResponse { @@ -774,6 +777,9 @@ message RespondActivityTaskFailedByIdRequest { temporal.api.common.v1.Payloads last_heartbeat_details = 7; // Resource ID for routing. Contains "workflow:workflow_id" or "activity:activity_id" for standalone activities. string resource_id = 8; + // Why did the activity task fail? Optional; when unset the failure is treated as a normal + // activity failure. See the type's doc for more. + temporal.api.enums.v1.ActivityTaskFailedCause cause = 9; } message RespondActivityTaskFailedByIdResponse { @@ -3517,6 +3523,10 @@ message StartNexusOperationExecutionRequest { // Defines how to resolve an operation id conflict with a *running* operation. // The default policy is NEXUS_OPERATION_ID_CONFLICT_POLICY_FAIL. temporal.api.enums.v1.NexusOperationIdConflictPolicy id_conflict_policy = 13; + // Defines actions to be done to the existing running standalone Nexus when the conflict policy + // NEXUS_OPERATION_ID_CONFLICT_POLICY_USE_EXISTING is used. If not set or set to a empty object + // (all options with default value), it will not modify the running operation. + temporal.api.nexusoperation.v1.OnConflictOptions on_conflict_options = 17; // Search attributes for indexing. temporal.api.common.v1.SearchAttributes search_attributes = 14; @@ -3529,6 +3539,12 @@ message StartNexusOperationExecutionRequest { map nexus_header = 15; // Metadata for use by user interfaces to display the fixed as-of-start summary and details of the operation. temporal.api.sdk.v1.UserMetadata user_metadata = 16; + + // Completion callbacks to be invoked once the Nexus operation reaches a terminal state. + repeated temporal.api.common.v1.Callback completion_callbacks = 18; + // Links to be associated with the Nexus operation. Callbacks may also have associated links; + // links already included with a callback should not be duplicated here. + repeated temporal.api.common.v1.Link links = 19; } message StartNexusOperationExecutionResponse { @@ -3577,6 +3593,10 @@ message DescribeNexusOperationExecutionResponse { // Token for follow-on long-poll requests. Absent only if the operation is complete. bytes long_poll_token = 6; + + // Completion callbacks to be invoked once the Nexus operation reaches a terminal state. + // They will remain in the CALLBACK_STATE_STANDBY state until the Nexus operation is finished. + repeated temporal.api.nexusoperation.v1.CallbackInfo completion_callbacks = 7; } message PollNexusOperationExecutionRequest { diff --git a/crates/protos/protos/api_upstream/temporal/api/workflowservice/v1/service.proto b/crates/protos/protos/api_upstream/temporal/api/workflowservice/v1/service.proto index 76c43e8a5..b2ec8a5e9 100644 --- a/crates/protos/protos/api_upstream/temporal/api/workflowservice/v1/service.proto +++ b/crates/protos/protos/api_upstream/temporal/api/workflowservice/v1/service.proto @@ -1089,7 +1089,6 @@ service WorkflowService { } // Describes a worker deployment. - // Experimental. This API might significantly change or be removed in a future release. // Deprecated. Replaced with `DescribeWorkerDeploymentVersion`. rpc DescribeDeployment (DescribeDeploymentRequest) returns (DescribeDeploymentResponse) { option (google.api.http) = { @@ -1101,7 +1100,6 @@ service WorkflowService { } // Describes a worker deployment version. - // Experimental. This API might significantly change or be removed in a future release. rpc DescribeWorkerDeploymentVersion (DescribeWorkerDeploymentVersionRequest) returns (DescribeWorkerDeploymentVersionResponse) { option (google.api.http) = { get: "/namespaces/{namespace}/worker-deployment-versions/{deployment_version.deployment_name}/{deployment_version.build_id}" @@ -1117,7 +1115,6 @@ service WorkflowService { // Lists worker deployments in the namespace. Optionally can filter based on deployment series // name. - // Experimental. This API might significantly change or be removed in a future release. // Deprecated. Replaced with `ListWorkerDeployments`. rpc ListDeployments (ListDeploymentsRequest) returns (ListDeploymentsResponse) { option (google.api.http) = { @@ -1134,7 +1131,6 @@ service WorkflowService { // Calculating reachability is relatively expensive. Therefore, server might return a recently // cached value. In such a case, the `last_update_time` will inform you about the actual // reachability calculation time. - // Experimental. This API might significantly change or be removed in a future release. // Deprecated. Replaced with `DrainageInfo` returned by `DescribeWorkerDeploymentVersion`. rpc GetDeploymentReachability (GetDeploymentReachabilityRequest) returns (GetDeploymentReachabilityResponse) { option (google.api.http) = { @@ -1146,7 +1142,6 @@ service WorkflowService { } // Returns the current deployment (and its info) for a given deployment series. - // Experimental. This API might significantly change or be removed in a future release. // Deprecated. Replaced by `current_version` returned by `DescribeWorkerDeployment`. rpc GetCurrentDeployment (GetCurrentDeploymentRequest) returns (GetCurrentDeploymentResponse) { option (google.api.http) = { @@ -1159,7 +1154,6 @@ service WorkflowService { // Sets a deployment as the current deployment for its deployment series. Can optionally update // the metadata of the deployment as well. - // Experimental. This API might significantly change or be removed in a future release. // Deprecated. Replaced by `SetWorkerDeploymentCurrentVersion`. rpc SetCurrentDeployment (SetCurrentDeploymentRequest) returns (SetCurrentDeploymentResponse) { option (google.api.http) = { @@ -1174,7 +1168,6 @@ service WorkflowService { // Set/unset the Current Version of a Worker Deployment. Automatically unsets the Ramping // Version if it is the Version being set as Current. - // Experimental. This API might significantly change or be removed in a future release. rpc SetWorkerDeploymentCurrentVersion (SetWorkerDeploymentCurrentVersionRequest) returns (SetWorkerDeploymentCurrentVersionResponse) { option (google.api.http) = { post: "/namespaces/{namespace}/worker-deployments/{deployment_name}/set-current-version" @@ -1191,7 +1184,6 @@ service WorkflowService { } // Describes a Worker Deployment. - // Experimental. This API might significantly change or be removed in a future release. rpc DescribeWorkerDeployment (DescribeWorkerDeploymentRequest) returns (DescribeWorkerDeploymentResponse) { option (google.api.http) = { get: "/namespaces/{namespace}/worker-deployments/{deployment_name}" @@ -1207,7 +1199,6 @@ service WorkflowService { // Deletes records of (an old) Deployment. A deployment can only be deleted if // it has no Version in it. - // Experimental. This API might significantly change or be removed in a future release. rpc DeleteWorkerDeployment (DeleteWorkerDeploymentRequest) returns (DeleteWorkerDeploymentResponse) { option (google.api.http) = { delete: "/namespaces/{namespace}/worker-deployments/{deployment_name}" @@ -1228,7 +1219,6 @@ service WorkflowService { // - It has no active pollers (none of the task queues in the Version have pollers) // - It is not draining (see WorkerDeploymentVersionInfo.drainage_info). This condition // can be skipped by passing `skip-drainage=true`. - // Experimental. This API might significantly change or be removed in a future release. rpc DeleteWorkerDeploymentVersion (DeleteWorkerDeploymentVersionRequest) returns (DeleteWorkerDeploymentVersionResponse) { option (google.api.http) = { delete: "/namespaces/{namespace}/worker-deployment-versions/{deployment_version.deployment_name}/{deployment_version.build_id}" @@ -1244,7 +1234,6 @@ service WorkflowService { // Set/unset the Ramping Version of a Worker Deployment and its ramp percentage. Can be used for // gradual ramp to unversioned workers too. - // Experimental. This API might significantly change or be removed in a future release. rpc SetWorkerDeploymentRampingVersion (SetWorkerDeploymentRampingVersionRequest) returns (SetWorkerDeploymentRampingVersionResponse) { option (google.api.http) = { post: "/namespaces/{namespace}/worker-deployments/{deployment_name}/set-ramping-version" @@ -1261,7 +1250,6 @@ service WorkflowService { } // Lists all Worker Deployments that are tracked in the Namespace. - // Experimental. This API might significantly change or be removed in a future release. rpc ListWorkerDeployments (ListWorkerDeploymentsRequest) returns (ListWorkerDeploymentsResponse) { option (google.api.http) = { get: "/namespaces/{namespace}/worker-deployments" @@ -1328,7 +1316,6 @@ service WorkflowService { } // Updates the user-given metadata attached to a Worker Deployment Version. - // Experimental. This API might significantly change or be removed in a future release. rpc UpdateWorkerDeploymentVersionMetadata (UpdateWorkerDeploymentVersionMetadataRequest) returns (UpdateWorkerDeploymentVersionMetadataResponse) { option (google.api.http) = { post: "/namespaces/{namespace}/worker-deployment-versions/{deployment_version.deployment_name}/{deployment_version.build_id}/update-metadata" diff --git a/crates/protos/protos/local/temporal/sdk/core/activity_result/activity_result.proto b/crates/protos/protos/local/temporal/sdk/core/activity_result/activity_result.proto index 198d5e7b9..3076ea8c1 100644 --- a/crates/protos/protos/local/temporal/sdk/core/activity_result/activity_result.proto +++ b/crates/protos/protos/local/temporal/sdk/core/activity_result/activity_result.proto @@ -37,6 +37,30 @@ message Success { // Used to report activity failure either when executing or resolving message Failure { temporal.api.failure.v1.Failure failure = 1; + // Only meaningful on ActivityExecutionResult (lang -> core); ignored on ActivityResolution, + // which reuses this message. + ActivityTaskFailedCause cause = 2; +} + +/* + * A well-known condition that caused an activity task to fail. Lang reports one alongside the + * failure so core can categorize activity failures instead of treating them all alike; it becomes + * the `failure_reason` metric label, so the set of values is deliberately small and bounded. + */ +enum ActivityTaskFailedCause { + ACTIVITY_TASK_FAILED_CAUSE_UNSPECIFIED = 0; + // A payload-bearing field on a request the worker sent for this activity task exceeded the + // per-field size limit configured on the server for the namespace. + // Check the activity task failure message for more information. + ACTIVITY_TASK_FAILED_CAUSE_PAYLOADS_TOO_LARGE = 1; + // The worker failed to offload a payload to, or retrieve one from, external storage while + // processing this activity task. + // Check the activity task failure message for more information. + ACTIVITY_TASK_FAILED_CAUSE_EXTERNAL_STORAGE_FAILURE = 2; + // The default cause for an activity task failure reported by a worker; a more specific cause + // takes precedence whenever the condition is recognized. + // Check the activity task failure message for more information. + ACTIVITY_TASK_FAILED_CAUSE_ACTIVITY_WORKER_UNHANDLED_FAILURE = 3; } /* diff --git a/crates/protos/protos/local/temporal/sdk/core/workflow_activation/workflow_activation.proto b/crates/protos/protos/local/temporal/sdk/core/workflow_activation/workflow_activation.proto index d86ff4c3c..03235adff 100644 --- a/crates/protos/protos/local/temporal/sdk/core/workflow_activation/workflow_activation.proto +++ b/crates/protos/protos/local/temporal/sdk/core/workflow_activation/workflow_activation.proto @@ -13,6 +13,7 @@ import "google/protobuf/empty.proto"; import "temporal/api/failure/v1/message.proto"; import "temporal/api/update/v1/message.proto"; import "temporal/api/common/v1/message.proto"; +import "temporal/api/enums/v1/failed_cause.proto"; import "temporal/api/enums/v1/workflow.proto"; import "temporal/api/notification/v1/message.proto"; import "temporal/sdk/core/activity_result/activity_result.proto"; @@ -295,6 +296,10 @@ message InitializeWorkflow { temporal.api.common.v1.WorkflowExecution root_workflow = 24; // Priority of this workflow execution temporal.api.common.v1.Priority priority = 25; + // The run id recorded on the `WORKFLOW_EXECUTION_STARTED` event. Unlike the execution's current + // run id, this value is preserved across workflow resets. Mirrors the `original_execution_run_id` + // field from `WorkflowExecutionStartedEventAttributes`. + string original_execution_run_id = 26; } // Notify a workflow that a timer has fired @@ -384,6 +389,8 @@ message SignalWorkflow { string identity = 3; // Headers attached to the signal map headers = 5; + // Event ID of the `WORKFLOW_EXECUTION_SIGNALED` history event that produced this job. + int64 originating_event_id = 6; } // Inform lang what the result of a call to `patched` or similar API should be -- this is always @@ -399,15 +406,19 @@ message ResolveSignalExternalWorkflow { // If populated, this signal either failed to be sent or was cancelled depending on failure // type / info. temporal.api.failure.v1.Failure failure = 2; + // The server-reported cause when the signal failed. Unspecified when the signal succeeded or + // was cancelled before being sent. + temporal.api.enums.v1.SignalExternalWorkflowExecutionFailedCause cause = 3; } message ResolveRequestCancelExternalWorkflow { // Sequence number as provided by lang in the corresponding // RequestCancelExternalWorkflowExecution command uint32 seq = 1; - // If populated, this signal either failed to be sent or was cancelled depending on failure - // type / info. + // If populated, the cancellation request failed. temporal.api.failure.v1.Failure failure = 2; + // The server-reported cause when the cancellation request failed. + temporal.api.enums.v1.CancelExternalWorkflowExecutionFailedCause cause = 3; } // Lang is requested to invoke an update handler on the workflow. Lang should invoke the update diff --git a/crates/protos/protos/local/temporal/sdk/core/workflow_commands/workflow_commands.proto b/crates/protos/protos/local/temporal/sdk/core/workflow_commands/workflow_commands.proto index cbac35fdc..d73cdfa1c 100644 --- a/crates/protos/protos/local/temporal/sdk/core/workflow_commands/workflow_commands.proto +++ b/crates/protos/protos/local/temporal/sdk/core/workflow_commands/workflow_commands.proto @@ -15,6 +15,7 @@ import "temporal/api/common/v1/message.proto"; import "temporal/api/enums/v1/workflow.proto"; import "temporal/api/failure/v1/message.proto"; import "temporal/api/sdk/v1/user_metadata.proto"; +import "temporal/api/sdk/v1/event_group_marker.proto"; import "temporal/sdk/core/child_workflow/child_workflow.proto"; import "temporal/sdk/core/nexus/nexus.proto"; import "temporal/sdk/core/common/common.proto"; @@ -26,6 +27,11 @@ message WorkflowCommand { // per-command basis where applicable. temporal.api.sdk.v1.UserMetadata user_metadata = 100; + // Event group markers attached to the command. These are forwarded onto + // the corresponding server-side Command, and consequently surfaced on the + // resulting HistoryEvent. See `temporal/api/sdk/v1/event_group_marker.proto`. + repeated temporal.api.sdk.v1.EventGroupMarker event_group_markers = 101; + oneof variant { StartTimer start_timer = 1; ScheduleActivity schedule_activity = 2; @@ -271,6 +277,10 @@ message ScheduleLocalActivity { // confirmed. Lang should default this to `WAIT_CANCELLATION_COMPLETED`, even though proto // will default to `TRY_CANCEL` automatically. ActivityCancellationType cancellation_type = 13; + // If set, the local activity arguments will be included in the resulting marker under the + // `input` key. This is disabled by default to avoid increasing history size unless the lang + // SDK explicitly chooses to expose it. + bool include_arguments_in_marker = 14; } enum ActivityCancellationType { @@ -354,7 +364,9 @@ message ContinueAsNewWorkflowExecution { // Indicate a workflow has completed as cancelled. Generally sent as a response to an activation // containing a cancellation job. -message CancelWorkflowExecution {} +message CancelWorkflowExecution { + temporal.api.common.v1.Payloads details = 1; +} // A request to set/check if a certain patch is present or not message SetPatchMarker { diff --git a/crates/protos/src/lib.rs b/crates/protos/src/lib.rs index 684657469..764bcb542 100644 --- a/crates/protos/src/lib.rs +++ b/crates/protos/src/lib.rs @@ -1,6 +1,10 @@ #![warn(missing_docs)] //! Compiled protobuf definitions for the Temporal Rust SDK. +//! +//! This crate remains on a `0.x` version because generated protobuf messages are not marked +//! `#[non_exhaustive]`. Becauase of this constructing a message with a struct literal is +//! unsupported. pub mod protos; diff --git a/crates/protos/src/protos/mod.rs b/crates/protos/src/protos/mod.rs index 9149b9339..e0ff9a7f1 100644 --- a/crates/protos/src/protos/mod.rs +++ b/crates/protos/src/protos/mod.rs @@ -89,7 +89,7 @@ pub mod coresdk { fn from(v: workflow_command::Variant) -> Self { Self { variant: Some(v), - user_metadata: None, + ..Default::default() } } } @@ -746,6 +746,7 @@ pub mod coresdk { Self { status: Some(aer::Status::Failed(Failure { failure: Some(fail), + cause: ActivityTaskFailedCause::ActivityWorkerUnhandledFailure as i32, })), } } @@ -830,7 +831,11 @@ pub mod coresdk { Self { status: match r { Ok(p) => Some(aer::Status::Completed(Success { result: Some(p) })), - Err(f) => Some(aer::Status::Failed(Failure { failure: Some(f) })), + Err(f) => Some(aer::Status::Failed(Failure { + failure: Some(f), + cause: ActivityTaskFailedCause::ActivityWorkerUnhandledFailure + as i32, + })), }, } } @@ -866,6 +871,7 @@ pub mod coresdk { match self.status { Some(activity_resolution::Status::Failed(Failure { failure: Some(ref f), + .. })) => f.is_timeout(), _ => None, } @@ -1407,13 +1413,16 @@ pub mod coresdk { } } - impl From for SignalWorkflow { - fn from(a: WorkflowExecutionSignaledEventAttributes) -> Self { + impl From<(WorkflowExecutionSignaledEventAttributes, i64)> for SignalWorkflow { + fn from( + (a, originating_event_id): (WorkflowExecutionSignaledEventAttributes, i64), + ) -> Self { Self { signal_name: a.signal_name, input: Vec::from_payloads(a.input), identity: a.identity, headers: a.header.map(Into::into).unwrap_or_default(), + originating_event_id, } } } @@ -1463,6 +1472,7 @@ pub mod coresdk { start_time: Some(start_time), root_workflow: attrs.root_workflow_execution, priority: attrs.priority, + original_execution_run_id: attrs.original_execution_run_id, } } } @@ -2251,9 +2261,9 @@ pub mod temporal { } impl From for command::Attributes { - fn from(_c: workflow_commands::CancelWorkflowExecution) -> Self { + fn from(c: workflow_commands::CancelWorkflowExecution) -> Self { Self::CancelWorkflowExecutionCommandAttributes( - CancelWorkflowExecutionCommandAttributes { details: None }, + CancelWorkflowExecutionCommandAttributes { details: c.details }, ) } } @@ -3126,6 +3136,11 @@ pub mod temporal { } } } + pub mod nexusoperation { + pub mod v1 { + tonic::include_proto!("temporal.api.nexusoperation.v1"); + } + } pub mod nexusservices { pub mod workerservice { pub mod v1 { diff --git a/crates/protos/src/protos/task_token.rs b/crates/protos/src/protos/task_token.rs index 1b7dc036c..4a2a55fee 100644 --- a/crates/protos/src/protos/task_token.rs +++ b/crates/protos/src/protos/task_token.rs @@ -17,9 +17,14 @@ static LOCAL_ACT_TASK_TOKEN_PREFIX: &[u8] = b"local_act_"; serde::Deserialize, )] /// Type-safe wrapper for task token bytes -pub struct TaskToken(pub Vec); +pub struct TaskToken(Vec); impl TaskToken { + /// Consumes this token and returns its underlying bytes. + pub fn into_inner(self) -> Vec { + self.0 + } + /// Task tokens for local activities are always prefixed with a special sigil so they can /// be identified easily pub fn new_local_activity_token(unique_data: impl IntoIterator) -> Self { diff --git a/crates/protos/src/protos/utilities.rs b/crates/protos/src/protos/utilities.rs index 65f5c7950..94a7de306 100644 --- a/crates/protos/src/protos/utilities.rs +++ b/crates/protos/src/protos/utilities.rs @@ -37,6 +37,12 @@ pub fn decode_status_detail(details: &[u8]) -> Option { T::decode(first_detail.value.as_slice()).ok() } +/// Encode a `google.rpc.Status` into the serialized bytes format carried by +/// `grpc-status-details-bin` (as expected by `tonic::Status::with_details`). +pub fn encode_status_details(status: &super::google::rpc::Status) -> Vec { + status.encode_to_vec() +} + /// Given a header map, lowercase all the keys and return it as a new map. /// Any keys that are duplicated after lowercasing will clobber each other in undefined ordering. pub fn normalize_http_headers(headers: HashMap) -> HashMap { diff --git a/crates/sdk-core-c-bridge/Cargo.toml b/crates/sdk-core-c-bridge/Cargo.toml index 6506823b2..91d8b7cf1 100644 --- a/crates/sdk-core-c-bridge/Cargo.toml +++ b/crates/sdk-core-c-bridge/Cargo.toml @@ -41,20 +41,21 @@ xz2 = { version = "0.1", optional = true } [dependencies.temporalio-client] path = "../client" -version = "0.6" +version = "~1.0.0" +features = ["experimental"] [dependencies.temporalio-sdk-core] path = "../sdk-core" -version = "0.6" +version = "=0.9.0" features = ["ephemeral-server", "otel"] [dependencies.temporalio-common] path = "../common" -version = "0.6" +version = "~1.0.0" features = ["core-based-sdk", "otel"] [dev-dependencies] -base64 = "0.22" +base64 = "0.23" futures-util = "0.3" thiserror = { workspace = true } diff --git a/crates/sdk-core-c-bridge/src/client.rs b/crates/sdk-core-c-bridge/src/client.rs index 9776327c5..a0ae53b82 100644 --- a/crates/sdk-core-c-bridge/src/client.rs +++ b/crates/sdk-core-c-bridge/src/client.rs @@ -1559,9 +1559,9 @@ impl From<&ClientDnsLoadBalancingOptions> for temporalio_client::DnsLoadBalancin } } -impl From<&ClientHttpConnectProxyOptions> for temporalio_client::proxy::HttpConnectProxyOptions { +impl From<&ClientHttpConnectProxyOptions> for temporalio_client::HttpConnectProxyOptions { fn from(opts: &ClientHttpConnectProxyOptions) -> Self { - temporalio_client::proxy::HttpConnectProxyOptions::new(opts.target_host.to_string()) + temporalio_client::HttpConnectProxyOptions::new(opts.target_host.to_string()) .maybe_basic_auth(if opts.username.size != 0 && opts.password.size != 0 { Some((opts.username.to_string(), opts.password.to_string())) } else { diff --git a/crates/sdk-core-c-bridge/src/envconfig.rs b/crates/sdk-core-c-bridge/src/envconfig.rs index 48c611d03..aa52e65fb 100644 --- a/crates/sdk-core-c-bridge/src/envconfig.rs +++ b/crates/sdk-core-c-bridge/src/envconfig.rs @@ -56,11 +56,17 @@ struct ClientEnvConfig { profiles: HashMap, } -impl From for ClientEnvConfig { - fn from(c: CoreClientConfig) -> Self { - Self { - profiles: c.profiles.into_iter().map(|(k, v)| (k, v.into())).collect(), - } +impl TryFrom for ClientEnvConfig { + type Error = String; + + fn try_from(c: CoreClientConfig) -> Result { + Ok(Self { + profiles: c + .profiles + .into_iter() + .map(|(name, profile)| Ok((name, profile.try_into()?))) + .collect::>()?, + }) } } @@ -80,16 +86,18 @@ struct ClientEnvConfigProfile { grpc_meta: HashMap, } -impl From for ClientEnvConfigProfile { - fn from(c: CoreClientConfigProfile) -> Self { - Self { +impl TryFrom for ClientEnvConfigProfile { + type Error = String; + + fn try_from(c: CoreClientConfigProfile) -> Result { + Ok(Self { address: c.address, namespace: c.namespace, api_key: c.api_key, - tls: c.tls.map(Into::into), + tls: c.tls.map(TryInto::try_into).transpose()?, codec: c.codec.map(Into::into), grpc_meta: c.grpc_meta, - } + }) } } @@ -107,15 +115,17 @@ struct ClientEnvConfigTLS { client_key: Option, } -impl From for ClientEnvConfigTLS { - fn from(c: CoreClientConfigTLS) -> Self { - Self { +impl TryFrom for ClientEnvConfigTLS { + type Error = String; + + fn try_from(c: CoreClientConfigTLS) -> Result { + Ok(Self { disabled: c.disabled, server_name: c.server_name, - server_ca_cert: c.server_ca_cert.map(Into::into), - client_cert: c.client_cert.map(Into::into), - client_key: c.client_key.map(Into::into), - } + server_ca_cert: c.server_ca_cert.map(TryInto::try_into).transpose()?, + client_cert: c.client_cert.map(TryInto::try_into).transpose()?, + client_key: c.client_key.map(TryInto::try_into).transpose()?, + }) } } @@ -144,9 +154,11 @@ struct DataSource { data: Option>, } -impl From for DataSource { - fn from(c: CoreDataSource) -> Self { - match c { +impl TryFrom for DataSource { + type Error = String; + + fn try_from(c: CoreDataSource) -> Result { + Ok(match c { CoreDataSource::Path(p) => Self { path: Some(p), data: None, @@ -155,7 +167,8 @@ impl From for DataSource { path: None, data: Some(d), }, - } + _ => return Err("Unsupported envconfig data source".to_string()), + }) } } @@ -227,7 +240,7 @@ pub extern "C" fn temporal_core_client_env_config_load( let core_config = envconfig::load_client_config(load_options, env_vars_map.as_ref()) .map_err(|e| e.to_string())?; - Ok(core_config.into()) + core_config.try_into() }; match result() { @@ -289,7 +302,7 @@ pub extern "C" fn temporal_core_client_env_config_profile_load( let profile = envconfig::load_client_config_profile(load_options, env_vars_map.as_ref()) .map_err(|e| e.to_string())?; - Ok(profile.into()) + profile.try_into() }; match result() { diff --git a/crates/sdk-core-c-bridge/src/tests/mod.rs b/crates/sdk-core-c-bridge/src/tests/mod.rs index eeb118f81..c4554d500 100644 --- a/crates/sdk-core-c-bridge/src/tests/mod.rs +++ b/crates/sdk-core-c-bridge/src/tests/mod.rs @@ -17,7 +17,7 @@ use context::Context; use prost::Message; use std::{ collections::HashMap, - sync::{Arc, LazyLock, Mutex}, + sync::{Arc, LazyLock, Mutex, PoisonError}, }; use temporalio_common::protos::temporal::api::{ failure::v1::Failure, @@ -30,8 +30,13 @@ use temporalio_common::protos::temporal::api::{ mod context; mod utils; +static EPHEMERAL_SERVER_TEST_MUTEX: Mutex<()> = Mutex::new(()); + #[test] fn test_get_system_info() { + let _server_guard = EPHEMERAL_SERVER_TEST_MUTEX + .lock() + .unwrap_or_else(PoisonError::into_inner); Context::with(|context| { context.runtime_new().unwrap(); context @@ -82,6 +87,9 @@ fn rpc_call_exists(context: &Arc, service: RpcService, rpc: &str) -> bo #[test] fn test_missing_rpc_call_has_expected_error_message() { + let _server_guard = EPHEMERAL_SERVER_TEST_MUTEX + .lock() + .unwrap_or_else(PoisonError::into_inner); Context::with(|context| { context.runtime_new().unwrap(); context @@ -146,6 +154,9 @@ fn all_rpc_calls_exist(context: &Arc, service: RpcService, proto: &str) #[test] fn test_all_rpc_calls_exist() { + let _server_guard = EPHEMERAL_SERVER_TEST_MUTEX + .lock() + .unwrap_or_else(PoisonError::into_inner); Context::with(|context| { context.runtime_new().unwrap(); context diff --git a/crates/sdk-core-c-bridge/src/worker.rs b/crates/sdk-core-c-bridge/src/worker.rs index 7130d5f38..328845562 100644 --- a/crates/sdk-core-c-bridge/src/worker.rs +++ b/crates/sdk-core-c-bridge/src/worker.rs @@ -1198,10 +1198,10 @@ impl TryFrom<&WorkerOptions> for temporalio_sdk_core::WorkerConfig { }; temporalio_sdk_core::WorkerVersioningStrategy::WorkerDeploymentBased( temporalio_common::worker::WorkerDeploymentOptions::new( - temporalio_common::worker::WorkerDeploymentVersion { - deployment_name: dopts.version.deployment_name.to_string(), - build_id: dopts.version.build_id.to_string(), - }, + temporalio_common::worker::WorkerDeploymentVersion::builder() + .deployment_name(dopts.version.deployment_name.to_string()) + .build_id(dopts.version.build_id.to_string()) + .build(), ) .use_worker_versioning(dopts.use_worker_versioning) .maybe_default_versioning_behavior(dvb) diff --git a/crates/sdk-core/CHANGELOG.md b/crates/sdk-core/CHANGELOG.md index 30c7c9c2a..a4557e377 100644 --- a/crates/sdk-core/CHANGELOG.md +++ b/crates/sdk-core/CHANGELOG.md @@ -33,19 +33,60 @@ relevant information. ## Unreleased +### Fixed +* Task-poll targets no longer decrease after cancelled or timed-out polls. Affected pollers still + retain their slot during backoff, while resource-exhaustion errors still reduce the target. +* Workflow poll balancing now lets non-sticky pollers use capacity after sticky pollers reach their + configured or autoscaled polling limit. +* Every path that fails a workflow task now only reports the failure to server + on the task's first attempt, and later attempts are left to time out. Previously `PayloadsTooLarge` + failures and history fetch failures were re-reported on every attempt. +* The `workflow_task_execution_failed` metric is now recorded for every failed workflow task + attempt, including attempts whose failure was not sent to the server, and its `failure_reason` + tag distinguishes `GrpcMessageTooLarge`, `PayloadsTooLarge`, and `RequestTooLarge` on every path. + +## [0.9.0] - 2026-09-04 + +## [0.8.0] - 2026-09-02 + ### Added * Added the Core protocol for replay-safe Workflow-originated external stream output, including exact Workflow Task History floors, compact staged-output marker proofs, and shared input/output replay segmentation. +* External workflow signal and cancellation resolution activations now include the typed server + failure cause alongside the existing failure. +* Language SDKs can opt in to recording local activity arguments in the local activity marker's + `input` detail. * Core console logs can now be emitted as newline-delimited JSON when an SDK selects the JSON log format. Configured log filters continue to apply to JSON output. +* Workflow completion-as-cancelled commands can now carry details for recording on the terminal + history event. * Worker heartbeats now report the SDK runtime, hosting environments, operating system, and architecture once per worker, retrying until the first successful delivery. Runtime options can disable the reporting. * Workers now log a `[TMPRL1104]` warning when a workflow task takes longer than 5 seconds. Set `TEMPORAL_WORKFLOW_TASK_DURATION_WARN_SECONDS` to change the threshold. +* Core now supports attaching `EventGroupMarker`s to most workflow commands. +* The `temporal_activity_execution_failed` and `temporal_local_activity_execution_failed` worker + metrics now carry a `failure_reason` attribute. Each is now split into one time series per + reason, which may affect existing dashboards. +* Workflow task completions larger than the gRPC request size limit are now paginated automatically when the namespace supports it. Paginated workflow task completions require Temporal Server 1.32.0 or later. ### Breaking Changes :boom: +* The following types are now non-exhaustive: `Priority`, `WorkerDeploymentVersion`, + `WorkerCallbacks`, `WorkflowExecutionInfo`, `ActivityCloseTimeouts`, + `ActivityExecutionDecodeHint`, child-workflow and signal decode hints, + `SerializationContext`, `SerializationContextData`, `PayloadConverter`, `IncomingError`, + `ScheduleSpec`, and `ScheduleOverlapPolicy`. Construct structs using their respective builders + or constructors (`WorkerCallbacks::new`, `ActivityExecutionDecodeHint::new`, or + `SerializationContext::new`); use `Default` for `PayloadConverter`; and add wildcard branches + when matching enums. +* Renamed `ActivityCloseTimeouts::Both` to `ActivityCloseTimeouts::ScheduleAndStartToClose`. +* Removed the unused `ActExitValue` type. Use `ActivityError::WillCompleteAsync` to mark an + activity for asynchronous completion. +* Removed the test-only `FailOnNondeterminismInterceptor` from the public API. +* `TaskToken` no longer exposes its underlying bytes directly. Use `TaskToken::into_inner()` to + consume a token into its bytes. * Activity failures now include the latest heartbeat details atomically instead of force-flushing a throttled heartbeat first. Temporal Server 1.16.0 or newer is required to guarantee those details are preserved on failure; workers warn when the server does not advertise support. @@ -64,6 +105,23 @@ relevant information. can only be honored while a task is held. * Workers now defensively buffer a replacement workflow task if it reaches a run that still owns one, preserving the outstanding task token in release builds. +* Worker shutdown now drains activity completions that are still flushing their result to the + server before finishing. Previously such a completion — typically one whose final heartbeat RPC + was still in flight — could be permanently stranded by shutdown: the activity's result was + never reported (the server had to time the attempt out before retrying it), and workers missed + shutdown's slot-permit release deadline, panicking in debug builds. +* The Prometheus exporter now appends `_total` to counter metric names when an SDK enables the + counter suffix option. +* Update-with-start `ExecuteMultiOperation` calls now use Core's long-poll timeout instead of the + normal RPC timeout, avoiding premature failures while waiting for an update to reach its + requested stage. +* An activity failure caused by oversized final heartbeat details is now counted in the + `temporal_activity_execution_failed` metric as `failure_reason="PayloadsTooLarge"`. Previously it + was counted under the reason for the failure the activity itself reported, and was not counted at + all when that failure was benign, even though a payload-limit failure was reported instead. +* Workers now warn when autoscaling task polling encounters errors continuously for one minute. + Repeated warnings use exponential backoff up to 15-minute intervals and stop after polling + recovers. * Workers no longer send worker heartbeats or appear in centralized heartbeat reports before they begin polling. * Ephemeral server processes no longer leak on failed start. @@ -78,3 +136,8 @@ relevant information. already elapsed, a sub-millisecond unit, or a multi-unit value like `1m30s`. Previously such a header was ignored entirely, so the handler was never told the task had timed out, and a task left unanswered could block worker shutdown indefinitely. +* Workers with a small workflow cache no longer briefly stop accepting new workflows. Sticky + workflow-task pollers could consume every workflow-cache permit and starve the non-sticky poller, + so the worker would stop picking up new workflows until a poll timed out (up to ~60s). The poll + balancer now reserves a non-sticky slot against the workflow cache size rather than the slot + supplier size. diff --git a/crates/sdk-core/Cargo.toml b/crates/sdk-core/Cargo.toml index 31ff1c508..1235bf286 100644 --- a/crates/sdk-core/Cargo.toml +++ b/crates/sdk-core/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "temporalio-sdk-core" -version = "0.6.0" +version = "0.9.0" authors = ["Temporal Technologies Inc. "] edition = "2024" license-file = { workspace = true } @@ -47,6 +47,8 @@ debug-plugin = ["dep:reqwest", "dep:hyper"] test-utilities = ["dep:assert_matches", "dep:bimap"] antithesis_assertions = ["dep:antithesis_sdk"] dynamic-tls = ["temporalio-client/dynamic-tls"] +# Used by the integration-test runner to make Cloud exclusions native libtest ignores. +cloud-test-mode = [] [dependencies] anyhow = "1.0" @@ -131,19 +133,19 @@ zip = { version = "8.4", optional = true, default-features = false, features = [ # 1st party local deps [dependencies.temporalio-common] path = "../common" -version = "0.6" +version = "~1.0.0" default-features = false features = ["core-telemetry-bridge", "serde_serialize"] [dependencies.temporalio-client] path = "../client" -version = "0.6" +version = "~1.0.0" default-features = false features = ["core-based-sdk"] [dependencies.temporalio-macros] path = "../macros" -version = "0.6" +version = "~1.0.0" [dev-dependencies] assert_matches = "1.5" @@ -163,9 +165,16 @@ hyper-util = { version = "0.1", features = [ ] } rstest = "0.26" semver = "1.0" -temporalio-sdk = { path = "../sdk", features = ["wasm-workflows"] } -temporalio-common = { path = "../common", version = "0.6", default-features = false } -temporalio-workflow = { path = "../workflow" } +# A registry version here creates a cycle when Cargo orders sdk-core before sdk for workspace publishing. +temporalio-sdk = { path = "../sdk", features = [ + "experimental", + "testing", + "wasm-workflows", +] } +temporalio-common = { path = "../common", version = "~1.0.0", default-features = false } +temporalio-workflow = { path = "../workflow", version = "~1.0.0", features = [ + "experimental", +] } tokio = { version = "1.47", default-features = false, features = [ "rt", "rt-multi-thread", @@ -186,10 +195,6 @@ name = "histfetch" path = "src/histfetch.rs" required-features = ["test-utilities"] -[[bin]] -name = "changelog-release-notes" -path = "src/changelog_release_notes.rs" - [[test]] name = "integ_tests" path = "tests/main.rs" diff --git a/crates/sdk-core/README.md b/crates/sdk-core/README.md new file mode 100644 index 000000000..18528a6f9 --- /dev/null +++ b/crates/sdk-core/README.md @@ -0,0 +1,11 @@ +# `temporalio-sdk-core` + +[![crates.io](https://img.shields.io/crates/v/temporalio-sdk-core.svg)](https://crates.io/crates/temporalio-sdk-core) +[![docs.rs](https://docs.rs/temporalio-sdk-core/badge.svg)](https://docs.rs/temporalio-sdk-core) + +Part of [Temporal](https://temporal.io)'s [Rust SDK](https://github.com/temporalio/sdk-rust). + +The shared worker runtime used to build Temporal language SDKs. + +Its APIs are intended for SDK implementers and may change without notice. Rust application authors +should use [`temporalio-sdk`](https://crates.io/crates/temporalio-sdk) instead. diff --git a/crates/sdk-core/src/abstractions.rs b/crates/sdk-core/src/abstractions.rs index 204f2e025..d08186ec5 100644 --- a/crates/sdk-core/src/abstractions.rs +++ b/crates/sdk-core/src/abstractions.rs @@ -108,6 +108,11 @@ where self.supplier.available_slots() } + /// Hard cap on extant permits (the workflow cache size), if any. + pub(crate) fn max_permits(&self) -> Option { + self.max_permits + } + pub(crate) fn slot_supplier_kind(&self) -> &SlotSupplierKind { &self.slot_supplier_kind } diff --git a/crates/sdk-core/src/core_tests/activity_tasks.rs b/crates/sdk-core/src/core_tests/activity_tasks.rs index 21a563861..85156e1da 100644 --- a/crates/sdk-core/src/core_tests/activity_tasks.rs +++ b/crates/sdk-core/src/core_tests/activity_tasks.rs @@ -1,5 +1,6 @@ use crate::{ - ActivityHeartbeat, CompleteActivityError, Worker, advance_fut, job_assert, prost_dur, + ActivityHeartbeat, CompleteActivityError, PollError, Worker, advance_fut, job_assert, + prost_dur, replay::{TestHistoryBuilder, canned_histories}, test_help::{ FakeWfResponses, MockPollCfg, MockWorkerInputs, MocksHolder, QueueResponse, WorkerExt, @@ -10,15 +11,17 @@ use crate::{ worker::{ PollerBehavior, WorkerVersioningStrategy, client::{ - WorkerClient, WorkerClientBag, - mocks::{mock_manual_worker_client, mock_worker_client}, + MockWorkerClient, WorkerClient, WorkerClientBag, + mocks::{DEFAULT_TEST_CAPABILITIES, mock_manual_worker_client, mock_worker_client}, }, }, }; use futures_util::FutureExt; use itertools::Itertools; use prost::Message; +use rstest::rstest; use std::{ + borrow::Borrow, collections::{HashMap, HashSet, VecDeque, hash_map::Entry}, future, sync::{ @@ -30,6 +33,7 @@ use std::{ use temporalio_client::{ Connection, ConnectionOptions, PayloadErrorLimits, SharedReplaceableClient, callback_based::{CallbackBasedGrpcService, GrpcSuccessResponse}, + worker::ClientWorkerSet, }; use temporalio_common::{ payload_limits::{LimitClass, LimitSeverity, PayloadLimitViolation}, @@ -37,8 +41,8 @@ use temporalio_common::{ coresdk::{ ActivityTaskCompletion, activity_result::{ - ActivityExecutionResult, ActivityResolution, Success, activity_execution_result, - activity_resolution, + self as activity_result, ActivityExecutionResult, ActivityResolution, + ActivityTaskFailedCause, Success, activity_execution_result, activity_resolution, }, activity_task::{ActivityCancelReason, ActivityTask, Cancel, activity_task}, workflow_activation::{ @@ -52,8 +56,8 @@ use temporalio_common::{ }, temporal::api::{ command::v1::{ScheduleActivityTaskCommandAttributes, command::Attributes}, - enums::v1::EventType, - failure::v1::failure::FailureInfo, + enums::v1::{ApplicationErrorCategory, EventType}, + failure::v1::{ApplicationFailureInfo, Failure, failure::FailureInfo}, workflowservice::v1::{ GetSystemInfoResponse, PollActivityTaskQueueResponse, RecordActivityTaskHeartbeatRequest, RecordActivityTaskHeartbeatResponse, @@ -65,7 +69,11 @@ use temporalio_common::{ }, worker::WorkerTaskTypes, }; -use tokio::{join, sync::oneshot, time::sleep}; +use tokio::{ + join, + sync::{Notify, oneshot}, + time::sleep, +}; use tokio_util::sync::CancellationToken; use uuid::Uuid; @@ -661,6 +669,146 @@ async fn can_heartbeat_acts_during_shutdown() { core.drain_activity_poller_and_shutdown().await; } +#[derive(PartialEq)] +enum HungRpc { + Heartbeat, + Fail, + Complete, +} + +// Worker-level regression test for the shutdown/completion race: an activity is completed while +// one of its server calls is still in flight, so the completion has left the outstanding map but +// still owns its slot permit. Shutdown must wait for the completion to finish reporting rather +// than tear down the heartbeat manager under it (which stranded the completion forever) and then +// trip the slot-permit release deadline. The hung-heartbeat case reproduces the original race +// (eviction defers its ack until the in-flight heartbeat returns); the hung-failure and +// hung-success cases pin the same invariant when the completion is parked on the result RPC +// itself. +#[rstest] +#[case::heartbeat_rpc_hangs(HungRpc::Heartbeat)] +#[case::fail_rpc_hangs(HungRpc::Fail)] +#[case::complete_rpc_hangs(HungRpc::Complete)] +#[tokio::test] +async fn worker_shutdown_awaits_activity_completion_flushing_result(#[case] hung: HungRpc) { + let rpc_entered = Arc::new(Notify::new()); + let rpc_release = Arc::new(Notify::new()); + let result_reported = Arc::new(AtomicBool::new(false)); + + let mut mock_client = mock_manual_worker_client(); + if hung == HungRpc::Heartbeat { + let entered = rpc_entered.clone(); + let release = rpc_release.clone(); + mock_client + .expect_record_activity_heartbeat() + .times(1) + .returning(move |_, _| { + let entered = entered.clone(); + let release = release.clone(); + async move { + entered.notify_one(); + release.notified().await; + Ok(RecordActivityTaskHeartbeatResponse::default()) + } + .boxed() + }); + } + let entered = rpc_entered.clone(); + let release = rpc_release.clone(); + let result_reported_clone = result_reported.clone(); + if hung == HungRpc::Complete { + mock_client + .expect_complete_activity_task() + .times(1) + .returning(move |_, _| { + let entered = entered.clone(); + let release = release.clone(); + let result_reported = result_reported_clone.clone(); + async move { + entered.notify_one(); + release.notified().await; + result_reported.store(true, Ordering::SeqCst); + Ok(RespondActivityTaskCompletedResponse::default()) + } + .boxed() + }); + } else { + let hold_fail_rpc = hung == HungRpc::Fail; + mock_client + .expect_fail_activity_task() + .times(1) + .returning(move |_, _, _, _| { + let entered = entered.clone(); + let release = release.clone(); + let result_reported = result_reported_clone.clone(); + async move { + if hold_fail_rpc { + entered.notify_one(); + release.notified().await; + } + result_reported.store(true, Ordering::SeqCst); + Ok(RespondActivityTaskFailedResponse::default()) + } + .boxed() + }); + } + + let core = mock_worker(MocksHolder::from_client_with_activities( + mock_client, + [PollActivityTaskQueueResponse { + task_token: vec![1], + activity_id: "act1".to_string(), + heartbeat_timeout: Some(prost_dur!(from_secs(100))), + ..Default::default() + } + .into()], + )); + + let act = core.poll_activity_task().await.unwrap(); + if hung == HungRpc::Heartbeat { + core.record_activity_heartbeat(ActivityHeartbeat { + task_token: act.task_token.clone(), + details: vec![], + }); + } + + let result = if hung == HungRpc::Complete { + ActivityExecutionResult::ok(vec![1].into()) + } else { + ActivityExecutionResult::fail("retry me".into()) + }; + join!( + async { + core.complete_activity_task(ActivityTaskCompletion { + task_token: act.task_token.clone(), + result: Some(result), + }) + .await + .unwrap(); + }, + async { + // Only begin shutdown once the hung RPC is in flight — for the heartbeat case that + // parks the completion's eviction behind it, the window in which shutdown used to + // slip through. + rpc_entered.notified().await; + core.initiate_shutdown(); + let shutdown_fut = async { + assert_matches!( + core.poll_activity_task().await.unwrap_err(), + PollError::ShutDown + ); + core.shutdown().await; + }; + advance_fut!(shutdown_fut); + rpc_release.notify_one(); + shutdown_fut.await; + assert!( + result_reported.load(Ordering::SeqCst), + "worker shutdown completed before the activity's result was reported to server" + ); + } + ); +} + /// Rapid heartbeats are not force-flushed before failure. The failure request carries the latest /// details atomically instead. #[tokio::test] @@ -679,7 +827,7 @@ async fn complete_act_with_fail_includes_latest_heartbeat() { }) }); mock_client.expect_fail_activity_task().times(1).returning( - move |_, _, last_heartbeat_details| { + move |_, _, _, last_heartbeat_details| { assert_eq!(last_heartbeat_details.unwrap().payloads[0].data, [last_hb]); Ok(RespondActivityTaskFailedResponse::default()) }, @@ -843,7 +991,7 @@ async fn activity_failure_distinguishes_no_heartbeat_from_empty_heartbeat() { Ok(RecordActivityTaskHeartbeatResponse::default()) }); mock_client.expect_fail_activity_task().times(1).returning( - move |_, _, last_heartbeat_details| { + move |_, _, _, last_heartbeat_details| { if explicit_empty_heartbeat { assert_eq!(last_heartbeat_details.unwrap().payloads, []); } else { @@ -958,7 +1106,8 @@ async fn oversized_activity_result_failure_includes_latest_heartbeat() { .times(1) .returning(|_, _| Err(payload_too_large_status())); mock_client.expect_fail_activity_task().times(1).returning( - |_, failure, last_heartbeat_details| { + |_, cause, failure, last_heartbeat_details| { + assert_eq!(cause, ActivityTaskFailedCause::PayloadsTooLarge); assert_payloads_too_large_retryable(&failure); assert_eq!(last_heartbeat_details.unwrap().payloads[0].data, [2]); Ok(RespondActivityTaskFailedResponse::default()) @@ -1005,7 +1154,8 @@ async fn oversized_cancel_details_fails_activity() { .times(1) .returning(|_, _| Err(payload_too_large_status())); mock_client.expect_fail_activity_task().times(1).returning( - |_, failure, last_heartbeat_details| { + |_, cause, failure, last_heartbeat_details| { + assert_eq!(cause, ActivityTaskFailedCause::PayloadsTooLarge); assert_payloads_too_large_retryable(&failure); assert_eq!(last_heartbeat_details.unwrap().payloads[0].data, [2]); Ok(RespondActivityTaskFailedResponse::default()) @@ -1040,6 +1190,102 @@ async fn oversized_cancel_details_fails_activity() { core.drain_activity_poller_and_shutdown().await; } +/// Oversized *final* heartbeat details replace whatever lang reported, so the cause and failure sent +/// to the server describe the payload-limit violation rather than the activity's own error, and the +/// oversized details are dropped so the server does not reject the whole request. +#[tokio::test] +async fn oversized_final_heartbeat_details_replace_reported_failure() { + // Manually created because we need non-default payload error limits, and mockall matches + // expectations in creation order, so they cannot be overridden after `mock_worker_client`. + // This will no longer be needed if https://github.com/asomers/mockall/issues/283 is implemented. + let mut mock_client = MockWorkerClient::new(); + let workers = Arc::new(ClientWorkerSet::new()); + mock_client + .expect_payload_error_limits() + .returning(|| Some(PayloadErrorLimits { blob: 10, memo: 10 })); + mock_client + .expect_capabilities() + .returning(|| Some(*DEFAULT_TEST_CAPABILITIES)); + mock_client + .expect_workers() + .returning(move || workers.clone()); + mock_client.expect_is_mock().returning(|| true); + mock_client + .expect_shutdown_worker() + .returning(|_, _, _, _| Ok(ShutdownWorkerResponse {})); + mock_client + .expect_sdk_name_and_version() + .returning(|| ("test-core".to_string(), "0.0.0".to_string())); + mock_client + .expect_identity() + .returning(|| "test-identity".to_string()); + mock_client + .expect_worker_grouping_key() + .returning(Uuid::new_v4); + mock_client + .expect_worker_instance_key() + .returning(Uuid::new_v4); + mock_client + .expect_set_heartbeat_client_fields() + .returning(|_| {}); + mock_client + .expect_record_activity_heartbeat() + .returning(|_, _| Ok(RecordActivityTaskHeartbeatResponse::default())); + mock_client.expect_fail_activity_task().times(1).returning( + |_, cause, failure, last_heartbeat_details| { + assert_eq!(cause, ActivityTaskFailedCause::PayloadsTooLarge); + assert_payloads_too_large_retryable(&failure); + assert!( + last_heartbeat_details.is_none(), + "oversized details must be dropped" + ); + Ok(RespondActivityTaskFailedResponse::default()) + }, + ); + + let core = mock_worker(MocksHolder::from_client_with_activities( + mock_client, + [PollActivityTaskQueueResponse { + task_token: vec![1], + activity_id: "act1".to_string(), + heartbeat_timeout: Some(prost_dur!(from_secs(10))), + ..Default::default() + } + .into()], + )); + + let act = core.poll_activity_task().await.unwrap(); + core.record_activity_heartbeat(ActivityHeartbeat { + task_token: act.task_token.clone(), + details: vec![vec![0_u8; 1024].into()], + }); + // A benign failure is normally not reported as an execution failure at all; what reaches the + // server here is a payload-limit failure instead, which is not benign. + core.complete_activity_task(ActivityTaskCompletion { + task_token: act.task_token, + result: Some(ActivityExecutionResult { + status: Some(activity_execution_result::Status::Failed( + activity_result::Failure { + failure: Some(Failure { + message: "benign".to_string(), + failure_info: Some(FailureInfo::ApplicationFailureInfo( + ApplicationFailureInfo { + category: ApplicationErrorCategory::Benign as i32, + ..Default::default() + }, + )), + ..Default::default() + }), + cause: ActivityTaskFailedCause::ActivityWorkerUnhandledFailure as i32, + }, + )), + }), + }) + .await + .unwrap(); + core.drain_activity_poller_and_shutdown().await; +} + /// An oversized heartbeat `details` payload must fail the activity task (retryably) and stop the /// running activity with a `Cancelled` cancel — replicating the server, which fails the activity /// task and returns `cancel_requested = true`. @@ -1051,7 +1297,8 @@ async fn oversized_heartbeat_fails_activity() { .times(1) .returning(|_, _| Err(payload_too_large_status())); mock_client.expect_fail_activity_task().times(1).returning( - |_, failure, last_heartbeat_details| { + |_, cause, failure, last_heartbeat_details| { + assert_eq!(cause, ActivityTaskFailedCause::PayloadsTooLarge); assert_payloads_too_large_retryable(&failure); assert!(last_heartbeat_details.is_none()); Ok(RespondActivityTaskFailedResponse::default()) @@ -1182,7 +1429,7 @@ async fn no_eager_activities_requested_when_worker_options_disable_it( let mut mock = mock_worker_client(); mock.expect_complete_workflow_task() .times(1) - .returning(move |req| { + .returning(move |req, _| { // Store the number of eager activities requested to be checked below let count = req .commands @@ -1269,7 +1516,7 @@ async fn activity_tasks_from_completion_are_delivered() { let mut mock = mock_worker_client(); mock.expect_complete_workflow_task() .times(1) - .returning(move |req| { + .returning(move |req, _| { // Store the number of eager activities requested to be checked below let count = req .commands @@ -1456,8 +1703,8 @@ async fn graceful_shutdown(#[values(true, false)] at_max_outstanding: bool) { Ok(RecordActivityTaskHeartbeatResponse::default()) }); mock_client.expect_fail_activity_task().times(3).returning( - |task_token, _, last_heartbeat_details| { - if task_token.0 == [1] { + |task_token, _, _, last_heartbeat_details| { + if task_token.borrow() == [1] { assert_eq!(last_heartbeat_details.unwrap().payloads[0].data, [2]); } else { assert!(last_heartbeat_details.is_none()); diff --git a/crates/sdk-core/src/core_tests/event_groups.rs b/crates/sdk-core/src/core_tests/event_groups.rs new file mode 100644 index 000000000..c1c7dc6c9 --- /dev/null +++ b/crates/sdk-core/src/core_tests/event_groups.rs @@ -0,0 +1,399 @@ +use crate::{ + replay::{TestHistoryBuilder, canned_histories, default_act_sched}, + test_help::{MockPollCfg, build_mock_pollers, mock_worker, start_timer_cmd}, +}; +use std::time::Duration; +use temporalio_common::protos::{ + coresdk::{ + AsJsonPayloadExt, + child_workflow::ChildWorkflowCancellationType, + nexus::NexusOperationCancellationType, + workflow_commands::{ + ActivityCancellationType, CancelChildWorkflowExecution, CancelTimer, + CompleteWorkflowExecution, RequestCancelActivity, RequestCancelNexusOperation, + ScheduleActivity, ScheduleNexusOperation, SetPatchMarker, StartChildWorkflowExecution, + WorkflowCommand, workflow_command, + }, + workflow_completion::{WorkflowActivationCompletion, workflow_activation_completion}, + }, + temporal::api::{ + command::v1::Command, + common::v1::Payload, + enums::v1::{CommandType, EventType}, + history::v1::{ + NexusOperationCancelRequestedEventAttributes, NexusOperationCanceledEventAttributes, + NexusOperationScheduledEventAttributes, history_event, + }, + sdk::v1::{ + EventGroupMarker, UserMetadata, + event_group_marker::{Label, Variant}, + }, + }, +}; + +fn plain(cmd: impl Into) -> WorkflowCommand { + cmd.into().into() +} + +/// Tag a command with a marker and a summary both derived from `group`, so that a single name +/// identifies the annotations expected downstream and both fields are checked to travel together. +fn annotate(cmd: impl Into, group: &str) -> WorkflowCommand { + let mut cmd = plain(cmd); + cmd.event_group_markers = vec![EventGroupMarker { + variant: Some(Variant::Label(Label { + id: group.to_string(), + label: Some(group.as_json_payload().unwrap()), + })), + }]; + cmd.user_metadata = Some(UserMetadata { + summary: Some(group.as_json_payload().unwrap()), + details: None, + }); + cmd +} + +#[track_caller] +fn assert_annotated(cmd: &Command, group: &str) { + let expected = annotate(CompleteWorkflowExecution::default(), group); + assert_eq!(cmd.event_group_markers, expected.event_group_markers); + assert_eq!(cmd.user_metadata, expected.user_metadata); +} + +fn complete(run_id: String, cmds: Vec) -> WorkflowActivationCompletion { + WorkflowActivationCompletion { + run_id, + status: Some(workflow_activation_completion::Status::Successful( + cmds.into(), + )), + ..Default::default() + } +} + +#[rstest::rstest] +#[tokio::test] +async fn cancel_timer_command_is_annotated(#[values(false, true)] lang_annotates_cancel: bool) { + let cancelled_timer_seq = 2; + let t = canned_histories::cancel_timer("1", &cancelled_timer_seq.to_string()); + let mut mock_cfg = MockPollCfg::from_hist_builder(t); + let expected_group = if lang_annotates_cancel { + "cancel-group" + } else { + "timer-group" + }; + mock_cfg.completion_asserts_from_expectations(|mut asserts| { + asserts.then(|_| {}).then(move |wft| { + assert_eq!(wft.commands[0].command_type(), CommandType::CancelTimer); + assert_annotated(&wft.commands[0], expected_group); + }); + }); + let mut mock = build_mock_pollers(mock_cfg); + mock.worker_cfg(|wc| wc.max_cached_workflows = 1); + let core = mock_worker(mock); + + let act = core.poll_workflow_activation().await.unwrap(); + core.complete_workflow_activation(complete( + act.run_id, + vec![ + annotate( + start_timer_cmd(cancelled_timer_seq, Duration::from_secs(1)), + "timer-group", + ), + plain(start_timer_cmd(1, Duration::from_secs(1))), + ], + )) + .await + .unwrap(); + + let cancel = CancelTimer { + seq: cancelled_timer_seq, + }; + let cancel = if lang_annotates_cancel { + annotate(cancel, "cancel-group") + } else { + plain(cancel) + }; + let act = core.poll_workflow_activation().await.unwrap(); + core.complete_workflow_activation(complete( + act.run_id, + vec![cancel, plain(CompleteWorkflowExecution::default())], + )) + .await + .unwrap(); +} + +#[rstest::rstest] +#[tokio::test] +async fn cancel_activity_command_is_annotated(#[values(false, true)] lang_annotates_cancel: bool) { + let activity_seq = 1; + let t = canned_histories::cancel_scheduled_activity_with_activity_task_cancel( + "fake_activity", + "signal", + ); + let mut mock_cfg = MockPollCfg::from_hist_builder(t); + let expected_group = if lang_annotates_cancel { + "cancel-group" + } else { + "activity-group" + }; + mock_cfg.completion_asserts_from_expectations(|mut asserts| { + asserts.then(|_| {}).then(move |wft| { + assert_eq!( + wft.commands[0].command_type(), + CommandType::RequestCancelActivityTask + ); + assert_annotated(&wft.commands[0], expected_group); + }); + }); + let mut mock = build_mock_pollers(mock_cfg); + mock.worker_cfg(|wc| wc.max_cached_workflows = 1); + let core = mock_worker(mock); + + let act = core.poll_workflow_activation().await.unwrap(); + core.complete_workflow_activation(complete( + act.run_id, + vec![annotate( + ScheduleActivity { + seq: activity_seq, + activity_id: "fake_activity".to_string(), + cancellation_type: ActivityCancellationType::WaitCancellationCompleted as i32, + ..default_act_sched() + }, + "activity-group", + )], + )) + .await + .unwrap(); + + let cancel = RequestCancelActivity { seq: activity_seq }; + let cancel = if lang_annotates_cancel { + annotate(cancel, "cancel-group") + } else { + plain(cancel) + }; + let act = core.poll_workflow_activation().await.unwrap(); + core.complete_workflow_activation(complete(act.run_id, vec![cancel])) + .await + .unwrap(); + + let act = core.poll_workflow_activation().await.unwrap(); + core.complete_workflow_activation(complete( + act.run_id, + vec![plain(CompleteWorkflowExecution::default())], + )) + .await + .unwrap(); +} + +/// Cancelling a child is doubly indirect: the child machine asks Core to create a whole other +/// machine for the external cancel, and that machine's command is the one the server sees. +#[rstest::rstest] +#[tokio::test] +async fn cancel_child_workflow_command_is_annotated( + #[values(false, true)] lang_annotates_cancel: bool, +) { + let child_wf_id = "child-1"; + let child_seq = 1; + let t = canned_histories::single_child_workflow_try_cancelled(child_wf_id); + let mut mock_cfg = MockPollCfg::from_hist_builder(t); + let expected_group = if lang_annotates_cancel { + "cancel-group" + } else { + "child-group" + }; + mock_cfg.completion_asserts_from_expectations(|mut asserts| { + asserts.then(|_| {}).then(move |wft| { + assert_eq!( + wft.commands[0].command_type(), + CommandType::RequestCancelExternalWorkflowExecution + ); + assert_annotated(&wft.commands[0], expected_group); + }); + }); + let mut mock = build_mock_pollers(mock_cfg); + mock.worker_cfg(|wc| wc.max_cached_workflows = 1); + let core = mock_worker(mock); + + let act = core.poll_workflow_activation().await.unwrap(); + core.complete_workflow_activation(complete( + act.run_id, + vec![annotate( + StartChildWorkflowExecution { + seq: child_seq, + workflow_id: child_wf_id.to_string(), + workflow_type: "child".to_string(), + cancellation_type: ChildWorkflowCancellationType::TryCancel as i32, + ..Default::default() + }, + "child-group", + )], + )) + .await + .unwrap(); + + let cancel = CancelChildWorkflowExecution { + child_workflow_seq: child_seq, + reason: "because".to_string(), + }; + let cancel = if lang_annotates_cancel { + annotate(cancel, "cancel-group") + } else { + plain(cancel) + }; + let act = core.poll_workflow_activation().await.unwrap(); + core.complete_workflow_activation(complete(act.run_id, vec![cancel])) + .await + .unwrap(); + + let act = core.poll_workflow_activation().await.unwrap(); + core.complete_workflow_activation(complete( + act.run_id, + vec![plain(CompleteWorkflowExecution::default())], + )) + .await + .unwrap(); +} + +#[rstest::rstest] +#[tokio::test] +async fn cancel_nexus_operation_command_is_annotated( + #[values(false, true)] lang_annotates_cancel: bool, +) { + let nexus_seq = 1; + let mut t = TestHistoryBuilder::default(); + t.add_by_type(EventType::WorkflowExecutionStarted); + t.add_full_wf_task(); + let scheduled_event_id = t.add(NexusOperationScheduledEventAttributes { + endpoint: "endpoint".to_string(), + service: "service".to_string(), + operation: "operation".to_string(), + ..Default::default() + }); + t.add_we_signaled( + "signal", + vec![Payload { + metadata: Default::default(), + data: b"hello ".to_vec(), + external_payloads: Default::default(), + }], + ); + t.add_full_wf_task(); + t.add( + history_event::Attributes::NexusOperationCancelRequestedEventAttributes( + NexusOperationCancelRequestedEventAttributes { + scheduled_event_id, + ..Default::default() + }, + ), + ); + t.add( + history_event::Attributes::NexusOperationCanceledEventAttributes( + NexusOperationCanceledEventAttributes { + scheduled_event_id, + ..Default::default() + }, + ), + ); + t.add_full_wf_task(); + t.add_workflow_execution_completed(); + + let mut mock_cfg = MockPollCfg::from_hist_builder(t); + let expected_group = if lang_annotates_cancel { + "cancel-group" + } else { + "nexus-group" + }; + mock_cfg.completion_asserts_from_expectations(|mut asserts| { + asserts.then(|_| {}).then(move |wft| { + assert_eq!( + wft.commands[0].command_type(), + CommandType::RequestCancelNexusOperation + ); + assert_annotated(&wft.commands[0], expected_group); + }); + }); + let mut mock = build_mock_pollers(mock_cfg); + mock.worker_cfg(|wc| wc.max_cached_workflows = 1); + let core = mock_worker(mock); + + let act = core.poll_workflow_activation().await.unwrap(); + core.complete_workflow_activation(complete( + act.run_id, + vec![annotate( + ScheduleNexusOperation { + seq: nexus_seq, + endpoint: "endpoint".to_string(), + service: "service".to_string(), + operation: "operation".to_string(), + cancellation_type: NexusOperationCancellationType::WaitCancellationCompleted as i32, + ..Default::default() + }, + "nexus-group", + )], + )) + .await + .unwrap(); + + let cancel = RequestCancelNexusOperation { seq: nexus_seq }; + let cancel = if lang_annotates_cancel { + annotate(cancel, "cancel-group") + } else { + plain(cancel) + }; + let act = core.poll_workflow_activation().await.unwrap(); + core.complete_workflow_activation(complete(act.run_id, vec![cancel])) + .await + .unwrap(); + + let act = core.poll_workflow_activation().await.unwrap(); + core.complete_workflow_activation(complete( + act.run_id, + vec![plain(CompleteWorkflowExecution::default())], + )) + .await + .unwrap(); +} + +/// The `TemporalChangeVersion` upsert exists only to make the patch searchable, so it belongs to +/// the same group as the patch marker rather than to no group at all. +#[tokio::test] +async fn patch_search_attribute_upsert_is_annotated() { + let patch_id = "the-patch"; + let mut t = TestHistoryBuilder::default(); + t.add_by_type(EventType::WorkflowExecutionStarted); + t.add_full_wf_task(); + t.add_has_change_marker(patch_id, false); + t.add_workflow_execution_completed(); + + let mut mock_cfg = MockPollCfg::from_hist_builder(t); + mock_cfg.completion_asserts_from_expectations(|mut asserts| { + asserts.then(|wft| { + assert_eq!(wft.commands[0].command_type(), CommandType::RecordMarker); + assert_annotated(&wft.commands[0], "patch-group"); + assert_eq!( + wft.commands[1].command_type(), + CommandType::UpsertWorkflowSearchAttributes + ); + assert_annotated(&wft.commands[1], "patch-group"); + }); + }); + let mut mock = build_mock_pollers(mock_cfg); + mock.worker_cfg(|wc| wc.max_cached_workflows = 1); + let core = mock_worker(mock); + + let act = core.poll_workflow_activation().await.unwrap(); + core.complete_workflow_activation(complete( + act.run_id, + vec![ + annotate( + SetPatchMarker { + patch_id: patch_id.to_string(), + deprecated: false, + }, + "patch-group", + ), + plain(CompleteWorkflowExecution::default()), + ], + )) + .await + .unwrap(); +} diff --git a/crates/sdk-core/src/core_tests/mod.rs b/crates/sdk-core/src/core_tests/mod.rs index 25385c8bc..0c3b4174f 100644 --- a/crates/sdk-core/src/core_tests/mod.rs +++ b/crates/sdk-core/src/core_tests/mod.rs @@ -1,4 +1,5 @@ mod activity_tasks; +mod event_groups; mod external_streams; mod queries; mod replay_flag; diff --git a/crates/sdk-core/src/core_tests/queries.rs b/crates/sdk-core/src/core_tests/queries.rs index fc70dcd17..a099223b8 100644 --- a/crates/sdk-core/src/core_tests/queries.rs +++ b/crates/sdk-core/src/core_tests/queries.rs @@ -468,7 +468,7 @@ async fn query_cache_miss_causes_page_fetch_dont_reply_wft_too_early( mock_client .expect_complete_workflow_task() .times(1) - .returning(|resp| { + .returning(|resp, _| { // Verify both the complete command and the query response are sent assert_eq!(resp.commands.len(), 1); assert_eq!(resp.query_responses.len(), 1); @@ -549,7 +549,7 @@ async fn query_replay_with_continue_as_new_doesnt_reply_empty_command() { mock_client .expect_complete_workflow_task() .times(1) - .returning(|resp| { + .returning(|resp, _| { // Verify both the complete command and the query response are sent assert_eq!(resp.commands.len(), 1); assert_eq!(resp.query_responses.len(), 1); @@ -754,7 +754,7 @@ async fn new_query_fail() { mock_client .expect_complete_workflow_task() .times(1) - .returning(|resp| { + .returning(|resp, _| { // Verify there is a failed query response along w/ start timer cmd assert_eq!(resp.commands.len(), 1); assert_matches!( @@ -1043,7 +1043,7 @@ async fn queries_arent_lost_in_buffer_void(#[values(false, true)] buffered_becau let mut mock = mock_worker_client(); mock.expect_complete_workflow_task() - .returning(|_| Ok(Default::default())); + .returning(|_, _| Ok(Default::default())); mock.expect_respond_legacy_query() .times(2) .returning(|_, _| Ok(Default::default())); diff --git a/crates/sdk-core/src/core_tests/updates.rs b/crates/sdk-core/src/core_tests/updates.rs index 2498d156d..2c254d619 100644 --- a/crates/sdk-core/src/core_tests/updates.rs +++ b/crates/sdk-core/src/core_tests/updates.rs @@ -110,7 +110,7 @@ async fn initial_request_sent_back(#[values(false, true)] reject: bool) { mock_client .expect_complete_workflow_task() .times(1) - .returning(move |mut resp| { + .returning(move |mut resp, _| { let msg = resp.messages.pop().unwrap(); let orig_req = if reject { let acceptance = msg.body.unwrap().to_msg::().unwrap(); @@ -322,3 +322,45 @@ async fn replay_with_signal_and_update_same_task() { .await .unwrap(); } + +#[tokio::test] +async fn update_activation_has_update_id() { + let wfid = "fakeid"; + let mut t = TestHistoryBuilder::default(); + t.add_by_type(EventType::WorkflowExecutionStarted); + t.add_workflow_task_scheduled_and_started(); + + let update_id = "upd-1"; + let mut poll_resp = hist_to_poll_resp(&t, wfid, ResponseType::AllHistory); + poll_resp.add_update_request(update_id, 1); + + let mut mock_client = mock_worker_client(); + mock_client + .expect_complete_workflow_task() + .times(1) + .returning(|_, _| Ok(RespondWorkflowTaskCompletedResponse::default())); + let mh = MockPollCfg::from_resp_batches(wfid, t, [poll_resp], mock_client); + let core = mock_worker(build_mock_pollers(mh)); + + let task = core.poll_workflow_activation().await.unwrap(); + let update = task + .jobs + .iter() + .find_map(|job| match job.variant.as_ref() { + Some(workflow_activation_job::Variant::DoUpdate(update)) => Some(update), + _ => None, + }) + .expect("activation should contain an update"); + assert_eq!(update.id, update_id); + + core.complete_workflow_activation(WorkflowActivationCompletion::from_cmd( + task.run_id, + UpdateResponse { + protocol_instance_id: update_id.to_string(), + response: Some(Response::Accepted(())), + } + .into(), + )) + .await + .unwrap(); +} diff --git a/crates/sdk-core/src/core_tests/workers.rs b/crates/sdk-core/src/core_tests/workers.rs index 01d858910..d5d324870 100644 --- a/crates/sdk-core/src/core_tests/workers.rs +++ b/crates/sdk-core/src/core_tests/workers.rs @@ -146,7 +146,7 @@ async fn worker_shutdown_during_poll_doesnt_deadlock() { let mut mock_client = mock_worker_client(); mock_client .expect_complete_workflow_task() - .returning(|_| Ok(RespondWorkflowTaskCompletedResponse::default())); + .returning(|_, _| Ok(RespondWorkflowTaskCompletedResponse::default())); let worker = mock_worker(MocksHolder::from_mock_worker(mock_client, mw)); let pollfut = worker.poll_workflow_activation(); let shutdownfut = async { @@ -206,7 +206,7 @@ async fn complete_with_task_not_found_during_shutdown() { let mut mock = mock_worker_client(); mock.expect_complete_workflow_task() .times(1) - .returning(|_| Err(tonic::Status::not_found("Workflow task not found."))); + .returning(|_, _| Err(tonic::Status::not_found("Workflow task not found."))); let mh = MockPollCfg::from_resp_batches("fakeid", t, [1], mock); let core = mock_worker(build_mock_pollers(mh)); @@ -277,7 +277,7 @@ async fn worker_does_not_panic_on_retry_exhaustion_of_nonfatal_net_err() { // Return a failure that counts as retryable, and hence we want to be swallowed mock.expect_complete_workflow_task() .times(1) - .returning(|_| Err(tonic::Status::internal("Some retryable error"))); + .returning(|_, _| Err(tonic::Status::internal("Some retryable error"))); let mut mh = MockPollCfg::from_resp_batches("fakeid", t, [1.into(), ResponseType::AllHistory], mock); mh.enforce_correct_number_of_polls = false; @@ -1291,7 +1291,7 @@ async fn graceful_shutdown_sends_shutdown_worker_rpc_during_initiate() { }); mock_client .expect_complete_workflow_task() - .returning(|_| Ok(RespondWorkflowTaskCompletedResponse::default())); + .returning(|_, _| Ok(RespondWorkflowTaskCompletedResponse::default())); // Polls block until shutdown_worker RPC releases them (simulating server holding polls // open until it receives the ShutdownWorker signal) diff --git a/crates/sdk-core/src/core_tests/workflow_cancels.rs b/crates/sdk-core/src/core_tests/workflow_cancels.rs index 197f70c37..a5390ce82 100644 --- a/crates/sdk-core/src/core_tests/workflow_cancels.rs +++ b/crates/sdk-core/src/core_tests/workflow_cancels.rs @@ -117,7 +117,7 @@ async fn immediate_cancel() { workflow_activation_job::Variant::InitializeWorkflow(_), workflow_activation_job::Variant::CancelWorkflow(_) ), - vec![CancelWorkflowExecution {}.into()], + vec![CancelWorkflowExecution::default().into()], )], ) .await; diff --git a/crates/sdk-core/src/core_tests/workflow_tasks.rs b/crates/sdk-core/src/core_tests/workflow_tasks.rs index dac74c801..9f4d95d91 100644 --- a/crates/sdk-core/src/core_tests/workflow_tasks.rs +++ b/crates/sdk-core/src/core_tests/workflow_tasks.rs @@ -4,7 +4,8 @@ use crate::{ job_assert, replay::{TestHistoryBuilder, canned_histories, default_act_sched, default_wes_attribs}, test_help::{ - FakeWfResponses, MockPollCfg, MocksHolder, ResponseType, WorkerExt, WorkerTestHelpers, + CounterRecordingMeter, FakeWfResponses, MockPollCfg, MocksHolder, ResponseType, WorkerExt, + WorkerTestHelpers, WorkflowCachingPolicy::{self, AfterEveryReply, NonSticky}, build_fake_worker, build_mock_pollers, build_multihist_mock_sg, fanout_tasks, gen_assert_and_fail, gen_assert_and_reply, hist_to_poll_resp, mock_worker, poll_and_reply, @@ -31,6 +32,7 @@ use std::{ }; use temporalio_client::MESSAGE_TOO_LARGE_KEY; use temporalio_common::{ + payload_limits::{LimitClass, LimitSeverity, PayloadLimitViolation}, protos::{ coresdk::{ ActivityTaskCompletion, @@ -317,7 +319,7 @@ async fn scheduled_activity_timeout(hist_batches: &'static [usize]) { seq, result: Some(ActivityResolution { status: Some(activity_resolution::Status::Failed(ar::Failure { - failure: Some(failure) + failure: Some(failure), .. })), }), .. } @@ -370,7 +372,7 @@ async fn started_activity_timeout(hist_batches: &'static [usize]) { seq, result: Some(ActivityResolution { status: Some(activity_resolution::Status::Failed(ar::Failure { - failure: Some(failure) + failure: Some(failure), .. })), }), .. } @@ -783,6 +785,38 @@ async fn simple_timer_fail_wf_execution(hist_batches: &'static [usize]) { .await; } +#[tokio::test] +async fn signal_activation_has_originating_event_id() { + let mut t = TestHistoryBuilder::default(); + t.add_by_type(EventType::WorkflowExecutionStarted); + t.add_full_wf_task(); + t.add_we_signaled("signal", vec![]); + let signal_event_id = t.current_event_id(); + t.add_full_wf_task(); + t.add_workflow_execution_completed(); + + let mock = MockPollCfg::from_resps(t, [ResponseType::AllHistory]); + let mut mock = build_mock_pollers(mock); + mock.worker_cfg(|wc| wc.max_cached_workflows = 1); + let core = mock_worker(mock); + + let task = core.poll_workflow_activation().await.unwrap(); + core.complete_workflow_activation(WorkflowActivationCompletion::empty(task.run_id)) + .await + .unwrap(); + + let task = core.poll_workflow_activation().await.unwrap(); + assert_matches!( + task.jobs.as_slice(), + [WorkflowActivationJob { + variant: Some(workflow_activation_job::Variant::SignalWorkflow(signal)), + }] => { + assert_eq!(signal.originating_event_id, signal_event_id); + } + ); + core.complete_execution(&task.run_id).await; +} + #[rstest(hist_batches, case::incremental(&[1, 2]), case::replay(&[2]))] #[tokio::test] async fn two_signals(hist_batches: &'static [usize]) { @@ -1105,9 +1139,9 @@ async fn sends_appropriate_sticky_task_queue_responses() { let t = canned_histories::single_timer("1"); let mut mock = mock_worker_client(); mock.expect_complete_workflow_task() - .withf(|comp| comp.sticky_attributes.is_some()) + .withf(|comp, _| comp.sticky_attributes.is_some()) .times(1) - .returning(|_| Ok(Default::default())); + .returning(|_, _| Ok(Default::default())); mock.expect_complete_workflow_task().times(0); let mut mock = single_hist_mock_sg(wfid, t, [1], mock, false); mock.worker_cfg(|wc| wc.max_cached_workflows = 10); @@ -1190,7 +1224,7 @@ async fn buffered_work_drained_on_shutdown() { ); let mut mock = mock_worker_client(); mock.expect_complete_workflow_task() - .returning(|_| Ok(RespondWorkflowTaskCompletedResponse::default())); + .returning(|_, _| Ok(RespondWorkflowTaskCompletedResponse::default())); let mut mock = MocksHolder::from_wft_stream(mock, stream::iter(tasks)); // Cache on to avoid being super repetitive mock.worker_cfg(|wc| wc.max_cached_workflows = 10); @@ -1364,10 +1398,10 @@ async fn lang_slower_than_wft_timeouts() { let mut mock = mock_worker_client(); mock.expect_complete_workflow_task() .times(1) - .returning(|_| Err(tonic::Status::not_found("Workflow task not found."))); + .returning(|_, _| Err(tonic::Status::not_found("Workflow task not found."))); mock.expect_complete_workflow_task() .times(1) - .returning(|_| Ok(Default::default())); + .returning(|_, _| Ok(Default::default())); let mut mock = single_hist_mock_sg(wfid, t, [1, 1], mock, true); let tasksmap = mock.outstanding_task_map.clone().unwrap(); mock.worker_cfg(|wc| { @@ -1783,10 +1817,10 @@ async fn tasks_from_completion_are_delivered() { }; mock.expect_complete_workflow_task() .times(1) - .returning(move |_| Ok(complete_resp.clone())); + .returning(move |_, _| Ok(complete_resp.clone())); mock.expect_complete_workflow_task() .times(1) - .returning(|_| Ok(Default::default())); + .returning(|_, _| Ok(Default::default())); let mut mock = single_hist_mock_sg(wfid, t, [1], mock, true); mock.worker_cfg(|wc| wc.max_cached_workflows = 2); let core = mock_worker(mock); @@ -1829,10 +1863,10 @@ async fn pagination_works_with_tasks_from_completion() { }; mock.expect_complete_workflow_task() .times(1) - .returning(move |_| Ok(complete_resp.clone())); + .returning(move |_, _| Ok(complete_resp.clone())); mock.expect_complete_workflow_task() .times(1) - .returning(|_| Ok(Default::default())); + .returning(|_, _| Ok(Default::default())); let get_exec_resp: GetWorkflowExecutionHistoryResponse = t.get_full_history_info().unwrap().into(); @@ -1878,7 +1912,7 @@ async fn poll_faster_than_complete_wont_overflow_cache() { mock_client .expect_complete_workflow_task() .times(3) - .returning(|_| Ok(Default::default())); + .returning(|_, _| Ok(Default::default())); let mut mock_cfg = MockPollCfg::new(tasks, true, 0); mock_cfg.mock_client = mock_client; let mut mock = build_mock_pollers(mock_cfg); @@ -2135,7 +2169,7 @@ async fn no_race_acquiring_permits() { .returning(move |_, _| async move { Ok(Default::default()) }.boxed()); mock_client .expect_complete_workflow_task() - .returning(|_| async move { Ok(Default::default()) }.boxed()); + .returning(|_, _| async move { Ok(Default::default()) }.boxed()); let worker = Worker::new_test( { @@ -2221,7 +2255,7 @@ async fn continue_as_new_preserves_some_values() { }; mock_client .expect_complete_workflow_task() - .returning(move |mut c| { + .returning(move |mut c, _| { let cmd = c.commands.pop().unwrap().attributes.unwrap(); if let Attributes::ContinueAsNewWorkflowExecutionCommandAttributes(a) = cmd { assert_eq!(a.workflow_type.unwrap().name, "meow"); @@ -2789,7 +2823,7 @@ async fn poller_wont_run_ahead_of_task_slots() { .returning(move |_, _| Ok(bunch_of_first_tasks.next().unwrap())); mock_client .expect_complete_workflow_task() - .returning(|_| Ok(Default::default())); + .returning(|_, _| Ok(Default::default())); let worker = Worker::new_test( { @@ -2899,7 +2933,7 @@ async fn use_compatible_version_flag( #[allow(deprecated)] mock_client .expect_complete_workflow_task() - .returning(move |mut c| { + .returning(move |mut c, _| { let can_cmd = c.commands.pop().unwrap().attributes.unwrap(); match can_cmd { Attributes::ContinueAsNewWorkflowExecutionCommandAttributes(a) => { @@ -2975,7 +3009,7 @@ async fn slot_provider_cant_hand_out_more_permits_than_cache_size() { .returning(move |_, _| Ok(bunch_of_first_tasks.next().unwrap())); mock_client .expect_complete_workflow_task() - .returning(|_| Ok(Default::default())); + .returning(|_, _| Ok(Default::default())); struct EndlessSupplier {} #[async_trait::async_trait] @@ -3137,8 +3171,8 @@ async fn both_normal_and_sticky_pollers_poll_concurrently() { let cc = Arc::clone(&counters); mock_client .expect_complete_workflow_task() - .returning(move |completion| { - if completion.task_token.0.ends_with(b"normal") { + .returning(move |completion, _| { + if completion.task_token.into_inner().ends_with(b"normal") { cc.normal_slots_active_count.fetch_sub(1, Ordering::Relaxed); } else { cc.sticky_slots_active_count.fetch_sub(1, Ordering::Relaxed); @@ -3239,6 +3273,8 @@ async fn grpc_message_too_large_doesnt_spam_task_fails() { let mut mock = build_mock_pollers(mh); mock.worker_cfg(|wc| wc.max_cached_workflows = 1); + let meter = Arc::new(CounterRecordingMeter::default()); + mock.set_temporal_meter(meter.clone().into_temporal_meter()); let core = mock_worker(mock); // Since the mock makes us fail 5 times, we should succeed on the sixth @@ -3253,4 +3289,216 @@ async fn grpc_message_too_large_doesnt_spam_task_fails() { core.complete_execution(&act.run_id).await; core.drain_pollers_and_shutdown().await; // Mock only expects 1 task failure, and would fail here if we spammed + // Every attempt counts as a failure though, even the unreported ones + assert_eq!( + meter.counter_total( + "workflow_task_execution_failed", + &[("failure_reason", "GrpcMessageTooLarge")] + ), + 5 + ); +} + +#[tokio::test] +async fn payloads_too_large_doesnt_spam_task_fails() { + let mut t = TestHistoryBuilder::default(); + t.add_by_type(EventType::WorkflowExecutionStarted); + t.add_workflow_task_scheduled_and_started(); + + let mut mh = MockPollCfg::from_resp_batches( + "fake_wf_id", + t, + [ + ResponseType::AllHistory, + ResponseType::AllHistory, + ResponseType::AllHistory, + ResponseType::AllHistory, + ResponseType::AllHistory, + ResponseType::AllHistory, + ], + mock_worker_client(), + ); + mh.num_expected_fails = 1; + let mut times = 1; + mh.completion_mock_fn = Some(Box::new(move |_| { + if times <= 5 { + let violation = PayloadLimitViolation { + path: "commands[0].input".to_string(), + class: LimitClass::Blob, + severity: LimitSeverity::Error, + size: 1024, + limit: 10, + }; + let mut err = tonic::Status::new(tonic::Code::InvalidArgument, violation.to_string()); + err.set_source(Arc::new(violation)); + times += 1; + Err(err) + } else { + Ok(Default::default()) + } + })); + let fails = Arc::new(AtomicUsize::new(0)); + let fails_clone = fails.clone(); + mh.expect_fail_wft_matcher = Box::new(move |_, cause, _| { + fails_clone.fetch_add(1, Ordering::Relaxed); + *cause == WorkflowTaskFailedCause::PayloadsTooLarge + }); + + let mut mock = build_mock_pollers(mh); + mock.worker_cfg(|wc| wc.max_cached_workflows = 1); + let meter = Arc::new(CounterRecordingMeter::default()); + mock.set_temporal_meter(meter.clone().into_temporal_meter()); + let core = mock_worker(mock); + + for _ in 1..=5 { + let act = core.poll_workflow_activation().await.unwrap(); + core.complete_workflow_activation(WorkflowActivationCompletion::empty(&act.run_id)) + .await + .unwrap(); + core.handle_eviction().await; + } + let act = core.poll_workflow_activation().await.unwrap(); + core.complete_execution(&act.run_id).await; + core.drain_pollers_and_shutdown().await; + assert_eq!(fails.load(Ordering::Relaxed), 1); + assert_eq!( + meter.counter_total( + "workflow_task_execution_failed", + &[("failure_reason", "PayloadsTooLarge")] + ), + 5 + ); +} + +/// A history fetch failure for a run that was never cached is reported through a path that has no +/// activation to complete. Later attempts of that same task must still not be re-reported. +#[tokio::test] +async fn unstored_wft_fetch_failure_doesnt_spam_task_fails() { + let mut t = TestHistoryBuilder::default(); + t.add_by_type(EventType::WorkflowExecutionStarted); + t.add_workflow_task_scheduled_and_started(); + t.add_workflow_task_completed(); + let mut need_fetch_resp = + hist_to_poll_resp(&t, "wfid".to_owned(), ResponseType::AllHistory).resp; + need_fetch_resp.next_page_token = vec![1]; + + let mut mock_client = mock_worker_client(); + mock_client + .expect_get_workflow_execution_history() + .returning(|_, _, _| Err(tonic::Status::not_found("Ahh broken"))) + .times(2); + // Identical responses are handed out with incrementing attempt numbers by the mock + let mut mh = MockPollCfg::from_resp_batches( + "wfid", + t, + [ + ResponseType::Raw(need_fetch_resp.clone()), + ResponseType::Raw(need_fetch_resp), + ], + mock_client, + ); + // Counted explicitly because a violated mock expectation inside the worker does not reliably + // fail the test. + let fails = Arc::new(AtomicUsize::new(0)); + let fails_clone = fails.clone(); + mh.num_expected_fails = 1; + mh.expect_fail_wft_matcher = Box::new(move |_, cause, _| { + fails_clone.fetch_add(1, Ordering::Relaxed); + *cause == WorkflowTaskFailedCause::WorkflowWorkerUnhandledFailure + }); + let mut mock = build_mock_pollers(mh); + let meter = Arc::new(CounterRecordingMeter::default()); + mock.set_temporal_meter(meter.clone().into_temporal_meter()); + let core = mock_worker(mock); + + // Both fetch failures are processed before the exhausted poller shuts the worker down + assert_matches!( + core.poll_workflow_activation().await.unwrap_err(), + PollError::ShutDown + ); + core.shutdown().await; + assert_eq!(fails.load(Ordering::Relaxed), 1); + assert_eq!( + meter.counter_total( + "workflow_task_execution_failed", + &[("failure_reason", "WorkflowError")] + ), + 2 + ); +} + +/// A history fetch failure for a cached run evicts it and reports the failure once the eviction +/// completes. Later attempts of that same task must still not be re-reported. +#[tokio::test] +async fn cached_run_fetch_failure_doesnt_spam_task_fails() { + let mut t = TestHistoryBuilder::default(); + t.add_by_type(EventType::WorkflowExecutionStarted); + t.add_workflow_task_scheduled_and_started(); + t.add_workflow_task_completed(); + let mut need_fetch_resp = + hist_to_poll_resp(&t, "wfid".to_owned(), ResponseType::AllHistory).resp; + need_fetch_resp.next_page_token = vec![1]; + + let mut mock_client = mock_worker_client(); + mock_client + .expect_get_workflow_execution_history() + .returning(|_, _, _| Err(tonic::Status::not_found("Ahh broken"))) + .times(2); + let mut mh = MockPollCfg::from_resp_batches( + "wfid", + t, + [ + ResponseType::ToTaskNum(1), + ResponseType::Raw(need_fetch_resp.clone()), + ResponseType::ToTaskNum(1), + ResponseType::Raw(need_fetch_resp), + ], + mock_client, + ); + // Counted explicitly because a violated mock expectation inside the worker does not reliably + // fail the test. + let fails = Arc::new(AtomicUsize::new(0)); + let fails_clone = fails.clone(); + mh.num_expected_fails = 1; + mh.expect_fail_wft_matcher = Box::new(move |_, cause, _| { + fails_clone.fetch_add(1, Ordering::Relaxed); + *cause == WorkflowTaskFailedCause::WorkflowWorkerUnhandledFailure + }); + let mut mock = build_mock_pollers(mh); + // Otherwise the poller runs dry and starts shutdown before the second eviction completes + mock.make_wft_stream_interminable(); + mock.worker_cfg(|wc| wc.max_cached_workflows = 10); + let meter = Arc::new(CounterRecordingMeter::default()); + mock.set_temporal_meter(meter.clone().into_temporal_meter()); + let core = mock_worker(mock); + + for _ in 0..2 { + let act = core.poll_workflow_activation().await.unwrap(); + assert_matches!( + act.jobs[0].variant, + Some(workflow_activation_job::Variant::InitializeWorkflow(_)) + ); + core.complete_workflow_activation(WorkflowActivationCompletion::empty(act.run_id)) + .await + .unwrap(); + let evict_act = core.poll_workflow_activation().await.unwrap(); + assert_matches!( + evict_act.jobs.as_slice(), + [WorkflowActivationJob { + variant: Some(workflow_activation_job::Variant::RemoveFromCache(r)), + }] => r.message.contains("Fetching history failed") + ); + core.complete_workflow_activation(WorkflowActivationCompletion::empty(evict_act.run_id)) + .await + .unwrap(); + } + core.shutdown().await; + assert_eq!(fails.load(Ordering::Relaxed), 1); + assert_eq!( + meter.counter_total( + "workflow_task_execution_failed", + &[("failure_reason", "WorkflowError")] + ), + 2 + ); } diff --git a/crates/sdk-core/src/ephemeral_server/mod.rs b/crates/sdk-core/src/ephemeral_server/mod.rs index 43e09bff0..608d8c7a1 100644 --- a/crates/sdk-core/src/ephemeral_server/mod.rs +++ b/crates/sdk-core/src/ephemeral_server/mod.rs @@ -1,16 +1,16 @@ //! This module implements support for downloading and running ephemeral test //! servers useful for testing. -use anyhow::anyhow; use flate2::read::GzDecoder; use futures_util::StreamExt; use serde::Deserialize; use std::{ + error::Error, fs::OpenOptions, io, path::{Path, PathBuf}, }; -use temporalio_client::{Connection, ConnectionOptions}; +use temporalio_client::{Connection, ConnectionOptions, errors::ClientConnectError}; use tokio::{ task::spawn_blocking, time::{Duration, sleep}, @@ -23,6 +23,98 @@ use zip::read::read_zipfile_from_stream; use std::os::unix::fs::OpenOptionsExt; use std::process::Stdio; +/// Errors encountered while downloading, starting, or stopping an ephemeral server. +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum EphemeralServerError { + /// A configured executable path does not exist. + #[error("ephemeral server executable does not exist: {}", path.display())] + ExecutableNotFound { + /// Missing executable path. + path: PathBuf, + }, + /// Downloading or caching an executable failed. + #[error("failed to download ephemeral server executable: {source}")] + Download { + /// Underlying download or cache failure. + #[source] + source: Box, + }, + /// Download metadata, platform support, or archive contents were invalid. + #[error("invalid ephemeral server download: {message}")] + InvalidDownload { + /// Description of the invalid content. + message: String, + /// Underlying validation or archive failure, when available. + #[source] + source: Option>, + }, + /// Starting the server process failed. + #[error("failed to start ephemeral server: {source}")] + ServerStart { + /// Underlying process or target failure. + #[source] + source: Box, + }, + /// The server process did not become available before its startup deadline. + #[error( + "ephemeral server at {target} did not start within {timeout:?}. Make sure another download isn't stuck and delete the temp file." + )] + StartupTimeout { + /// Server target that could not be reached. + target: String, + /// Amount of time spent waiting. + timeout: Duration, + /// Last connection error, when one was observed. + #[source] + last_error: Option, + }, + /// Stopping the server process failed. + #[error("failed to stop ephemeral server: {source}")] + ServerShutdown { + /// Underlying process failure. + #[source] + source: Box, + }, +} + +impl EphemeralServerError { + fn download_error(source: impl Error + Send + Sync + 'static) -> Self { + Self::Download { + source: Box::new(source), + } + } + + fn invalid_download(message: impl Into) -> Self { + Self::InvalidDownload { + message: message.into(), + source: None, + } + } + + fn invalid_download_with_source( + message: impl Into, + source: impl Error + Send + Sync + 'static, + ) -> Self { + Self::InvalidDownload { + message: message.into(), + source: Some(Box::new(source)), + } + } + + fn server_start_error(source: impl Error + Send + Sync + 'static) -> Self { + Self::ServerStart { + source: Box::new(source), + } + } + + fn server_shutdown_error(source: impl Error + Send + Sync + 'static) -> Self { + Self::ServerShutdown { + source: Box::new(source), + } + } +} + /// Configuration for Temporal CLI dev server. #[derive(Debug, Clone, bon::Builder)] #[builder(on(String, into))] @@ -44,7 +136,7 @@ pub struct TemporalDevServerConfig { /// Whether to enable the UI. If ui_port is set, assumes true. #[builder(default)] pub ui: bool, - /// Log format and level + /// Log format and level. #[builder(default = ("pretty".to_owned(), "warn".to_owned()))] pub log: (String, String), /// Additional arguments to Temporal dev server. @@ -54,7 +146,7 @@ pub struct TemporalDevServerConfig { impl TemporalDevServerConfig { /// Start a Temporal CLI dev server. - pub async fn start_server(&self) -> anyhow::Result { + pub async fn start_server(&self) -> Result { self.start_server_with_output(Stdio::inherit(), Stdio::inherit()) .await } @@ -64,7 +156,7 @@ impl TemporalDevServerConfig { &self, output: Stdio, err_output: Stdio, - ) -> anyhow::Result { + ) -> Result { // Get exe path let exe_path = self .exe @@ -74,7 +166,7 @@ impl TemporalDevServerConfig { // Get free port if not already given let port = match self.port { Some(p) => p, - None => get_free_port(&self.ip)?, + None => get_free_port(&self.ip).map_err(EphemeralServerError::server_start_error)?, }; // Build arg set @@ -142,7 +234,7 @@ pub struct TestServerConfig { impl TestServerConfig { /// Start a test server. - pub async fn start_server(&self) -> anyhow::Result { + pub async fn start_server(&self) -> Result { self.start_server_with_output(Stdio::inherit(), Stdio::inherit()) .await } @@ -152,7 +244,7 @@ impl TestServerConfig { &self, output: Stdio, err_output: Stdio, - ) -> anyhow::Result { + ) -> Result { // Get exe path let exe_path = self .exe @@ -162,7 +254,7 @@ impl TestServerConfig { // Get free port if not already given let port = match self.port { Some(p) => p, - None => get_free_port("0.0.0.0")?, + None => get_free_port("0.0.0.0").map_err(EphemeralServerError::server_start_error)?, }; // Build arg set @@ -202,7 +294,7 @@ pub struct EphemeralServer { } impl EphemeralServer { - async fn start(config: EphemeralServerConfig) -> anyhow::Result { + async fn start(config: EphemeralServerConfig) -> Result { // Start process. kill_on_drop ensures the process cannot outlive this // handle if start fails before an EphemeralServer (whose shutdown is // the normal kill path) is returned to the caller. @@ -212,15 +304,18 @@ impl EphemeralServer { .stdout(config.output) .stderr(config.err_output) .kill_on_drop(true) - .spawn()?; + .spawn() + .map_err(EphemeralServerError::server_start_error)?; let target = format!("127.0.0.1:{}", config.port); let target_url = format!("http://{target}"); - let connection_options = ConnectionOptions::new(Url::parse(&target_url)?) - .identity("online_checker".to_owned()) - .client_name("online-checker".to_owned()) - .client_version("0.1.0".to_owned()) - .build(); + let connection_options = ConnectionOptions::new( + Url::parse(&target_url).map_err(EphemeralServerError::server_start_error)?, + ) + .identity("online_checker".to_owned()) + .client_name("online-checker".to_owned()) + .client_version("0.1.0".to_owned()) + .build(); // Try to connect every 100ms for 5s // TODO(cretz): Some other way, e.g. via stdout, to know whether the @@ -243,22 +338,28 @@ impl EphemeralServer { // which does not wait for the kill to complete) so it cannot linger // holding inherited stdout/stderr pipes. let _ = child.kill().await; - Err(anyhow!( - "Failed connecting to test server after 5 seconds, last error: {last_error:?}" - )) + Err(EphemeralServerError::StartupTimeout { + target, + timeout: Duration::from_secs(5), + last_error, + }) } /// Shutdown the server (i.e. kill the child process). This does not attempt /// a kill if the child process appears completed, but such a check is not /// atomic so a kill could still fail as completed if completed just before /// kill. - pub async fn shutdown(&mut self) -> anyhow::Result<()> { + pub async fn shutdown(&mut self) -> Result<(), EphemeralServerError> { // Only kill if there is a PID if self.child.id().is_some() { - Ok(self.child.kill().await?) + self.child + .kill() + .await + .map_err(EphemeralServerError::server_shutdown_error)?; } else { - Ok(()) + return Ok(()); } + Ok(()) } /// Get the process ID of the child. This will be None if the process is @@ -324,12 +425,12 @@ impl EphemeralExe { artifact_name: &str, downloaded_name_prefix: &str, preferred_format: Option<&str>, - ) -> anyhow::Result { + ) -> Result { match self { EphemeralExe::ExistingPath(exe_path) => { let path = PathBuf::from(exe_path); if !path.exists() { - return Err(anyhow!("Exe path does not exist")); + return Err(EphemeralServerError::ExecutableNotFound { path }); } Ok(path) } @@ -370,7 +471,11 @@ impl EphemeralExe { let arch = match std::env::consts::ARCH { "x86_64" => "amd64", "arm" | "aarch64" => "arm64", - other => return Err(anyhow!("Unsupported arch: {other}")), + other => { + return Err(EphemeralServerError::invalid_download(format!( + "unsupported architecture: {other}" + ))); + } }; let mut get_info_params = vec![("arch", arch), ("platform", platform)]; if let Some(format) = preferred_format { @@ -396,9 +501,14 @@ impl EphemeralExe { )) .query(&get_info_params) .send() - .await? - .error_for_status()?; - let info: DownloadInfo = resp.json().await?; + .await + .map_err(EphemeralServerError::download_error)? + .error_for_status() + .map_err(EphemeralServerError::download_error)?; + let info: DownloadInfo = resp + .json() + .await + .map_err(EphemeralServerError::download_error)?; // Attempt download, looping because it could have waited for // concurrent one to finish @@ -478,7 +588,7 @@ async fn lazy_download_exe( file_to_extract: &Path, dest: &Path, already_tried_cleaning_old: bool, -) -> anyhow::Result { +) -> Result { // If it already exists, do not extract if dest.exists() { return Ok(true); @@ -488,7 +598,13 @@ async fn lazy_download_exe( // kind of global lock, we'll just create the file eagerly w/ a temp // filename and delete it on failure or move it on success. If the temp file // already exists, we'll wait a bit and re-run this. - let temp_dest_str = format!("{}{}", dest.to_str().unwrap(), ".downloading"); + let Some(dest_str) = dest.to_str() else { + return Err(EphemeralServerError::invalid_download(format!( + "download path is not UTF-8: {}", + dest.display() + ))); + }; + let temp_dest_str = format!("{dest_str}.downloading"); let temp_dest = Path::new(&temp_dest_str); // Try to open file, using a file mode on unix families #[cfg(target_family = "unix")] @@ -514,32 +630,36 @@ async fn lazy_download_exe( loop { let since_progress = match temp_dest.metadata() { Err(_) => return Ok(false), - Ok(meta) => meta.modified()?.elapsed()?.as_secs(), + Ok(meta) => meta + .modified() + .map_err(EphemeralServerError::download_error)? + .elapsed() + .map_err(EphemeralServerError::download_error)? + .as_secs(), }; if since_progress > DOWNLOAD_STALE_SECS { // No progress for a while; assume the downloader was // abandoned. Reclaim it once; if it goes stale again, fail // loudly rather than looping forever. if already_tried_cleaning_old { - return Err(anyhow!( - "Temp download file at {} made no progress for over {} \ - seconds. Make sure another download isn't stuck and \ - delete the temp file.", + return Err(EphemeralServerError::invalid_download(format!( + "temporary file at {} made no progress for over {} seconds", temp_dest.display(), DOWNLOAD_STALE_SECS, - )); + ))); } - std::fs::remove_file(temp_dest)?; + std::fs::remove_file(temp_dest) + .map_err(EphemeralServerError::download_error)?; return Box::pin(lazy_download_exe(client, uri, file_to_extract, dest, true)) .await; } sleep(Duration::from_secs(1)).await; } } - Err(err) => Err(err.into()), + Err(err) => Err(EphemeralServerError::download_error(err)), // If the dest was added since, just remove temp file Ok(_) if dest.exists() => { - std::fs::remove_file(temp_dest)?; + std::fs::remove_file(temp_dest).map_err(EphemeralServerError::download_error)?; return Ok(true); } // Download and extract the binary @@ -560,7 +680,7 @@ async fn lazy_download_exe( } }?; // Now that file should be dropped, we can rename - std::fs::rename(temp_dest, dest)?; + std::fs::rename(temp_dest, dest).map_err(EphemeralServerError::download_error)?; Ok(true) } @@ -569,10 +689,16 @@ async fn download_and_extract( uri: &str, file_to_extract: &Path, dest: &mut std::fs::File, -) -> anyhow::Result<()> { +) -> Result<(), EphemeralServerError> { // Start download. We are using streaming here to extract the file from the // tarball or zip instead of loading into memory for Cursor/Seek. - let resp = client.get(uri).send().await?.error_for_status()?; + let resp = client + .get(uri) + .send() + .await + .map_err(EphemeralServerError::download_error)? + .error_for_status() + .map_err(EphemeralServerError::download_error)?; // We have to map the error type to an io error let stream = resp .bytes_stream() @@ -586,43 +712,89 @@ async fn download_and_extract( } else if uri.ends_with(".zip") { false } else { - return Err(anyhow!("URI not .tar.gz or .zip")); + return Err(EphemeralServerError::invalid_download(format!( + "archive URL has unsupported format: {uri}" + ))); }; let file_to_extract = file_to_extract.to_path_buf(); - let mut dest = dest.try_clone()?; + let mut dest = dest + .try_clone() + .map_err(EphemeralServerError::download_error)?; - spawn_blocking(move || { + spawn_blocking(move || -> Result<(), EphemeralServerError> { if tarball { - for entry in tar::Archive::new(GzDecoder::new(reader)).entries()? { - let mut entry = entry?; - if entry.path()? == file_to_extract { - std::io::copy(&mut entry, &mut dest)?; + for entry in tar::Archive::new(GzDecoder::new(reader)) + .entries() + .map_err(|source| { + EphemeralServerError::invalid_download_with_source( + "could not read tar archive", + source, + ) + })? + { + let mut entry = entry.map_err(|source| { + EphemeralServerError::invalid_download_with_source( + "could not read tar archive entry", + source, + ) + })?; + if entry.path().map_err(|source| { + EphemeralServerError::invalid_download_with_source( + "tar archive entry path is invalid", + source, + ) + })? == file_to_extract + { + std::io::copy(&mut entry, &mut dest).map_err(|source| { + EphemeralServerError::invalid_download_with_source( + "could not extract tar archive entry", + source, + ) + })?; return Ok(()); } } - Err(anyhow!("Unable to find file in tarball")) + Err(EphemeralServerError::invalid_download( + "requested executable was not found in tar archive", + )) } else { loop { // This is the way to stream a zip file without creating an archive // that requires Seek. - if let Some(mut file) = read_zipfile_from_stream(&mut reader)? { + if let Some(mut file) = read_zipfile_from_stream(&mut reader).map_err(|source| { + EphemeralServerError::invalid_download_with_source( + "could not read ZIP archive", + source, + ) + })? { // If this is the file we're expecting, extract it if file.enclosed_name().as_ref() == Some(&file_to_extract) { - std::io::copy(&mut file, &mut dest)?; + std::io::copy(&mut file, &mut dest).map_err(|source| { + EphemeralServerError::invalid_download_with_source( + "could not extract ZIP archive entry", + source, + ) + })?; return Ok(()); } } else { - return Err(anyhow!("Unable to find file in zip")); + return Err(EphemeralServerError::invalid_download( + "requested executable was not found in ZIP archive", + )); } } } }) - .await? + .await + .map_err(EphemeralServerError::download_error)? } /// Remove the file if it's older than the TTL. Returns true if the current file can be re-used, /// returns false if it was removed or should otherwise be re-downloaded. -fn remove_file_past_ttl(ttl: &Option, dest: &PathBuf) -> Result { +fn remove_file_past_ttl( + ttl: &Option, + dest: &PathBuf, +) -> Result { match ttl { None => return Ok(true), Some(ttl) => { @@ -631,7 +803,7 @@ fn remove_file_past_ttl(ttl: &Option, dest: &PathBuf) -> Result Result<(), anyhow::Error> { .nth(1) .expect("must provide workflow id as only argument"); let run_id = std::env::args().nth(2).filter(|s| !s.is_empty()); - let handle = WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: wf_id.clone(), - run_id, - first_execution_run_id: None, - } - .bind_untyped(client); + let handle = WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(wf_id.clone()) + .maybe_run_id(run_id) + .build() + .bind_untyped(client); let events = handle .fetch_history(WorkflowFetchHistoryOptions::default()) - .await? - .into_events(); + .into_events() + .await?; let hist = History { events }; // Serialize history to file let byteified = hist.encode_to_vec(); diff --git a/crates/sdk-core/src/lib.rs b/crates/sdk-core/src/lib.rs index b6cd14e13..11ca4d498 100644 --- a/crates/sdk-core/src/lib.rs +++ b/crates/sdk-core/src/lib.rs @@ -195,14 +195,12 @@ pub struct RuntimeOptions { #[builder(default)] disable_environment_info: bool, /// Runtime information supplied by language SDK bridges. - #[doc(hidden)] #[builder(skip = vec![environment::native_runtime()])] runtimes: Vec, } impl RuntimeOptions { /// Supplies runtime information from a language SDK bridge. - #[doc(hidden)] pub fn with_runtimes(mut self, runtimes: Vec) -> Self { self.runtimes = runtimes; self diff --git a/crates/sdk-core/src/pollers/poll_buffer.rs b/crates/sdk-core/src/pollers/poll_buffer.rs index 45fe81d03..b35249bc1 100644 --- a/crates/sdk-core/src/pollers/poll_buffer.rs +++ b/crates/sdk-core/src/pollers/poll_buffer.rs @@ -59,6 +59,11 @@ const THROTTLE_POLL_BACKOFF: ExponentialBuilder = ExponentialBuilder::new() .with_factor(2.0) .with_max_delay(Duration::from_secs(10)) .without_max_times(); +const PERSISTENT_POLL_ERROR_WARN_BACKOFF: ExponentialBuilder = ExponentialBuilder::new() + .with_min_delay(Duration::from_secs(60)) + .with_factor(2.0) + .with_max_delay(Duration::from_secs(15 * 60)) + .without_max_times(); type PollReceiver = Mutex)>>>; @@ -132,7 +137,10 @@ impl LongPollBuffer { ); if let Some(wftps) = options.wft_poller_shared.as_ref() { if is_sticky { - wftps.set_sticky_active(poll_scaler.active_rx.clone()); + wftps.set_sticky_active( + poll_scaler.active_rx.clone(), + poll_scaler.report_handle.target.subscribe(), + ); } else { wftps.set_non_sticky_active(poll_scaler.active_rx.clone()); }; @@ -539,10 +547,11 @@ where initial, } => (minimum, maximum, initial), }; + let target = watch::Sender::new(target); let report_handle = Arc::new(PollScalerReportHandle { max, min, - target: AtomicUsize::new(target), + target, ever_saw_scaling_decision: AtomicBool::default(), capabilities, behavior, @@ -550,6 +559,7 @@ where ingested_last_period: Default::default(), scale_up_allowed: AtomicBool::new(true), last_successful_poll_time, + persistent_error_warning_state: Default::default(), exponential_backoff: parking_lot::Mutex::new(TASK_POLL_BACKOFF.build()), resource_exhausted_backoff: parking_lot::Mutex::new(THROTTLE_POLL_BACKOFF.build()), }); @@ -590,10 +600,7 @@ where async fn wait_until_allowed(&mut self) -> ActiveCounter> { self.active_rx - .wait_for(|v| { - *v < self.report_handle.max - && *v < self.report_handle.target.load(Ordering::Relaxed) - }) + .wait_for(|v| *v < self.report_handle.max && *v < *self.report_handle.target.borrow()) .await .expect("Poll allow does not panic"); ActiveCounter::new(self.active_tx.clone(), self.num_pollers_handler.clone()) @@ -607,7 +614,7 @@ where struct PollScalerReportHandle { max: usize, min: usize, - target: AtomicUsize, + target: watch::Sender, ever_saw_scaling_decision: AtomicBool, capabilities: Arc, behavior: PollerBehavior, @@ -616,6 +623,7 @@ struct PollScalerReportHandle { ingested_last_period: AtomicUsize, scale_up_allowed: AtomicBool, last_successful_poll_time: Arc>>, + persistent_error_warning_state: parking_lot::Mutex, // Exponential backoff for normal errors and resource exhausted errors exponential_backoff: parking_lot::Mutex, @@ -634,10 +642,10 @@ impl PollScalerReportHandle { Ok(res) => { self.last_successful_poll_time .store(Some(SystemTime::now())); - // Reset backoff on successful poll *self.exponential_backoff.lock() = TASK_POLL_BACKOFF.build(); *self.resource_exhausted_backoff.lock() = THROTTLE_POLL_BACKOFF.build(); + self.persistent_error_warning_state.lock().reset(); if let PollerBehavior::SimpleMaximum(_) = self.behavior { // We don't do auto-scaling with the simple max @@ -673,6 +681,18 @@ impl PollScalerReportHandle { } Err(e) => { if matches!(self.behavior, PollerBehavior::Autoscaling { .. }) { + if let Some(error_duration) = self + .persistent_error_warning_state + .lock() + .record_error(Instant::now()) + { + warn!( + error = ?e, + ?error_duration, + "Task polling has encountered errors continuously; the worker will continue retrying" + ); + } + // Follow the same backoff logic as the retry client let mut backoff_duration = self .exponential_backoff @@ -695,15 +715,10 @@ impl PollScalerReportHandle { .metadata() .contains_key(ERROR_RETURNED_DUE_TO_SHORT_CIRCUIT); - if self.can_scale_down() { + if self.can_scale_down() && e.code() == Code::ResourceExhausted { debug!("Got error from server while polling: {:?}", e); - if e.code() == Code::ResourceExhausted { - // Scale down significantly for resource exhaustion - self.change_target(usize::saturating_div, 2); - } else { - // Other codes that would normally have made us back off briefly can reclaim this poller - self.change_target(usize::saturating_sub, 1); - } + // Scale down significantly for resource exhaustion + self.change_target(usize::saturating_div, 2); } return (should_forward, backoff_duration); } @@ -714,11 +729,15 @@ impl PollScalerReportHandle { #[inline] fn change_target(&self, change: fn(usize, usize) -> usize, change_by: usize) { - self.target - .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |v| { - Some(change(v, change_by).clamp(self.min, self.max)) - }) - .expect("Cannot fail because always returns Some"); + self.target.send_if_modified(|target| { + let new_target = change(*target, change_by).clamp(self.min, self.max); + if *target == new_target { + return false; + } + + *target = new_target; + true + }); } /// We want to avoid scaling down on empty polls if the server has never made any scaling @@ -854,6 +873,44 @@ where } } +#[derive(Debug)] +struct PersistentPollErrorWarningState { + started_at: Option, + next_warning_at: Option, + warning_backoff: backon::ExponentialBackoff, +} + +impl Default for PersistentPollErrorWarningState { + fn default() -> Self { + Self { + started_at: None, + next_warning_at: None, + warning_backoff: PERSISTENT_POLL_ERROR_WARN_BACKOFF.build(), + } + } +} + +impl PersistentPollErrorWarningState { + fn record_error(&mut self, now: Instant) -> Option { + let Some(started_at) = self.started_at else { + self.started_at = Some(now); + self.next_warning_at = self.warning_backoff.next().map(|delay| now + delay); + return None; + }; + let next_warning_at = self.next_warning_at?; + if now < next_warning_at { + return None; + } + + self.next_warning_at = self.warning_backoff.next().map(|delay| now + delay); + Some(now.saturating_duration_since(started_at)) + } + + fn reset(&mut self) { + *self = Self::default(); + } +} + #[cfg(test)] mod tests { use super::*; @@ -863,7 +920,7 @@ mod tests { }; use futures_util::FutureExt; use rstest::rstest; - use std::time::Duration; + use std::{future::pending, time::Duration}; use temporalio_common::protos::temporal::api::namespace::v1::namespace_info::Capabilities; use tokio::{select, sync::Notify}; @@ -1078,11 +1135,14 @@ mod tests { #[rstest] #[case::resource_exhausted(Code::ResourceExhausted)] - #[case::internal(Code::Internal)] + #[case::cancelled(Code::Cancelled)] #[tokio::test] async fn autoscaler_applies_backoff_on_errors(#[case] error_code: Code) { use temporalio_common::protos::temporal::api::taskqueue::v1::PollerScalingDecision; + const INITIAL_POLLERS: usize = 10; + const MAX_POLLS_DURING_BACKOFF: usize = INITIAL_POLLERS + 1; + let call_count = Arc::new(AtomicUsize::new(0)); let call_count_clone = call_count.clone(); let first_poll_done = Arc::new(AtomicBool::new(false)); @@ -1124,13 +1184,13 @@ mod tests { PollerBehavior::Autoscaling { minimum: 5, maximum: 100, - initial: 10, + initial: INITIAL_POLLERS, }, - fixed_size_permit_dealer(10), + fixed_size_permit_dealer(INITIAL_POLLERS), CancellationToken::new(), None::, WorkflowTaskOptions { - wft_poller_shared: Some(Arc::new(WFTPollerShared::new(Some(10)))), + wft_poller_shared: Some(Arc::new(WFTPollerShared::new(Some(INITIAL_POLLERS)))), }, Arc::new(AtomicCell::new(None)), Arc::new(NamespaceCapabilities::default()), @@ -1149,11 +1209,10 @@ mod tests { tokio::time::sleep(Duration::from_millis(100)).await; let hot_loop_calls = call_count.load(Ordering::SeqCst); - // Without backoff, this was producing ~6300 polls in ~100ms on my machine. - // With exponential backoff, I'm getting exactly 10 (initial poller count). + // One replacement may follow the successful setup poll, but errors must not create a hot loop. assert!( - hot_loop_calls == 10, - "Expected proper backoff with == 10 polls in 100ms, but got {} polls.", + hot_loop_calls <= MAX_POLLS_DURING_BACKOFF, + "Expected at most {MAX_POLLS_DURING_BACKOFF} polls during backoff, got {}.", hot_loop_calls ); @@ -1163,6 +1222,85 @@ mod tests { .await; } + #[rstest] + #[case::cancelled(Code::Cancelled)] + #[case::deadline_exceeded(Code::DeadlineExceeded)] + #[tokio::test(start_paused = true)] + async fn transient_error_keeps_target(#[case] error_code: Code) { + const INITIAL_POLLERS: usize = 4; + const REPLACEMENT_CALLS: usize = INITIAL_POLLERS + 1; + const BACKOFF_SETTLE_TIME: Duration = Duration::from_millis(250); + + let call_count = Arc::new(AtomicUsize::new(0)); + let call_count_clone = call_count.clone(); + let fail_poll = Arc::new(Notify::new()); + let fail_poll_clone = fail_poll.clone(); + + let mut mock_client = mock_manual_worker_client(); + mock_client + .expect_poll_workflow_task() + .returning(move |_, _| { + let call_number = call_count_clone.fetch_add(1, Ordering::SeqCst) + 1; + let fail_poll = fail_poll_clone.clone(); + + async move { + if call_number == INITIAL_POLLERS { + fail_poll.notified().await; + + return Err(tonic::Status::new(error_code, "simulated poll error")); + } + + pending().await + } + .boxed() + }); + + let (active_tx, mut active_rx) = watch::channel(0); + let pb = LongPollBuffer::new_workflow_task( + Arc::new(mock_client), + "normal".to_string(), + Some("sticky".to_string()), + PollerBehavior::Autoscaling { + minimum: 1, + maximum: INITIAL_POLLERS, + initial: INITIAL_POLLERS, + }, + fixed_size_permit_dealer(INITIAL_POLLERS), + CancellationToken::new(), + Some(move |active| { + active_tx.send_replace(active); + }), + WorkflowTaskOptions { + wft_poller_shared: None, + }, + Arc::new(AtomicCell::new(None)), + Arc::new(NamespaceCapabilities::resolved(Capabilities { + poller_autoscaling: true, + ..Default::default() + })), + ); + + let _ = pb.starter.send(()); + active_rx + .wait_for(|active| *active == INITIAL_POLLERS) + .await + .unwrap(); + + fail_poll.notify_one(); + tokio::task::yield_now().await; + + assert_eq!(call_count.load(Ordering::SeqCst), INITIAL_POLLERS); + assert_eq!(*active_rx.borrow_and_update(), INITIAL_POLLERS); + + tokio::time::advance(BACKOFF_SETTLE_TIME).await; + tokio::task::yield_now().await; + + assert_eq!(call_count.load(Ordering::SeqCst), REPLACEMENT_CALLS); + assert_eq!(*active_rx.borrow_and_update(), INITIAL_POLLERS); + + pb.shutdown().await; + } + #[rstest] #[case::graceful(true)] #[case::legacy(false)] @@ -1267,7 +1405,7 @@ mod tests { let handle = Arc::new(PollScalerReportHandle { max: 10, min: minimum, - target: AtomicUsize::new(10), + target: watch::channel(10).0, ever_saw_scaling_decision: AtomicBool::new(false), capabilities: Arc::new(NamespaceCapabilities::resolved(Capabilities { poller_autoscaling: supports_autoscaling, @@ -1282,6 +1420,7 @@ mod tests { ingested_last_period: Default::default(), scale_up_allowed: AtomicBool::new(true), last_successful_poll_time: Arc::new(AtomicCell::new(None)), + persistent_error_warning_state: Default::default(), exponential_backoff: parking_lot::Mutex::new(TASK_POLL_BACKOFF.build()), resource_exhausted_backoff: parking_lot::Mutex::new(THROTTLE_POLL_BACKOFF.build()), }); @@ -1292,7 +1431,27 @@ mod tests { handle.poll_result(&empty_resp); } - assert_eq!(handle.target.load(Ordering::Relaxed), expected_target); + assert_eq!(*handle.target.borrow(), expected_target); assert!(!handle.ever_saw_scaling_decision.load(Ordering::Relaxed)); } + + #[test] + fn persistent_poll_error_warning_uses_exponential_backoff_and_resets() { + let mut state = PersistentPollErrorWarningState::default(); + let started_at = Instant::now(); + let minute = Duration::from_secs(60); + + assert_eq!(state.record_error(started_at), None); + assert_eq!( + state.record_error(started_at + minute - Duration::from_secs(1)), + None + ); + for warning_minute in [1, 3, 7, 15, 30, 45] { + let elapsed = minute * warning_minute; + assert_eq!(state.record_error(started_at + elapsed), Some(elapsed)); + } + + state.reset(); + assert_eq!(state.record_error(started_at + minute * 45), None); + } } diff --git a/crates/sdk-core/src/protosext/mod.rs b/crates/sdk-core/src/protosext/mod.rs index b2adcde96..506b953c6 100644 --- a/crates/sdk-core/src/protosext/mod.rs +++ b/crates/sdk-core/src/protosext/mod.rs @@ -38,7 +38,7 @@ use temporalio_common::protos::{ failure::v1::Failure, history::v1::{History, HistoryEvent, MarkerRecordedEventAttributes, history_event}, query::v1::WorkflowQuery, - sdk::v1::UserMetadata, + sdk::v1::{EventGroupMarker, UserMetadata}, workflowservice::v1::PollWorkflowTaskQueueResponse, }, utilities::TryIntoOrNone, @@ -121,7 +121,7 @@ impl TryFrom for ValidPollWFTQResponse { let messages = messages.into_iter().map(TryInto::try_into).try_collect()?; Ok(Self { - task_token: TaskToken(task_token), + task_token: task_token.into(), task_queue: tq.name, workflow_execution, workflow_type: workflow_type.name, @@ -322,7 +322,9 @@ pub(crate) struct ValidScheduleLA { pub(crate) retry_policy: ValidatedRetryPolicy, pub(crate) local_retry_threshold: Duration, pub(crate) cancellation_type: ActivityCancellationType, + pub(crate) include_arguments_in_marker: bool, pub(crate) user_metadata: Option, + pub(crate) event_group_markers: Vec, } #[derive(Debug, Clone, Copy)] @@ -355,6 +357,7 @@ impl ValidScheduleLA { pub(crate) fn from_schedule_la( v: ScheduleLocalActivity, user_metadata: Option, + event_group_markers: Vec, ) -> Result { let original_schedule_time = v .original_schedule_time @@ -430,7 +433,9 @@ impl ValidScheduleLA { retry_policy, local_retry_threshold, cancellation_type, + include_arguments_in_marker: v.include_arguments_in_marker, user_metadata, + event_group_markers, }) } } diff --git a/crates/sdk-core/src/replay/mod.rs b/crates/sdk-core/src/replay/mod.rs index ea38736b5..6e4cadc50 100644 --- a/crates/sdk-core/src/replay/mod.rs +++ b/crates/sdk-core/src/replay/mod.rs @@ -129,9 +129,11 @@ where .boxed() }); - client.expect_complete_workflow_task().returning(move |_a| { - async move { Ok(RespondWorkflowTaskCompletedResponse::default()) }.boxed() - }); + client + .expect_complete_workflow_task() + .returning(move |_a, _b| { + async move { Ok(RespondWorkflowTaskCompletedResponse::default()) }.boxed() + }); client .expect_fail_workflow_task() .returning(move |_, _, _| { diff --git a/crates/sdk-core/src/telemetry/metrics.rs b/crates/sdk-core/src/telemetry/metrics.rs index cc569670d..41d5864e1 100644 --- a/crates/sdk-core/src/telemetry/metrics.rs +++ b/crates/sdk-core/src/telemetry/metrics.rs @@ -10,7 +10,10 @@ use std::{ time::Duration, }; use temporalio_common::{ - protos::temporal::api::{enums::v1::WorkflowTaskFailedCause, failure::v1::Failure}, + protos::{ + coresdk::activity_result::ActivityTaskFailedCause, + temporal::api::{enums::v1::WorkflowTaskFailedCause, failure::v1::Failure}, + }, telemetry::metrics::{core::*, *}, }; @@ -755,22 +758,28 @@ pub(crate) fn eager(is_eager: bool) -> MetricKeyValue { pub(crate) enum FailureReason { Nondeterminism, Workflow, + Activity, Timeout, NexusOperation(String), NexusHandlerError(String), GrpcMessageTooLarge, PayloadsTooLarge, + ExternalStorageError, + RequestTooLarge, } impl Display for FailureReason { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { let str = match self { FailureReason::Nondeterminism => "NonDeterminismError".to_owned(), FailureReason::Workflow => "WorkflowError".to_owned(), + FailureReason::Activity => "ActivityError".to_owned(), FailureReason::Timeout => "timeout".to_owned(), FailureReason::NexusOperation(op) => format!("operation_{op}"), FailureReason::NexusHandlerError(op) => format!("handler_error_{op}"), FailureReason::GrpcMessageTooLarge => "GrpcMessageTooLarge".to_owned(), FailureReason::PayloadsTooLarge => "PayloadsTooLarge".to_owned(), + FailureReason::ExternalStorageError => "ExternalStorageError".to_owned(), + FailureReason::RequestTooLarge => "RequestTooLarge".to_owned(), }; write!(f, "{str}") } @@ -779,10 +788,23 @@ impl From for FailureReason { fn from(v: WorkflowTaskFailedCause) -> Self { match v { WorkflowTaskFailedCause::NonDeterministicError => FailureReason::Nondeterminism, + WorkflowTaskFailedCause::GrpcMessageTooLarge => FailureReason::GrpcMessageTooLarge, + WorkflowTaskFailedCause::PayloadsTooLarge => FailureReason::PayloadsTooLarge, + WorkflowTaskFailedCause::RequestTooLarge => FailureReason::RequestTooLarge, _ => FailureReason::Workflow, } } } +impl From for FailureReason { + fn from(v: ActivityTaskFailedCause) -> Self { + match v { + ActivityTaskFailedCause::PayloadsTooLarge => FailureReason::PayloadsTooLarge, + ActivityTaskFailedCause::ExternalStorageFailure => FailureReason::ExternalStorageError, + ActivityTaskFailedCause::Unspecified + | ActivityTaskFailedCause::ActivityWorkerUnhandledFailure => FailureReason::Activity, + } + } +} pub(crate) fn failure_reason(reason: FailureReason) -> MetricKeyValue { MetricKeyValue::new(KEY_TASK_FAILURE_TYPE, reason.to_string()) } diff --git a/crates/sdk-core/src/test_help/integ_helpers.rs b/crates/sdk-core/src/test_help/integ_helpers.rs index 5847f83df..0e446be95 100644 --- a/crates/sdk-core/src/test_help/integ_helpers.rs +++ b/crates/sdk-core/src/test_help/integ_helpers.rs @@ -752,7 +752,7 @@ pub fn build_mock_pollers(mut cfg: MockPollCfg) -> MocksHolder { let rid = t.workflow_execution.as_ref().unwrap().run_id.clone(); if !outstanding.has_run(&rid) { let t = tasks.pop_front().unwrap(); - outstanding.put_token(rid, TaskToken(t.task_token.clone())); + outstanding.put_token(rid, t.task_token.clone().into()); resp = Some(t); break; } @@ -808,7 +808,7 @@ pub fn build_mock_pollers(mut cfg: MockPollCfg) -> MocksHolder { } else if cfg.completion_mock_fn.is_some() { expect_completes.times(1..); } - expect_completes.returning(move |comp| { + expect_completes.returning(move |comp, _| { let r = if let Some(ass) = cfg.completion_mock_fn.as_mut() { // tee hee ass(&comp) diff --git a/crates/sdk-core/src/test_help/unit_helpers.rs b/crates/sdk-core/src/test_help/unit_helpers.rs index 4d3220c13..e782c0485 100644 --- a/crates/sdk-core/src/test_help/unit_helpers.rs +++ b/crates/sdk-core/src/test_help/unit_helpers.rs @@ -1,11 +1,132 @@ //! Unit test helpers - only available in unit tests (cfg(test)) use futures_util::{StreamExt, stream::FuturesUnordered}; -use std::{collections::HashSet, future::Future}; -use temporalio_common::protos::coresdk::{ - workflow_activation::workflow_activation_job, - workflow_completion::{WorkflowActivationCompletion, workflow_activation_completion}, +use std::{ + collections::{HashMap, HashSet}, + future::Future, + sync::{Arc, Mutex}, }; +use temporalio_common::{ + protos::coresdk::{ + workflow_activation::workflow_activation_job, + workflow_completion::{WorkflowActivationCompletion, workflow_activation_completion}, + }, + telemetry::{ + TaskQueueLabelStrategy, + metrics::{ + CoreMeter, Counter, CounterBase, Gauge, GaugeF64, Histogram, HistogramDuration, + HistogramF64, MetricAttributable, MetricAttributes, MetricParameters, NewAttributes, + NoOpCoreMeter, TemporalMeter, UpDownCounter, + }, + }, +}; + +/// A meter that records counter increments in memory so unit tests can assert on them. All other +/// instrument kinds are discarded. +#[derive(Debug, Default)] +pub struct CounterRecordingMeter { + adds: Arc>>, +} +#[derive(Debug)] +struct CounterAdd { + name: String, + labels: HashMap, + value: u64, +} +impl CounterRecordingMeter { + pub fn into_temporal_meter(self: Arc) -> TemporalMeter { + TemporalMeter::new( + self, + NewAttributes::default(), + TaskQueueLabelStrategy::UseNormal, + ) + } + + /// Sum of all increments to the named counter whose labels include every provided pair + pub fn counter_total(&self, name: &str, labels: &[(&str, &str)]) -> u64 { + self.adds + .lock() + .unwrap() + .iter() + .filter(|a| { + a.name == name + && labels + .iter() + .all(|(k, v)| a.labels.get(*k).map(String::as_str) == Some(*v)) + }) + .map(|a| a.value) + .sum() + } +} +struct RecordingCounter { + name: String, + adds: Arc>>, +} +impl MetricAttributable> for RecordingCounter { + fn with_attributes( + &self, + attributes: &MetricAttributes, + ) -> Result, Box> { + let MetricAttributes::NoOp(labels) = attributes else { + panic!("CounterRecordingMeter only works with NoOp attributes"); + }; + Ok(Box::new(BoundRecordingCounter { + name: self.name.clone(), + labels: (**labels).clone(), + adds: self.adds.clone(), + })) + } +} +struct BoundRecordingCounter { + name: String, + labels: HashMap, + adds: Arc>>, +} +impl CounterBase for BoundRecordingCounter { + fn adds(&self, value: u64) { + self.adds.lock().unwrap().push(CounterAdd { + name: self.name.clone(), + labels: self.labels.clone(), + value, + }); + } +} +impl CoreMeter for CounterRecordingMeter { + fn new_attributes(&self, attribs: NewAttributes) -> MetricAttributes { + NoOpCoreMeter.new_attributes(attribs) + } + fn extend_attributes( + &self, + existing: MetricAttributes, + attribs: NewAttributes, + ) -> MetricAttributes { + NoOpCoreMeter.extend_attributes(existing, attribs) + } + fn counter(&self, params: MetricParameters) -> Counter { + Counter::new(Arc::new(RecordingCounter { + name: params.name.to_string(), + adds: self.adds.clone(), + })) + } + fn histogram(&self, params: MetricParameters) -> Histogram { + NoOpCoreMeter.histogram(params) + } + fn histogram_f64(&self, params: MetricParameters) -> HistogramF64 { + NoOpCoreMeter.histogram_f64(params) + } + fn histogram_duration(&self, params: MetricParameters) -> HistogramDuration { + NoOpCoreMeter.histogram_duration(params) + } + fn gauge(&self, params: MetricParameters) -> Gauge { + NoOpCoreMeter.gauge(params) + } + fn gauge_f64(&self, params: MetricParameters) -> GaugeF64 { + NoOpCoreMeter.gauge_f64(params) + } + fn up_down_counter(&self, params: MetricParameters) -> UpDownCounter { + NoOpCoreMeter.up_down_counter(params) + } +} /// Given a desired number of concurrent executions and a provided function that produces a future, /// run that many instances of the future concurrently. diff --git a/crates/sdk-core/src/worker/activities.rs b/crates/sdk-core/src/worker/activities.rs index 9d881fdfc..c41991103 100644 --- a/crates/sdk-core/src/worker/activities.rs +++ b/crates/sdk-core/src/worker/activities.rs @@ -9,16 +9,18 @@ pub(crate) use local_activities::{ use crate::{ TaskToken, abstractions::{ - ClosableMeteredPermitDealer, MeteredPermitDealer, TrackedOwnedMeteredSemPermit, - UsedMeteredSemPermit, + ActiveCounter, ClosableMeteredPermitDealer, MeteredPermitDealer, + TrackedOwnedMeteredSemPermit, UsedMeteredSemPermit, }, pollers::{BoxedActPoller, PermittedTqResp, TrackedPermittedTqResp, new_activity_task_poller}, telemetry::metrics::{ - MetricsContext, activity_type, eager, should_record_failure_metric, workflow_type, + FailureReason, MetricsContext, activity_type, eager, failure_reason, + should_record_failure_metric, workflow_type, }, worker::{ ActivitySlotKind, PollError, - activities::activity_heartbeat_manager::ActivityHeartbeatError, client::WorkerClient, + activities::activity_heartbeat_manager::ActivityHeartbeatError, + client::{WorkerClient, payload_limit_violation_from}, }, }; use activity_heartbeat_manager::ActivityHeartbeatManager; @@ -36,13 +38,15 @@ use std::{ }, time::{Duration, Instant, SystemTime}, }; -use temporalio_client::{payload_limit_violation_from, worker::CancelActivityCallback}; +use temporalio_client::{PayloadErrorLimits, worker::CancelActivityCallback}; use temporalio_common::{ - payload_limits::PayloadLimitViolation, + payload_limits::{PayloadLimitViolation, PayloadLimits, validate_known_payload_limits}, protos::{ coresdk::{ ActivityHeartbeat, ActivitySlotInfo, - activity_result::{self as ar, activity_execution_result as aer}, + activity_result::{ + self as ar, ActivityTaskFailedCause, activity_execution_result as aer, + }, activity_task::{ActivityCancelReason, ActivityCancellationDetails, ActivityTask}, }, temporal::api::{ @@ -50,7 +54,9 @@ use temporalio_common::{ failure::v1::{ ApplicationFailureInfo, CanceledFailureInfo, Failure, failure::FailureInfo, }, - workflowservice::v1::PollActivityTaskQueueResponse, + workflowservice::v1::{ + PollActivityTaskQueueResponse, RecordActivityTaskHeartbeatRequest, + }, }, }, }; @@ -59,6 +65,7 @@ use tokio::{ sync::{ Mutex, Notify, mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel}, + watch, }, task::JoinHandle, }; @@ -176,8 +183,13 @@ pub(crate) struct WorkerActivityTasks { max_heartbeat_throttle_interval: Duration, default_heartbeat_throttle_interval: Duration, - /// Wakes every time an activity is removed from the outstanding map - complete_notify: Arc, + /// Counts completions which have already taken their task out of + /// `outstanding_activity_tasks` but are still flushing the result to server. Such a + /// completion still owns the activity's slot permit, so shutdown must not treat the empty + /// map as "all activities finished" while this is nonzero — otherwise the heartbeat manager + /// can be torn down out from under an in-flight eviction (stranding it forever) and worker + /// shutdown can proceed while the result was never reported. + completions_in_flight: watch::Sender, /// Token to notify when poll returned a shutdown error poll_returned_shutdown_token: CancellationToken, /// Used to inject external cancellations (e.g. from nexus worker commands) @@ -219,7 +231,7 @@ impl WorkerActivityTasks { let (cancels_tx, cancels_rx) = unbounded_channel(); let external_cancels_tx = cancels_tx.clone(); let heartbeat_manager = ActivityHeartbeatManager::new(client, cancels_tx.clone()); - let complete_notify = Arc::new(Notify::new()); + let (completions_in_flight, completions_in_flight_rx) = watch::channel(0); let source_stream = stream::select_with_strategy( UnboundedReceiverStream::new(cancels_rx).map(ActivityTaskSource::from), starts_stream.map(|a| ActivityTaskSource::from(Box::new(a))), @@ -230,7 +242,7 @@ impl WorkerActivityTasks { source_stream, outstanding_tasks: outstanding_activity_tasks.clone(), start_tasks_stream_complete, - complete_notify: complete_notify.clone(), + completions_in_flight: completions_in_flight_rx, grace_period: graceful_shutdown, cancels_tx, local_timeout_buffer, @@ -245,7 +257,7 @@ impl WorkerActivityTasks { heartbeat_manager, activity_task_stream: Mutex::new(activity_task_stream.boxed()), eager_activities_semaphore, - complete_notify, + completions_in_flight, metrics, max_heartbeat_throttle_interval, default_heartbeat_throttle_interval, @@ -334,6 +346,12 @@ impl WorkerActivityTasks { status: aer::Status, client: &dyn WorkerClient, ) { + // Counted before taking the task out of the outstanding map so shutdown can never + // observe the map empty without also seeing this completion in flight. Declared first so + // it drops after `act_info` — and thus after the slot permit — even if this future is + // cancelled or panics mid-completion. + let _completion_guard = + ActiveCounter::::new(self.completions_in_flight.clone(), None); let act_info = { let mut outstanding_activity_tasks = self.outstanding_activity_tasks.lock(); outstanding_activity_tasks.remove(&task_token) @@ -363,13 +381,24 @@ impl WorkerActivityTasks { .evict(task_token.clone(), should_flush) .await; - let last_heartbeat_details = act_info - .last_heartbeat_details - .map(|payloads| Payloads { payloads }); - // No need to report activities which we already know the server doesn't care about if !known_not_found { let _flushing_guard = self.completers_lock.read().await; + + let mut last_heartbeat_details = act_info + .last_heartbeat_details + .map(|payloads| Payloads { payloads }); + // Only failure reports carry these to the server, and oversized details would make + // the server reject such a request outright, so drop them and report the violation + // as the failure instead. Checked here rather than in the client so that the + // reported failure, the cause, and the metric can't disagree about what happened. + let heartbeat_details_violation = last_heartbeat_details.as_ref().and_then(|d| { + heartbeat_details_limit_violation(d, client.payload_error_limits()) + }); + if heartbeat_details_violation.is_some() { + last_heartbeat_details = None; + } + let maybe_net_err = match status { aer::Status::WillCompleteAsync(_) => None, aer::Status::Completed(ar::Success { result }) => { @@ -390,10 +419,15 @@ impl WorkerActivityTasks { } Err(e) => { if let Some(violation) = payload_limit_violation_from(&e) { - act_metrics.act_execution_failed(); + act_metrics + .with_new_attrs([failure_reason( + FailureReason::PayloadsTooLarge, + )]) + .act_execution_failed(); client .fail_activity_task( task_token.clone(), + ActivityTaskFailedCause::PayloadsTooLarge, Some(make_payloads_too_large_failure(violation)), last_heartbeat_details.clone(), ) @@ -405,12 +439,42 @@ impl WorkerActivityTasks { } } } - aer::Status::Failed(ar::Failure { failure }) => { - if should_record_failure_metric(&failure) { - act_metrics.act_execution_failed(); - } + aer::Status::Failed(fail) => { + let (cause, failure) = if let Some(violation) = + heartbeat_details_violation.as_ref() + { + // What reaches the server is no longer whatever lang reported, so + // the metric must be recorded even for an otherwise benign failure. + act_metrics + .with_new_attrs([failure_reason(FailureReason::PayloadsTooLarge)]) + .act_execution_failed(); + ( + ActivityTaskFailedCause::PayloadsTooLarge, + Some(make_payloads_too_large_failure(violation)), + ) + } else { + // An SDK reporting no cause recognized nothing more specific, which is + // normalized to unhandled failure so all SDKs do not have to specify it. + let cause = match fail.cause() { + ActivityTaskFailedCause::Unspecified => { + ActivityTaskFailedCause::ActivityWorkerUnhandledFailure + } + c => c, + }; + if should_record_failure_metric(&fail.failure) { + act_metrics + .with_new_attrs([failure_reason(cause.into())]) + .act_execution_failed(); + } + (cause, fail.failure) + }; client - .fail_activity_task(task_token.clone(), failure, last_heartbeat_details) + .fail_activity_task( + task_token.clone(), + cause, + failure, + last_heartbeat_details, + ) .await .err() } @@ -422,10 +486,21 @@ impl WorkerActivityTasks { // We report cancels for graceful shutdown as failures, so we // don't wait for the whole timeout to elapse, which is what would // happen anyway. + let (cause, failure) = match heartbeat_details_violation.as_ref() { + Some(violation) => ( + ActivityTaskFailedCause::PayloadsTooLarge, + make_payloads_too_large_failure(violation), + ), + None => ( + ActivityTaskFailedCause::ActivityWorkerUnhandledFailure, + worker_shutdown_failure(), + ), + }; client .fail_activity_task( task_token.clone(), - Some(worker_shutdown_failure()), + cause, + Some(failure), last_heartbeat_details, ) .await @@ -453,10 +528,15 @@ impl WorkerActivityTasks { Ok(_) => None, Err(e) => { if let Some(violation) = payload_limit_violation_from(&e) { - act_metrics.act_execution_failed(); + act_metrics + .with_new_attrs([failure_reason( + FailureReason::PayloadsTooLarge, + )]) + .act_execution_failed(); client .fail_activity_task( task_token.clone(), + ActivityTaskFailedCause::PayloadsTooLarge, Some(make_payloads_too_large_failure(violation)), last_heartbeat_details, ) @@ -487,8 +567,6 @@ impl WorkerActivityTasks { &task_token ); } - - self.complete_notify.notify_waiters(); } /// Attempt to record an activity heartbeat @@ -499,8 +577,9 @@ impl WorkerActivityTasks { // TODO: Propagate these back as cancels. Silent fails is too nonobvious let (heartbeat_timeout, timeout_resetter) = { let mut outstanding_activity_tasks = self.outstanding_activity_tasks.lock(); + let task_token: TaskToken = details.task_token.clone().into(); let at_info = outstanding_activity_tasks - .get_mut(&TaskToken(details.task_token.clone())) + .get_mut(&task_token) .ok_or(ActivityHeartbeatError::UnknownActivity)?; at_info.last_heartbeat_details = Some(details.details.clone()); (at_info.heartbeat_timeout, at_info.timeout_resetter.clone()) @@ -566,7 +645,7 @@ struct ActivityTaskStream { source_stream: SrcStrm, outstanding_tasks: OutstandingActMap, start_tasks_stream_complete: CancellationToken, - complete_notify: Arc, + completions_in_flight: watch::Receiver, grace_period: Option, cancels_tx: UnboundedSender, /// The extra time we'll wait for local timeouts before firing them, to avoid racing with server @@ -611,7 +690,7 @@ where details.known_not_found = true; } Some(Ok(ActivityTask::cancel_from_ids( - next_pc.task_token.0, + next_pc.task_token.into_inner(), next_pc.reason, next_pc.details, ))) @@ -745,11 +824,23 @@ where join!( async { self.start_tasks_stream_complete.cancelled().await; - while { - let outstanding_tasks = outstanding_tasks_clone.lock(); - !outstanding_tasks.is_empty() - } { - self.complete_notify.notified().await + let mut completions_in_flight = self.completions_in_flight; + loop { + let no_outstanding = outstanding_tasks_clone.lock().is_empty(); + // An empty map alone isn't "all activities finished": completions + // flushing to server have already left the map but still hold their + // slot permit, and still need the heartbeat manager alive. Tasks only + // ever leave the map inside a counted completion, so every relevant + // transition ends in a counter change and waiting on the counter + // alone can't miss one. + if no_outstanding && *completions_in_flight.borrow_and_update() == 0 { + break; + } + if completions_in_flight.changed().await.is_err() { + // Sender closed: the manager (and any completion guards, which + // hold sender clones) are gone, so nothing further can flush. + break; + } } // If we were waiting for the grace period but everything already finished, // we don't need to keep waiting. @@ -807,6 +898,26 @@ fn worker_shutdown_failure() -> Failure { } } +/// Validates final heartbeat details against the worker's payload error limits, since attaching +/// oversized details to a failure request would make the server reject the request as a whole. +fn heartbeat_details_limit_violation( + details: &Payloads, + limits: Option, +) -> Option { + let limits = limits?; + validate_known_payload_limits( + &RecordActivityTaskHeartbeatRequest { + details: Some(details.clone()), + ..Default::default() + }, + &PayloadLimits { + blob_error: limits.blob, + memo_error: limits.memo, + ..Default::default() + }, + ) +} + /// The failure is deliberately retryable: catching the violation client-side exists precisely to /// turn what the server would hard-fail into a recoverable activity task failure, so fixing and /// redeploying the activity lets the next attempt succeed. @@ -1057,13 +1168,13 @@ mod tests { shutdown_token.cancel(); // Need to complete the tasks so shutdown will resolve atm.complete( - TaskToken(t1.task_token), + t1.task_token.into(), ActivityExecutionResult::ok(vec![1].into()).status.unwrap(), mock_client.as_ref(), ) .await; atm.complete( - TaskToken(t2.task_token), + t2.task_token.into(), ActivityExecutionResult::ok(vec![1].into()).status.unwrap(), mock_client.as_ref(), ) @@ -1126,7 +1237,7 @@ mod tests { // Make sure it didn't take wayyy too long. Our long timeouts specified above are huge assert!(start.elapsed() < Duration::from_secs(5)); atm.complete( - TaskToken(t.task_token), + t.task_token.into(), ActivityExecutionResult::fail("unimportant".into()) .status .unwrap(), @@ -1191,7 +1302,7 @@ mod tests { join!(heartbeater, poller); atm.complete( - TaskToken(t.task_token), + t.task_token.into(), ActivityExecutionResult::fail("unimportant".into()) .status .unwrap(), @@ -1259,7 +1370,7 @@ mod tests { assert!(activity_task.is_timeout()); atm.complete( - TaskToken(t.task_token), + t.task_token.into(), ActivityExecutionResult::fail("unimportant".into()) .status .unwrap(), diff --git a/crates/sdk-core/src/worker/activities/activity_heartbeat_manager.rs b/crates/sdk-core/src/worker/activities/activity_heartbeat_manager.rs index 7295e124d..8be4b8ba1 100644 --- a/crates/sdk-core/src/worker/activities/activity_heartbeat_manager.rs +++ b/crates/sdk-core/src/worker/activities/activity_heartbeat_manager.rs @@ -3,7 +3,7 @@ use crate::{ abstractions::take_cell::TakeCell, worker::{ activities::{PendingActivityCancel, make_payloads_too_large_failure}, - client::WorkerClient, + client::{WorkerClient, payload_limit_violation_from}, }, }; use futures_util::StreamExt; @@ -12,10 +12,10 @@ use std::{ sync::Arc, time::{Duration, Instant}, }; -use temporalio_client::payload_limit_violation_from; use temporalio_common::protos::{ coresdk::{ ActivityHeartbeat, IntoPayloadsExt, + activity_result::ActivityTaskFailedCause, activity_task::{ActivityCancelReason, ActivityCancellationDetails, ActivityTask}, }, temporal::api::{ @@ -201,6 +201,7 @@ impl ActivityHeartbeatManager { if let Err(fe) = sg .fail_activity_task( tt.clone(), + ActivityTaskFailedCause::PayloadsTooLarge, Some(make_payloads_too_large_failure(violation)), None, ) @@ -252,7 +253,7 @@ impl ActivityHeartbeatManager { ) -> Result<(), ActivityHeartbeatError> { self.heartbeat_tx .send(HeartbeatAction::SendHeartbeat(ValidActivityHeartbeat { - task_token: TaskToken(hb.task_token), + task_token: hb.task_token.into(), details: hb.details, throttle_interval, timeout_resetter, diff --git a/crates/sdk-core/src/worker/activities/local_activities.rs b/crates/sdk-core/src/worker/activities/local_activities.rs index bc5347b2c..4b0869336 100644 --- a/crates/sdk-core/src/worker/activities/local_activities.rs +++ b/crates/sdk-core/src/worker/activities/local_activities.rs @@ -2,7 +2,9 @@ use crate::{ MetricsContext, TaskToken, abstractions::{MeteredPermitDealer, OwnedMeteredSemPermit, UsedMeteredSemPermit, dbg_panic}, protosext::ValidScheduleLA, - telemetry::metrics::{activity_type, should_record_failure_metric, workflow_type}, + telemetry::metrics::{ + FailureReason, activity_type, failure_reason, should_record_failure_metric, workflow_type, + }, worker::{LocalActivitySlotKind, workflow::HeartbeatTimeoutMsg}, }; use futures_util::{ @@ -80,6 +82,7 @@ impl LocalActivityExecutionResult { )), ..Default::default() }), + ..Default::default() }) } @@ -528,7 +531,7 @@ impl LocalActivityManager { ]) .la_executed(); return Some(NextPendingLAAction::Dispatch(ActivityTask { - task_token: tt.0, + task_token: tt.into_inner(), variant: Some(activity_task::Variant::Start(Start { workflow_namespace: self.namespace.clone(), workflow_type: new_la.workflow_type, @@ -612,14 +615,18 @@ impl LocalActivityManager { let outcome = match &status { LocalActivityExecutionResult::Failed(fail) => { if should_record_failure_metric(&fail.failure) { - la_metrics.la_execution_failed() + la_metrics + .with_new_attrs([failure_reason(fail.cause().into())]) + .la_execution_failed() } Outcome::FailurePath { backoff: calc_backoff!(fail), } } LocalActivityExecutionResult::TimedOut(fail) => { - la_metrics.la_execution_failed(); + la_metrics + .with_new_attrs([failure_reason(FailureReason::Timeout)]) + .la_execution_failed(); is_timeout = true; // Start to close timeouts are retryable, other timeout types aren't. if matches!(status.get_timeout_type(), Some(TimeoutType::StartToClose)) { @@ -658,7 +665,7 @@ impl LocalActivityManager { // We want to generate a cancel task if the reason for failure was a timeout. let task = if is_timeout { Some(ActivityTask::cancel_from_ids( - task_token.clone().0, + task_token.clone().into_inner(), ActivityCancelReason::TimedOut, ActivityTask::primary_reason_to_cancellation_details( ActivityCancelReason::TimedOut, @@ -818,7 +825,7 @@ impl LocalActivityManager { self.cancels_req_tx .send(CancelOrTimeout::Cancel(ActivityTask::cancel_from_ids( - lai.task_token.0.clone(), + lai.task_token.clone().into_inner(), ActivityCancelReason::Cancelled, ActivityTask::primary_reason_to_cancellation_details( ActivityCancelReason::Cancelled, @@ -1055,7 +1062,7 @@ mod tests { activity_task::Variant::Start(Start {activity_id, ..}) if activity_id == i.to_string() ); - let next_tt = TaskToken(next.task_token); + let next_tt: TaskToken = next.task_token.into(); let complete_branch = async { lam.complete( &next_tt, @@ -1090,7 +1097,7 @@ mod tests { lam.workflows_have_shutdown(); let task = lam.next_pending().await.unwrap().unwrap(); - let task_token = TaskToken(task.task_token); + let task_token: TaskToken = task.task_token.into(); lam.complete( &task_token, LocalActivityExecutionResult::Completed(Default::default()), @@ -1114,7 +1121,7 @@ mod tests { .into()]); let next = lam.next_pending().await.unwrap().unwrap(); - let tt = TaskToken(next.task_token); + let tt: TaskToken = next.task_token.into(); tokio::select! { biased; @@ -1235,7 +1242,7 @@ mod tests { .into()]); let next = lam.next_pending().await.unwrap().unwrap(); - let tt = TaskToken(next.task_token); + let tt: TaskToken = next.task_token.into(); let res = lam.complete( &tt, LocalActivityExecutionResult::Failed(Default::default()), @@ -1270,7 +1277,7 @@ mod tests { .into()]); let next = lam.next_pending().await.unwrap().unwrap(); - let tt = TaskToken(next.task_token); + let tt: TaskToken = next.task_token.into(); let res = lam.complete( &tt, LocalActivityExecutionResult::Failed(ActFail { @@ -1284,6 +1291,7 @@ mod tests { )), ..Default::default() }), + ..Default::default() }), ); assert_matches!(res, LACompleteAction::Report { .. }); @@ -1317,7 +1325,7 @@ mod tests { .into()]); let next = lam.next_pending().await.unwrap().unwrap(); - let tt = TaskToken(next.task_token); + let tt: TaskToken = next.task_token.into(); lam.complete( &tt, LocalActivityExecutionResult::Failed(Default::default()), @@ -1364,7 +1372,7 @@ mod tests { .into()]); let next = lam.next_pending().await.unwrap().unwrap(); - let tt = TaskToken(next.task_token); + let tt: TaskToken = next.task_token.into(); lam.complete( &tt, LocalActivityExecutionResult::Failed(Default::default()), @@ -1516,7 +1524,7 @@ mod tests { let spinfail = || async { for _ in 1..=10 { let next = lam.next_pending().await.unwrap().unwrap(); - let tt = TaskToken(next.task_token); + let tt: TaskToken = next.task_token.into(); lam.complete( &tt, LocalActivityExecutionResult::Failed(Default::default()), diff --git a/crates/sdk-core/src/worker/client.rs b/crates/sdk-core/src/worker/client.rs index 4760205bd..856effc18 100644 --- a/crates/sdk-core/src/worker/client.rs +++ b/crates/sdk-core/src/worker/client.rs @@ -5,7 +5,10 @@ use crate::{ protosext::legacy_query_failure, worker::{WorkerVersioningStrategy, worker_control_task_queue}, }; +use backon::{BackoffBuilder, ExponentialBuilder}; +use futures_util::{StreamExt, TryStreamExt, stream}; use parking_lot::Mutex; +use prost::Message; use prost_types::Duration as PbDuration; use std::{ collections::HashMap, @@ -20,7 +23,11 @@ use temporalio_client::{ }; use temporalio_common::protos::{ TaskToken, - coresdk::{workflow_commands::QueryResult, workflow_completion}, + coresdk::{ + activity_result::ActivityTaskFailedCause, workflow_commands::QueryResult, + workflow_completion, + }, + google::rpc::Status as RpcStatus, temporal::api::{ command::v1::Command, common::v1::{ @@ -32,6 +39,7 @@ use temporalio_common::protos::{ TaskQueueKind, TaskQueueType, VersioningBehavior, WorkerVersioningMode, WorkflowTaskFailedCause, }, + errordetails::v1::WorkflowTaskCompletionBufferLostFailure, failure::v1::Failure, nexus::{self, v1::NexusTaskFailure}, protocol::v1::Message as ProtocolMessage, @@ -42,11 +50,148 @@ use temporalio_common::protos::{ workflowservice::v1::{get_system_info_response::Capabilities, *}, }, }; -use tonic::IntoRequest; +use tokio::time::sleep; +use tokio_util::sync::CancellationToken; +use tonic::{IntoRequest, metadata::MetadataValue}; use uuid::Uuid; type Result = std::result::Result; +pub(crate) fn payload_limit_violation_from( + status: &tonic::Status, +) -> Option<&temporalio_common::payload_limits::PayloadLimitViolation> { + std::error::Error::source(status).and_then(|source| source.downcast_ref()) +} + +/// Maximum encoded size of a single completion page, kept below the ~4 MiB gRPC frame limit. This +/// per-page cap is distinct from the server's namespace-wide limit on the recombined completion +/// size. +/// +/// Pages are packed by summing command body sizes only; the 512 KiB of headroom below 4 MiB absorbs +/// everything that sum omits: the per-request overhead (task token, identity, namespace) and the +/// per-command wire framing (a field tag plus a length varint, up to 6 bytes each). At the server's +/// default per-workflow history-count limit (~51,200 events), worst-case framing is ~300 KiB, so +/// this headroom covers even a page of many tiny commands and lets us skip per-command accounting. +const MAX_WFT_COMPLETION_PAGE_SIZE: usize = 4 * 1024 * 1024 - 512 * 1024; +// Conservative heuristic, not a tuned value: caps the client-side burst (concurrent request bodies +// and streams); the cost is only extra serial rounds for completions over this many pages. +const MAX_CONCURRENT_WFT_COMPLETION_PAGES: usize = 3; +// Backoff between resends of lost pages. Values are a conservative heuristic, not tuned; +// `without_max_times` leaves the number of resends to the loop (bounded by a stale token or +// shutdown), not the backoff. +const WFT_COMPLETION_PAGE_RESEND_BACKOFF: ExponentialBuilder = ExponentialBuilder::new() + .with_min_delay(Duration::from_millis(100)) + .with_factor(2.0) + .with_max_delay(Duration::from_secs(5)) + .without_max_times(); +/// Marker set on the error returned when a completion is failed proactively for exceeding the +/// namespace's recombined completion-size limit, so the workflow layer reports it as +/// `REQUEST_TOO_LARGE`. +pub(crate) static REQUEST_TOO_LARGE_KEY: &str = "request-too-large"; + +/// How a workflow task completion should be delivered, produced by [paginate_wft_completion]. +enum WftCompletionPages { + /// Send as a single request: it fits within a page, or it cannot be split. + Single(RespondWorkflowTaskCompletedRequest), + /// The server buffers only the commands of intermediate pages, so all messages and metadata + /// ride on the final page. + Paginated { + intermediate_pages: Vec, + final_page: RespondWorkflowTaskCompletedRequest, + }, +} + +/// Split a completion that may exceed `max_page_bytes` into pages that each stay under it, by +/// distributing its commands across intermediate pages in order. +/// +/// Falls back to [WftCompletionPages::Single] when the request already fits, has no commands to +/// distribute, or has a single command that alone exceeds a page (which the server then rejects). +fn paginate_wft_completion( + mut request: RespondWorkflowTaskCompletedRequest, + max_page_bytes: usize, +) -> WftCompletionPages { + if request.encoded_len() <= max_page_bytes { + return WftCompletionPages::Single(request); + } + + let intermediate_template = RespondWorkflowTaskCompletedRequest { + task_token: request.task_token.clone(), + identity: request.identity.clone(), + namespace: request.namespace.clone(), + intermediate_page: true, + ..Default::default() + }; + + // Pages are packed purely by command body size; MAX_WFT_COMPLETION_PAGE_SIZE reserves headroom + // for the per-request and per-command overhead this ignores. Only commands can be split across + // pages, so pagination cannot help when there are none, or when a single command alone exceeds + // a page. + if request.commands.is_empty() + || request + .commands + .iter() + .any(|c| c.encoded_len() > max_page_bytes) + { + return WftCompletionPages::Single(request); + } + + let commands = std::mem::take(&mut request.commands); + let mut intermediate = Vec::new(); + let mut current = Vec::new(); + let mut current_len = 0; + for command in commands { + let command_len = command.encoded_len(); + if !current.is_empty() && current_len + command_len > max_page_bytes { + let mut page = intermediate_template.clone(); + page.commands = std::mem::take(&mut current); + page.page_number = intermediate.len() as i32; + intermediate.push(page); + current_len = 0; + } + current_len += command_len; + current.push(command); + } + if !current.is_empty() { + let mut page = intermediate_template.clone(); + page.commands = current; + page.page_number = intermediate.len() as i32; + intermediate.push(page); + } + + request.page_number = intermediate.len() as i32; + request.intermediate_page = false; + WftCompletionPages::Paginated { + intermediate_pages: intermediate, + final_page: request, + } +} + +/// Returns true if `status` carries a `WorkflowTaskCompletionBufferLostFailure` detail, the +/// server's signal that it dropped the buffered pages and they must be resent from page 0. +fn is_workflow_task_completion_buffer_lost(status: &tonic::Status) -> bool { + RpcStatus::decode(status.details()) + .map(|rpc_status| { + rpc_status.details.iter().any(|detail| { + detail + .to_msg::() + .is_ok() + }) + }) + .unwrap_or(false) +} + +/// Wraps a completion page in a request that opts out of the client-layer retry for buffer loss, +/// which `complete_workflow_task` recovers itself by resending every page. +fn wft_completion_page_request( + page: RespondWorkflowTaskCompletedRequest, +) -> tonic::Request { + let mut request = page.into_request(); + request.extensions_mut().insert(NoRetryOnMatching { + predicate: is_workflow_task_completion_buffer_lost, + }); + request +} + /// The result of a legacy query sent via `respond_legacy_query`. pub enum LegacyQueryResult { /// The query handler returned a result successfully. @@ -182,6 +327,7 @@ pub trait WorkerClient: Sync + Send { async fn complete_workflow_task( &self, request: WorkflowTaskCompletion, + shutdown_token: CancellationToken, ) -> Result; /// Complete an activity task async fn complete_activity_task( @@ -211,6 +357,7 @@ pub trait WorkerClient: Sync + Send { async fn fail_activity_task( &self, task_token: TaskToken, + cause: ActivityTaskFailedCause, failure: Option, last_heartbeat_details: Option, ) -> Result; @@ -282,6 +429,10 @@ pub trait WorkerClient: Sync + Send { fn set_heartbeat_client_fields(&self, heartbeat: &mut WorkerHeartbeat); /// Set the worker's payload/memo error limits fn set_payload_error_limits(&self, _limits: Option) {} + /// Get the worker's payload/memo error limits + fn payload_error_limits(&self) -> Option { + None + } } /// Configuration options shared by workflow, activity, and Nexus polling calls @@ -457,7 +608,10 @@ impl WorkerClient for WorkerClientBag { async fn complete_workflow_task( &self, request: WorkflowTaskCompletion, + shutdown_token: CancellationToken, ) -> Result { + let pagination_enabled = request.pagination_enabled; + let wft_completion_size_limit = request.wft_completion_size_limit; #[allow(deprecated)] // want to list all fields explicitly let request = RespondWorkflowTaskCompletedRequest { task_token: request.task_token.into(), @@ -499,17 +653,98 @@ impl WorkerClient for WorkerClientBag { worker_instance_key: self.worker_instance_key.to_string(), worker_control_task_queue: self.worker_control_task_queue(), resource_id: Default::default(), - // Pagination fields: default to a single, final page. Pagination logic will - // populate these when splitting large completions. page_number: 0, intermediate_page: false, }; - Ok(self - .client - .clone() - .respond_workflow_task_completed(request.into_request()) - .await? - .into_inner()) + + let pages = if pagination_enabled { + paginate_wft_completion(request, MAX_WFT_COMPLETION_PAGE_SIZE) + } else { + WftCompletionPages::Single(request) + }; + let (intermediate_pages, final_page) = match pages { + WftCompletionPages::Single(request) => { + return Ok(self + .client + .clone() + .respond_workflow_task_completed(request.into_request()) + .await? + .into_inner()); + } + WftCompletionPages::Paginated { + intermediate_pages, + final_page, + } => (intermediate_pages, final_page), + }; + + // The server rejects the completion with REQUEST_TOO_LARGE and terminates the workflow once + // the buffered command bytes exceed the namespace limit, so fail here instead of sending + // doomed pages. Only buffered command bytes count toward that limit, not messages or + // metadata, which aren't buffered. + if let Some(limit) = wft_completion_size_limit { + let buffered_command_bytes: usize = intermediate_pages + .iter() + .flat_map(|page| page.commands.iter()) + .map(|command| command.encoded_len()) + .sum(); + if buffered_command_bytes > limit { + let mut status = tonic::Status::resource_exhausted( + "workflow task completion exceeds the namespace's recombined size limit", + ); + status + .metadata_mut() + .insert(REQUEST_TOO_LARGE_KEY, MetadataValue::from(0)); + return Err(status); + } + } + + // Buffer loss is transient, so resend the whole set from page 0 with exponential backoff. + // The server bounds the loop: once the task times out it starts a new attempt, and the next + // resend fails the token check with a non-buffer-lost error. Worker shutdown ends the loop + // sooner, so a stream of buffer losses can't hold shutdown's drain open until the task times + // out (a completion that is not resending still drains normally). Recovery has to live here + // because the client's retry layer would resend only the single failed page, which cannot + // rebuild the buffer the server dropped; it is told to pass buffer loss straight through + // (see `wft_completion_page_request`). + let mut backoff = WFT_COMPLETION_PAGE_RESEND_BACKOFF.build(); + loop { + let send_all = async { + // Cancel in-flight pages on the first error rather than awaiting them: any failure + // means we fail the task or resend from page 0, so the rest is wasted work. + stream::iter(intermediate_pages.iter().cloned()) + .map(|page| { + let mut client = self.client.clone(); + async move { + client + .respond_workflow_task_completed(wft_completion_page_request(page)) + .await + } + }) + .buffer_unordered(MAX_CONCURRENT_WFT_COMPLETION_PAGES) + .try_collect::>() + .await?; + // The final page must be sent only after every intermediate page has been + // buffered: it triggers the server-side merge, which requires pages 0..N-1 to all + // be present and otherwise returns a buffer-lost error. + self.client + .clone() + .respond_workflow_task_completed(wft_completion_page_request( + final_page.clone(), + )) + .await + }; + match send_all.await { + Ok(response) => return Ok(response.into_inner()), + Err(e) if is_workflow_task_completion_buffer_lost(&e) => { + let delay = backoff.next().expect("resend backoff is unbounded"); + tokio::select! { + _ = shutdown_token.cancelled() => return Err(e), + _ = sleep(delay) => {} + } + } + Err(e) => return Err(e), + } + } } async fn complete_activity_task( @@ -523,7 +758,7 @@ impl WorkerClient for WorkerClientBag { .respond_activity_task_completed( #[allow(deprecated)] // want to list all fields explicitly RespondActivityTaskCompletedRequest { - task_token: task_token.0, + task_token: task_token.into_inner(), result, identity: self.identity(), namespace: self.namespace.clone(), @@ -551,7 +786,7 @@ impl WorkerClient for WorkerClientBag { RespondNexusTaskCompletedRequest { namespace: self.namespace.clone(), identity: self.identity(), - task_token: task_token.0, + task_token: task_token.into_inner(), response: Some(response), poller_group_id: Default::default(), } @@ -571,7 +806,7 @@ impl WorkerClient for WorkerClientBag { .clone() .record_activity_task_heartbeat( RecordActivityTaskHeartbeatRequest { - task_token: task_token.0, + task_token: task_token.into_inner(), details, identity: self.identity(), namespace: self.namespace.clone(), @@ -594,7 +829,7 @@ impl WorkerClient for WorkerClientBag { .respond_activity_task_canceled( #[allow(deprecated)] // want to list all fields explicitly RespondActivityTaskCanceledRequest { - task_token: task_token.0, + task_token: task_token.into_inner(), details, identity: self.identity(), namespace: self.namespace.clone(), @@ -613,41 +848,17 @@ impl WorkerClient for WorkerClientBag { async fn fail_activity_task( &self, task_token: TaskToken, - mut failure: Option, - mut last_heartbeat_details: Option, + cause: ActivityTaskFailedCause, + failure: Option, + last_heartbeat_details: Option, ) -> Result { - let payload_error_limits = self.client.error_limits(); - if let (Some(details), Some(limits)) = - (last_heartbeat_details.as_ref(), payload_error_limits) - { - let heartbeat_request = RecordActivityTaskHeartbeatRequest { - details: Some(details.clone()), - ..Default::default() - }; - let payload_limits = temporalio_common::payload_limits::PayloadLimits { - blob_error: limits.blob, - memo_error: limits.memo, - ..Default::default() - }; - if let Some(violation) = - temporalio_common::payload_limits::validate_known_payload_limits( - &heartbeat_request, - &payload_limits, - ) - { - failure = Some(crate::worker::activities::make_payloads_too_large_failure( - &violation, - )); - last_heartbeat_details = None; - } - } Ok(self .client .clone() .respond_activity_task_failed( #[allow(deprecated)] // want to list all fields explicitly RespondActivityTaskFailedRequest { - task_token: task_token.0, + task_token: task_token.into_inner(), failure, identity: self.identity(), namespace: self.namespace.clone(), @@ -657,6 +868,7 @@ impl WorkerClient for WorkerClientBag { deployment: None, deployment_options: self.deployment_options(), resource_id: Default::default(), + cause: cause as i32, } .into_request(), ) @@ -672,7 +884,7 @@ impl WorkerClient for WorkerClientBag { ) -> Result { #[allow(deprecated)] // want to list all fields explicitly let request = RespondWorkflowTaskFailedRequest { - task_token: task_token.0, + task_token: task_token.into_inner(), cause: cause as i32, failure, identity: self.identity(), @@ -711,7 +923,7 @@ impl WorkerClient for WorkerClientBag { RespondNexusTaskFailedRequest { namespace: self.namespace.clone(), identity: self.identity(), - task_token: task_token.0, + task_token: task_token.into_inner(), failure, error, poller_group_id: Default::default(), @@ -938,6 +1150,10 @@ impl WorkerClient for WorkerClientBag { fn set_payload_error_limits(&self, limits: Option) { self.client.set_error_limits(limits); } + + fn payload_error_limits(&self) -> Option { + self.client.error_limits() + } } impl NamespacedClient for WorkerClientBag { @@ -952,7 +1168,7 @@ impl NamespacedClient for WorkerClientBag { /// A version of [RespondWorkflowTaskCompletedRequest] that will finish being filled out by the /// server client -#[derive(Debug, Clone, PartialEq)] +#[derive(Debug, Clone)] pub struct WorkflowTaskCompletion { /// The task token that would've been received from polling for a workflow activation pub task_token: TaskToken, @@ -974,6 +1190,13 @@ pub struct WorkflowTaskCompletion { pub metering_metadata: MeteringMetadata, /// Versioning behavior of the workflow, if any. pub versioning_behavior: VersioningBehavior, + /// Whether the namespace permits paginating this completion across multiple page requests when + /// it would otherwise exceed the server's gRPC request size limit. + pub pagination_enabled: bool, + /// The namespace's limit on the recombined size of a paginated completion, if the server + /// advertises one. A paginated completion larger than this is rejected server-side with + /// `REQUEST_TOO_LARGE`, so the worker fails it proactively instead of sending doomed pages. + pub wft_completion_size_limit: Option, } #[derive(Clone, Default)] @@ -1068,7 +1291,8 @@ mod tests { client .fail_activity_task( - TaskToken(vec![1]), + vec![1].into(), + ActivityTaskFailedCause::ActivityWorkerUnhandledFailure, None, Some(last_heartbeat_details.clone()), ) @@ -1088,10 +1312,12 @@ mod tests { ( "deployment", WorkerVersioningStrategy::WorkerDeploymentBased( - WorkerDeploymentOptions::new(WorkerDeploymentVersion { - deployment_name: "deployment".to_string(), - build_id: "deployment-build".to_string(), - }) + WorkerDeploymentOptions::new( + WorkerDeploymentVersion::builder() + .deployment_name("deployment".to_string()) + .build_id("deployment-build".to_string()) + .build(), + ) .use_worker_versioning(true) .build(), ), @@ -1216,4 +1442,645 @@ mod tests { ); } } + + mod pagination { + use super::*; + use temporalio_common::protos::temporal::api::{ + command::v1::{CompleteWorkflowExecutionCommandAttributes, command}, + common::v1::{Payload, Payloads}, + errordetails::v1::WorkflowExecutionAlreadyStartedFailure, + }; + + fn command_with_payload(data_size: usize) -> Command { + Command { + attributes: Some( + command::Attributes::CompleteWorkflowExecutionCommandAttributes( + CompleteWorkflowExecutionCommandAttributes { + result: Some(Payloads { + payloads: vec![Payload { + metadata: Default::default(), + data: vec![0u8; data_size], + ..Default::default() + }], + }), + }, + ), + ), + ..Default::default() + } + } + + fn request_with(commands: Vec) -> RespondWorkflowTaskCompletedRequest { + RespondWorkflowTaskCompletedRequest { + task_token: b"task-token".to_vec(), + identity: "identity".to_string(), + namespace: "namespace".to_string(), + commands, + ..Default::default() + } + } + + #[test] + fn completion_within_limit_is_a_single_final_page() { + let request = request_with(vec![command_with_payload(16)]); + let WftCompletionPages::Single(page) = paginate_wft_completion(request, 4096) else { + panic!("expected a single page"); + }; + assert_eq!(page.page_number, 0); + assert!(!page.intermediate_page); + assert_eq!(page.commands.len(), 1); + } + + #[test] + fn large_completion_splits_commands_across_pages() { + let max = 1024; + let command_count = 6; + let commands: Vec<_> = (0..command_count) + .map(|_| command_with_payload(400)) + .collect(); + let request = request_with(commands); + assert!(request.encoded_len() > max); + + let WftCompletionPages::Paginated { + intermediate_pages: intermediate, + final_page, + } = paginate_wft_completion(request, max) + else { + panic!("expected multiple pages"); + }; + + assert!(!final_page.intermediate_page); + assert!(final_page.commands.is_empty()); + assert_eq!(final_page.page_number as usize, intermediate.len()); + assert!(final_page.encoded_len() <= max); + assert_eq!(final_page.task_token, b"task-token"); + + let mut total_commands = 0; + for (idx, page) in intermediate.iter().enumerate() { + assert!(page.intermediate_page); + assert_eq!(page.page_number as usize, idx); + assert_eq!(page.task_token, b"task-token"); + assert!( + page.encoded_len() <= max, + "intermediate page {idx} over limit" + ); + total_commands += page.commands.len(); + } + // Every command is preserved exactly once across the intermediate pages. + assert_eq!(total_commands, command_count); + } + + #[test] + fn single_command_larger_than_a_page_is_not_split() { + let max = 1024; + let request = request_with(vec![command_with_payload(4096)]); + // Cannot be split, so it is left as one (oversized) request for the server to reject. + let WftCompletionPages::Single(page) = paginate_wft_completion(request, max) else { + panic!("expected a single page"); + }; + assert_eq!(page.commands.len(), 1); + assert!(!page.intermediate_page); + } + + // Pack the detail with `Any::from_msg` so its `type_url` is derived from the message name, + // the way the server sets it, rather than a hand-written string. + fn status_with_detail(detail: &M) -> tonic::Status { + let rpc_status = RpcStatus { + code: tonic::Code::Aborted as i32, + message: String::new(), + details: vec![prost_types::Any::from_msg(detail).expect("detail encodes")], + }; + tonic::Status::with_details(tonic::Code::Aborted, "", rpc_status.encode_to_vec().into()) + } + + #[test] + fn detects_buffer_lost_failure_detail() { + let status = status_with_detail(&WorkflowTaskCompletionBufferLostFailure {}); + assert!(is_workflow_task_completion_buffer_lost(&status)); + + let unrelated = tonic::Status::new(tonic::Code::Internal, "boom"); + assert!(!is_workflow_task_completion_buffer_lost(&unrelated)); + } + + #[test] + fn buffer_lost_detection_ignores_unrelated_detail() { + // A different error detail carried on the same gRPC code must not be mistaken for a + // buffer-lost failure. + let status = status_with_detail(&WorkflowExecutionAlreadyStartedFailure { + start_request_id: "req".to_string(), + run_id: "run".to_string(), + ..Default::default() + }); + assert!(!is_workflow_task_completion_buffer_lost(&status)); + } + + #[tokio::test] + async fn paginated_completion_sends_ordered_pages_sharing_a_token() { + let captured = Arc::new(Mutex::new(Vec::new())); + let captured_clone = captured.clone(); + let service_override = CallbackBasedGrpcService { + callback: Arc::new(move |request| { + let captured = captured_clone.clone(); + Box::pin(async move { + let proto = match request.rpc.as_str() { + "GetSystemInfo" => GetSystemInfoResponse { + capabilities: Some(Capabilities::default()), + ..Default::default() + } + .encode_to_vec(), + "RespondWorkflowTaskCompleted" => { + captured.lock().unwrap().push( + RespondWorkflowTaskCompletedRequest::decode(request.proto) + .expect("completion request is valid"), + ); + RespondWorkflowTaskCompletedResponse::default().encode_to_vec() + } + rpc => panic!("unexpected RPC: {rpc}"), + }; + Ok(GrpcSuccessResponse { + headers: Default::default(), + proto, + }) + }) + }), + }; + let connection = Connection::connect( + ConnectionOptions::new(url::Url::parse("http://localhost:7233").unwrap()) + .service_override(service_override) + .dns_load_balancing(None) + .build(), + ) + .await + .unwrap(); + let client = WorkerClientBag::new( + SharedReplaceableClient::new(connection), + "namespace".to_string(), + WorkerVersioningStrategy::LegacyBuildIdBased { + build_id: "test-build".to_string(), + }, + Uuid::new_v4(), + ); + + // Roughly 4 MiB of commands forces splitting under the ~3 MiB page target. + let commands: Vec<_> = (0..8).map(|_| command_with_payload(512 * 1024)).collect(); + let completion = WorkflowTaskCompletion { + task_token: b"shared-token".to_vec().into(), + commands, + messages: vec![], + sticky_attributes: None, + query_responses: vec![], + return_new_workflow_task: false, + force_create_new_workflow_task: false, + sdk_metadata: Default::default(), + metering_metadata: Default::default(), + versioning_behavior: VersioningBehavior::Unspecified, + pagination_enabled: true, + wft_completion_size_limit: None, + }; + client + .complete_workflow_task(completion, CancellationToken::new()) + .await + .unwrap(); + + let sent = captured.lock().unwrap(); + assert!( + sent.len() >= 2, + "expected multiple pages, got {}", + sent.len() + ); + // Every page shares the one task token. + assert!(sent.iter().all(|r| r.task_token == b"shared-token")); + // Exactly one final page, numbered after all the intermediate ones. + let finals: Vec<_> = sent.iter().filter(|r| !r.intermediate_page).collect(); + assert_eq!(finals.len(), 1); + assert_eq!(finals[0].page_number as usize, sent.len() - 1); + assert!(finals[0].commands.is_empty()); + // Intermediate pages carry sequential page numbers 0..N-1. + let mut intermediate_numbers: Vec<_> = sent + .iter() + .filter(|r| r.intermediate_page) + .map(|r| r.page_number) + .collect(); + intermediate_numbers.sort_unstable(); + assert_eq!( + intermediate_numbers, + (0..(sent.len() as i32 - 1)).collect::>() + ); + } + + #[tokio::test] + async fn failed_page_cancels_other_inflight_pages() { + // Page 0 fails immediately; every other intermediate page hangs forever. The call can + // only return if the failed page short-circuits the send and the hung pages are + // dropped (cancelled) rather than awaited. + let never = Arc::new(tokio::sync::Notify::new()); + let never_cb = never.clone(); + let service_override = CallbackBasedGrpcService { + callback: Arc::new(move |request| { + let never = never_cb.clone(); + Box::pin(async move { + match request.rpc.as_str() { + "GetSystemInfo" => Ok(GrpcSuccessResponse { + headers: Default::default(), + proto: GetSystemInfoResponse { + capabilities: Some(Capabilities::default()), + ..Default::default() + } + .encode_to_vec(), + }), + "RespondWorkflowTaskCompleted" => { + let page = + RespondWorkflowTaskCompletedRequest::decode(request.proto) + .expect("completion request is valid"); + if page.intermediate_page && page.page_number == 0 { + // InvalidArgument is non-retryable, so it is forwarded at once. + Err(tonic::Status::new(tonic::Code::InvalidArgument, "boom")) + } else { + never.notified().await; + unreachable!("a cancelled page must not resume"); + } + } + rpc => panic!("unexpected RPC: {rpc}"), + } + }) + }), + }; + let connection = Connection::connect( + ConnectionOptions::new(url::Url::parse("http://localhost:7233").unwrap()) + .service_override(service_override) + .dns_load_balancing(None) + .build(), + ) + .await + .unwrap(); + let client = WorkerClientBag::new( + SharedReplaceableClient::new(connection), + "namespace".to_string(), + WorkerVersioningStrategy::LegacyBuildIdBased { + build_id: "test-build".to_string(), + }, + Uuid::new_v4(), + ); + + // Enough commands to yield at least two intermediate pages (one fails, one hangs). + let commands: Vec<_> = (0..8).map(|_| command_with_payload(512 * 1024)).collect(); + let completion = WorkflowTaskCompletion { + task_token: b"shared-token".to_vec().into(), + commands, + messages: vec![], + sticky_attributes: None, + query_responses: vec![], + return_new_workflow_task: false, + force_create_new_workflow_task: false, + sdk_metadata: Default::default(), + metering_metadata: Default::default(), + versioning_behavior: VersioningBehavior::Unspecified, + pagination_enabled: true, + wft_completion_size_limit: None, + }; + + // Without cancellation this would hang on the never-completing page; the timeout guards + // against that regression instead of relying on a sleep. + let outcome = tokio::time::timeout( + Duration::from_secs(10), + client.complete_workflow_task(completion, CancellationToken::new()), + ) + .await + .expect("completion resolved without waiting on the hung page"); + assert!( + outcome.is_err(), + "the failed page should surface as an error" + ); + } + + #[tokio::test] + async fn completion_over_namespace_limit_fails_proactively_without_sending() { + let sent = Arc::new(Mutex::new(0usize)); + let sent_cb = sent.clone(); + let service_override = CallbackBasedGrpcService { + callback: Arc::new(move |request| { + let sent = sent_cb.clone(); + Box::pin(async move { + let proto = match request.rpc.as_str() { + "GetSystemInfo" => GetSystemInfoResponse { + capabilities: Some(Capabilities::default()), + ..Default::default() + } + .encode_to_vec(), + "RespondWorkflowTaskCompleted" => { + *sent.lock().unwrap() += 1; + RespondWorkflowTaskCompletedResponse::default().encode_to_vec() + } + rpc => panic!("unexpected RPC: {rpc}"), + }; + Ok(GrpcSuccessResponse { + headers: Default::default(), + proto, + }) + }) + }), + }; + let connection = Connection::connect( + ConnectionOptions::new(url::Url::parse("http://localhost:7233").unwrap()) + .service_override(service_override) + .dns_load_balancing(None) + .build(), + ) + .await + .unwrap(); + let client = WorkerClientBag::new( + SharedReplaceableClient::new(connection), + "namespace".to_string(), + WorkerVersioningStrategy::LegacyBuildIdBased { + build_id: "test-build".to_string(), + }, + Uuid::new_v4(), + ); + + // ~4 MiB total (so it would be paginated) but the namespace caps the recombined size + // at 1 MiB, so the server would reject it, and the worker must fail it without sending. + let commands: Vec<_> = (0..8).map(|_| command_with_payload(512 * 1024)).collect(); + let completion = WorkflowTaskCompletion { + task_token: b"shared-token".to_vec().into(), + commands, + messages: vec![], + sticky_attributes: None, + query_responses: vec![], + return_new_workflow_task: false, + force_create_new_workflow_task: false, + sdk_metadata: Default::default(), + metering_metadata: Default::default(), + versioning_behavior: VersioningBehavior::Unspecified, + pagination_enabled: true, + wft_completion_size_limit: Some(1024 * 1024), + }; + let err = client + .complete_workflow_task(completion, CancellationToken::new()) + .await + .expect_err("completion over the namespace limit must fail"); + assert!(err.metadata().contains_key(REQUEST_TOO_LARGE_KEY)); + assert_eq!(*sent.lock().unwrap(), 0, "no pages should have been sent"); + } + + #[tokio::test] + async fn buffer_loss_resends_all_pages_until_it_succeeds() { + // The server reports buffer loss on the final page of the first two attempts, then + // accepts the third. The whole set must be resent from page 0 each time, and the + // completion must ultimately succeed. Because buffer loss is marked non-retryable at the + // client layer, each attempt sends the final page exactly once, so a count of three + // proves the resend loop, not the client's retry policy, did the retrying. + let final_attempts = Arc::new(Mutex::new(0usize)); + let final_attempts_cb = final_attempts.clone(); + let service_override = CallbackBasedGrpcService { + callback: Arc::new(move |request| { + let final_attempts = final_attempts_cb.clone(); + Box::pin(async move { + match request.rpc.as_str() { + "GetSystemInfo" => Ok(GrpcSuccessResponse { + headers: Default::default(), + proto: GetSystemInfoResponse { + capabilities: Some(Capabilities::default()), + ..Default::default() + } + .encode_to_vec(), + }), + "RespondWorkflowTaskCompleted" => { + let page = + RespondWorkflowTaskCompletedRequest::decode(request.proto) + .expect("completion request is valid"); + if !page.intermediate_page { + let mut attempts = final_attempts.lock().unwrap(); + *attempts += 1; + if *attempts <= 2 { + return Err(status_with_detail( + &WorkflowTaskCompletionBufferLostFailure {}, + )); + } + } + Ok(GrpcSuccessResponse { + headers: Default::default(), + proto: RespondWorkflowTaskCompletedResponse::default() + .encode_to_vec(), + }) + } + rpc => panic!("unexpected RPC: {rpc}"), + } + }) + }), + }; + let connection = Connection::connect( + ConnectionOptions::new(url::Url::parse("http://localhost:7233").unwrap()) + .service_override(service_override) + .dns_load_balancing(None) + .build(), + ) + .await + .unwrap(); + let client = WorkerClientBag::new( + SharedReplaceableClient::new(connection), + "namespace".to_string(), + WorkerVersioningStrategy::LegacyBuildIdBased { + build_id: "test-build".to_string(), + }, + Uuid::new_v4(), + ); + + let commands: Vec<_> = (0..8).map(|_| command_with_payload(512 * 1024)).collect(); + let completion = WorkflowTaskCompletion { + task_token: b"shared-token".to_vec().into(), + commands, + messages: vec![], + sticky_attributes: None, + query_responses: vec![], + return_new_workflow_task: false, + force_create_new_workflow_task: false, + sdk_metadata: Default::default(), + metering_metadata: Default::default(), + versioning_behavior: VersioningBehavior::Unspecified, + pagination_enabled: true, + wft_completion_size_limit: None, + }; + client + .complete_workflow_task(completion, CancellationToken::new()) + .await + .expect("completion eventually succeeds after the buffer is re-established"); + assert_eq!( + *final_attempts.lock().unwrap(), + 3, + "the final page is sent once per resend, with no client-layer retry" + ); + } + + #[tokio::test] + async fn shutdown_stops_buffer_loss_resends() { + // The server never re-establishes the buffer, so without shutdown handling the resend + // loop would run until the task times out on the server, potentially minutes, holding + // shutdown open. A cancelled shutdown token must end it promptly instead. The token is + // cancelled before the call, so the first attempt still sends fully (in-flight work is + // never abandoned) and the loop bails as soon as that attempt reports buffer loss. + let final_attempts = Arc::new(Mutex::new(0usize)); + let final_attempts_cb = final_attempts.clone(); + let service_override = CallbackBasedGrpcService { + callback: Arc::new(move |request| { + let final_attempts = final_attempts_cb.clone(); + Box::pin(async move { + match request.rpc.as_str() { + "GetSystemInfo" => Ok(GrpcSuccessResponse { + headers: Default::default(), + proto: GetSystemInfoResponse { + capabilities: Some(Capabilities::default()), + ..Default::default() + } + .encode_to_vec(), + }), + "RespondWorkflowTaskCompleted" => { + let page = + RespondWorkflowTaskCompletedRequest::decode(request.proto) + .expect("completion request is valid"); + if page.intermediate_page { + Ok(GrpcSuccessResponse { + headers: Default::default(), + proto: RespondWorkflowTaskCompletedResponse::default() + .encode_to_vec(), + }) + } else { + *final_attempts.lock().unwrap() += 1; + Err(status_with_detail( + &WorkflowTaskCompletionBufferLostFailure {}, + )) + } + } + rpc => panic!("unexpected RPC: {rpc}"), + } + }) + }), + }; + let connection = Connection::connect( + ConnectionOptions::new(url::Url::parse("http://localhost:7233").unwrap()) + .service_override(service_override) + .dns_load_balancing(None) + .build(), + ) + .await + .unwrap(); + let client = WorkerClientBag::new( + SharedReplaceableClient::new(connection), + "namespace".to_string(), + WorkerVersioningStrategy::LegacyBuildIdBased { + build_id: "test-build".to_string(), + }, + Uuid::new_v4(), + ); + + let shutdown_token = CancellationToken::new(); + shutdown_token.cancel(); + let commands: Vec<_> = (0..8).map(|_| command_with_payload(512 * 1024)).collect(); + let completion = WorkflowTaskCompletion { + task_token: b"shared-token".to_vec().into(), + commands, + messages: vec![], + sticky_attributes: None, + query_responses: vec![], + return_new_workflow_task: false, + force_create_new_workflow_task: false, + sdk_metadata: Default::default(), + metering_metadata: Default::default(), + versioning_behavior: VersioningBehavior::Unspecified, + pagination_enabled: true, + wft_completion_size_limit: None, + }; + + let err = tokio::time::timeout( + Duration::from_secs(10), + client.complete_workflow_task(completion, shutdown_token), + ) + .await + .expect("shutdown ends the resend loop instead of waiting for the server timeout") + .expect_err("the last buffer-loss error is surfaced"); + assert!(is_workflow_task_completion_buffer_lost(&err)); + assert_eq!( + *final_attempts.lock().unwrap(), + 1, + "the first attempt still sends; shutdown prevents any resend" + ); + } + + #[tokio::test] + async fn cancelled_shutdown_does_not_interrupt_successful_completion() { + // The shutdown token is only consulted while resending after buffer loss. A completion + // that never hits buffer loss must still succeed even when the worker is shutting down, + // so graceful drain can finish outstanding completions rather than abandon them. + let final_pages = Arc::new(Mutex::new(0usize)); + let final_pages_cb = final_pages.clone(); + let service_override = CallbackBasedGrpcService { + callback: Arc::new(move |request| { + let final_pages = final_pages_cb.clone(); + Box::pin(async move { + let proto = match request.rpc.as_str() { + "GetSystemInfo" => GetSystemInfoResponse { + capabilities: Some(Capabilities::default()), + ..Default::default() + } + .encode_to_vec(), + "RespondWorkflowTaskCompleted" => { + let page = + RespondWorkflowTaskCompletedRequest::decode(request.proto) + .expect("completion request is valid"); + if !page.intermediate_page { + *final_pages.lock().unwrap() += 1; + } + RespondWorkflowTaskCompletedResponse::default().encode_to_vec() + } + rpc => panic!("unexpected RPC: {rpc}"), + }; + Ok(GrpcSuccessResponse { + headers: Default::default(), + proto, + }) + }) + }), + }; + let connection = Connection::connect( + ConnectionOptions::new(url::Url::parse("http://localhost:7233").unwrap()) + .service_override(service_override) + .dns_load_balancing(None) + .build(), + ) + .await + .unwrap(); + let client = WorkerClientBag::new( + SharedReplaceableClient::new(connection), + "namespace".to_string(), + WorkerVersioningStrategy::LegacyBuildIdBased { + build_id: "test-build".to_string(), + }, + Uuid::new_v4(), + ); + + let shutdown_token = CancellationToken::new(); + shutdown_token.cancel(); + let commands: Vec<_> = (0..8).map(|_| command_with_payload(512 * 1024)).collect(); + let completion = WorkflowTaskCompletion { + task_token: b"shared-token".to_vec().into(), + commands, + messages: vec![], + sticky_attributes: None, + query_responses: vec![], + return_new_workflow_task: false, + force_create_new_workflow_task: false, + sdk_metadata: Default::default(), + metering_metadata: Default::default(), + versioning_behavior: VersioningBehavior::Unspecified, + pagination_enabled: true, + wft_completion_size_limit: None, + }; + + client + .complete_workflow_task(completion, shutdown_token) + .await + .expect("a completion without buffer loss succeeds despite shutdown"); + // The paginated path ran to its final page rather than being cut short. + assert_eq!(*final_pages.lock().unwrap(), 1); + } + } } diff --git a/crates/sdk-core/src/worker/client/mocks.rs b/crates/sdk-core/src/worker/client/mocks.rs index 88cf9fbab..421662899 100644 --- a/crates/sdk-core/src/worker/client/mocks.rs +++ b/crates/sdk-core/src/worker/client/mocks.rs @@ -23,6 +23,7 @@ pub(crate) static DEFAULT_TEST_CAPABILITIES: &Capabilities = &Capabilities { /// Create a mock client primed with basic necessary expectations pub fn mock_worker_client() -> MockWorkerClient { let mut r = MockWorkerClient::new(); + r.expect_payload_error_limits().returning(|| None); let workers = Arc::new(ClientWorkerSet::new()); r.expect_capabilities() .returning(|| Some(*DEFAULT_TEST_CAPABILITIES)); @@ -86,6 +87,7 @@ mockall::mock! { fn complete_workflow_task<'a, 'b>( &self, request: WorkflowTaskCompletion, + shutdown_token: CancellationToken, ) -> impl Future> + Send + 'b where 'a: 'b, Self: 'b; @@ -113,6 +115,7 @@ mockall::mock! { fn fail_activity_task<'a, 'b>( &self, task_token: TaskToken, + cause: ActivityTaskFailedCause, failure: Option, last_heartbeat_details: Option, ) -> impl Future> + Send + 'b diff --git a/crates/sdk-core/src/worker/heartbeat.rs b/crates/sdk-core/src/worker/heartbeat.rs index fedcf28f2..b3fa1047e 100644 --- a/crates/sdk-core/src/worker/heartbeat.rs +++ b/crates/sdk-core/src/worker/heartbeat.rs @@ -303,7 +303,7 @@ async fn handle_worker_command_task( for command in &exec_req.commands { let result_type = match &command.r#type { Some(WorkerCommandType::CancelActivity(cancel_cmd)) => { - let tt = TaskToken(cancel_cmd.task_token.clone()); + let tt: TaskToken = cancel_cmd.task_token.clone().into(); let cancel_callbacks: Vec<_> = callbacks_map .read() .values() @@ -437,24 +437,20 @@ mod tests { .unwrap(); shared_worker.register_callback( Uuid::new_v4(), - WorkerCallbacks { - heartbeat: Arc::new(|| None), - heartbeat_success: None, - cancel_activity: None, - }, + WorkerCallbacks::new(Arc::new(|| None), None, None), ); shared_worker.register_callback( started_worker_key, - WorkerCallbacks { - heartbeat: Arc::new(move || { + WorkerCallbacks::new( + Arc::new(move || { Some(WorkerHeartbeat { worker_instance_key: started_worker_key.to_string(), ..Default::default() }) }), - heartbeat_success: None, - cancel_activity: None, - }, + None, + None, + ), ); tokio::time::timeout(Duration::from_secs(5), recorded_rx) @@ -659,16 +655,16 @@ mod tests { let callback_tx = Mutex::new(Some(callback_tx)); shared_worker.register_callback( Uuid::new_v4(), - WorkerCallbacks { - heartbeat: Arc::new(move || { + WorkerCallbacks::new( + Arc::new(move || { if let Some(tx) = callback_tx.lock().unwrap().take() { let _ = tx.send(()); } None }), - heartbeat_success: None, - cancel_activity: None, - }, + None, + None, + ), ); tokio::time::timeout(Duration::from_secs(5), callback_rx) .await diff --git a/crates/sdk-core/src/worker/mod.rs b/crates/sdk-core/src/worker/mod.rs index 9f5cdbab2..8d3627156 100644 --- a/crates/sdk-core/src/worker/mod.rs +++ b/crates/sdk-core/src/worker/mod.rs @@ -320,16 +320,16 @@ impl WorkerConfig { pub(crate) fn computed_deployment_version(&self) -> Option { let wdv = match self.versioning_strategy { - WorkerVersioningStrategy::None { ref build_id } => WorkerDeploymentVersion { - deployment_name: "".to_owned(), - build_id: build_id.clone(), - }, + WorkerVersioningStrategy::None { ref build_id } => WorkerDeploymentVersion::builder() + .deployment_name("") + .build_id(build_id.clone()) + .build(), WorkerVersioningStrategy::WorkerDeploymentBased(ref opts) => opts.version.clone(), WorkerVersioningStrategy::LegacyBuildIdBased { ref build_id } => { - WorkerDeploymentVersion { - deployment_name: "".to_owned(), - build_id: build_id.clone(), - } + WorkerDeploymentVersion::builder() + .deployment_name("") + .build_id(build_id.clone()) + .build() } }; if wdv.is_empty() { None } else { Some(wdv) } @@ -529,6 +529,25 @@ impl NamespaceCapabilities { self.capabilities() .is_some_and(|capabilities| capabilities.worker_commands) } + + /// Returns true if the namespace accepts paginated `RespondWorkflowTaskCompleted` requests, so + /// large completions may be split across multiple page requests sharing one task token. + pub fn workflow_task_completion_pagination(&self) -> bool { + self.capabilities() + .is_some_and(|capabilities| capabilities.workflow_task_completion_pagination) + } + + /// The namespace's limit on the recombined size of a paginated workflow task completion, if one + /// is configured. `None` when unset (the server advertises `0` for no explicit limit). + pub fn workflow_task_completion_size_limit(&self) -> Option { + self.description + .get() + .and_then(|description| description.namespace_info.as_ref()) + .and_then(|namespace_info| namespace_info.limits.as_ref()) + .map(|limits| limits.workflow_task_completion_size_limit_error) + .filter(|limit| *limit > 0) + .map(|limit| limit as usize) + } } /// Resolve the effective poller behavior. When no behavior was configured (`None`), pollers are @@ -978,7 +997,7 @@ impl Worker { shutdown_token.child_token(), Some(move |np| np_metrics.record_num_pollers(np)), nexus_last_suc_poll_time, - capabilities, + capabilities.clone(), shared_namespace_worker, )) as BoxedNexusPoller) } else { @@ -1079,6 +1098,7 @@ impl Worker { worker_shutdown_token: shutdown_token.clone(), metrics, server_capabilities: client.capabilities().unwrap_or_default(), + namespace_capabilities: capabilities.clone(), sdk_name: sdk_name_and_ver.0, sdk_version: sdk_name_and_ver.1, default_versioning_behavior: config @@ -1394,7 +1414,7 @@ impl Worker { /// options. pub fn record_activity_heartbeat(&self, details: ActivityHeartbeat) { if let Some(at_mgr) = self.task_subsystems.at_task_mgr.as_ref() { - let tt = TaskToken(details.task_token.clone()); + let tt: TaskToken = details.task_token.clone().into(); if let Err(e) = at_mgr.record_heartbeat(details) { warn!(task_token = %tt, details = ?e, "Activity heartbeat failed."); } @@ -1520,7 +1540,7 @@ impl Worker { &self, completion: ActivityTaskCompletion, ) -> Result<(), CompleteActivityError> { - let task_token = TaskToken(completion.task_token); + let task_token: TaskToken = completion.task_token.into(); let status = if let Some(s) = completion.result.and_then(|r| r.status) { s } else { @@ -1645,7 +1665,7 @@ impl Worker { reason: "Nexus completion had empty status field".to_owned(), }); }; - let tt = TaskToken(completion.task_token); + let tt: TaskToken = completion.task_token.into(); tracing::Span::current().record("task_token", tt.to_string()); tracing::Span::current().record("status", status.to_string()); @@ -2891,10 +2911,12 @@ mod tests { .namespace("default") .task_queue("test-queue") .versioning_strategy(WorkerVersioningStrategy::WorkerDeploymentBased( - WorkerDeploymentOptions::new(WorkerDeploymentVersion { - deployment_name: "deployment".to_string(), - build_id: "1.0".to_string(), - }) + WorkerDeploymentOptions::new( + WorkerDeploymentVersion::builder() + .deployment_name("deployment") + .build_id("1.0") + .build(), + ) .default_versioning_behavior(VersioningBehavior::AutoUpgrade.into()) .build(), )) diff --git a/crates/sdk-core/src/worker/nexus.rs b/crates/sdk-core/src/worker/nexus.rs index a8e39ce8d..732b12384 100644 --- a/crates/sdk-core/src/worker/nexus.rs +++ b/crates/sdk-core/src/worker/nexus.rs @@ -2,7 +2,10 @@ use crate::{ abstractions::UsedMeteredSemPermit, pollers::{BoxedNexusPoller, NexusPollItem, new_nexus_task_poller}, telemetry::metrics::{self, FailureReason, MetricsContext}, - worker::{CompleteNexusError, NexusSlotKind, PollError, client::WorkerClient}, + worker::{ + CompleteNexusError, NexusSlotKind, PollError, + client::{WorkerClient, payload_limit_violation_from}, + }, }; use anyhow::anyhow; use futures_util::{ @@ -18,7 +21,6 @@ use std::{ }, time::{Duration, Instant, SystemTime}, }; -use temporalio_client::payload_limit_violation_from; use temporalio_common::{ payload_limits::PayloadLimitViolation, protos::{ @@ -398,7 +400,7 @@ where } } - let tt = TaskToken(t.resp.task_token.clone()); + let tt: TaskToken = t.resp.task_token.clone().into(); let mut timeout_task = None; let mut request_deadline: Option = None; if let Some(timeout_str) = t @@ -418,7 +420,7 @@ where "Timing out nexus task due to elapsed local timeout timer" ); let _ = cancels_tx.send(CancelNexusTask { - task_token: tt_clone.0, + task_token: tt_clone.into_inner(), reason: NexusTaskCancelReason::TimedOut.into(), }); })); @@ -508,7 +510,7 @@ where tokio::time::sleep(gp).await; for (tt, _) in outstanding_task_clone.lock().iter() { let _ = cancels_tx_clone.send(CancelNexusTask { - task_token: tt.0.clone(), + task_token: tt.clone().into_inner(), reason: NexusTaskCancelReason::WorkerShutdown.into(), }); } diff --git a/crates/sdk-core/src/worker/tuner/resource_based.rs b/crates/sdk-core/src/worker/tuner/resource_based.rs index 4ea9401ee..e5fc9b6f6 100644 --- a/crates/sdk-core/src/worker/tuner/resource_based.rs +++ b/crates/sdk-core/src/worker/tuner/resource_based.rs @@ -778,7 +778,6 @@ mod tests { Arc, atomic::{AtomicU64, Ordering}, }, - thread::sleep, }; struct FakeMIS { @@ -1205,8 +1204,8 @@ mod tests { let to_allocate: usize = (half_total).saturating_sub(cur_used) as usize; let _buf = black_box(vec![1u8; to_allocate]); - // make sure we sleep enough to let real_sys_info need a refresh - sleep(Duration::from_millis(200)); + // Refresh synchronously so the assertion cannot race the background sampler. + sys_info.inner.refresh(); let percentage = sys_info.used_mem_percent(); let diff = (percentage - expected_percentage).abs(); diff --git a/crates/sdk-core/src/worker/workflow/driven_workflow.rs b/crates/sdk-core/src/worker/workflow/driven_workflow.rs index 654f49d11..c3e57a4d5 100644 --- a/crates/sdk-core/src/worker/workflow/driven_workflow.rs +++ b/crates/sdk-core/src/worker/workflow/driven_workflow.rs @@ -102,12 +102,8 @@ impl DrivenWorkflow { /// from a buffer that the language side sinks into when it calls [crate::Core::complete_task] pub(super) fn fetch_workflow_iteration_output(&mut self) -> Vec { let in_cmds = self.incoming_commands.try_recv(); - let in_cmds = in_cmds.unwrap_or_else(|_| { - vec![WFCommand { - variant: WFCommandVariant::NoCommandsFromLang, - metadata: None, - }] - }); + let in_cmds = + in_cmds.unwrap_or_else(|_| vec![WFCommand::new(WFCommandVariant::NoCommandsFromLang)]); debug!(in_cmds = %in_cmds.display(), "wf bridge iteration fetch"); in_cmds } diff --git a/crates/sdk-core/src/worker/workflow/machines/activity_state_machine.rs b/crates/sdk-core/src/worker/workflow/machines/activity_state_machine.rs index 15f1842cb..0d9a2dc5f 100644 --- a/crates/sdk-core/src/worker/workflow/machines/activity_state_machine.rs +++ b/crates/sdk-core/src/worker/workflow/machines/activity_state_machine.rs @@ -7,7 +7,10 @@ use super::{ use crate::{ abstractions::dbg_panic, internal_flags::CoreInternalFlags, - worker::workflow::{InternalFlagsRef, fatal, machines::HistEventData, nondeterminism}, + worker::workflow::{ + CommandAnnotations, InternalFlagsRef, ProtoCommandExt, fatal, machines::HistEventData, + nondeterminism, + }, }; use std::convert::{TryFrom, TryInto}; use temporalio_common::protos::{ @@ -112,6 +115,7 @@ impl ActivityMachine { attrs: ScheduleActivity, internal_flags: InternalFlagsRef, use_compatible_version: bool, + annotations: CommandAnnotations, ) -> NewMachineWithCommand { let mut s = Self::from_parts( Created {}.into(), @@ -123,6 +127,7 @@ impl ActivityMachine { scheduled_event_id: 0, started_event_id: 0, cancelled_before_sent: false, + annotations, }, ); OnEventWrapper::on_event_mut(&mut s, ActivityMachineEvents::Schedule) @@ -164,7 +169,11 @@ impl ActivityMachine { } } - pub(super) fn cancel(&mut self) -> Result, MachineError> { + pub(super) fn cancel( + &mut self, + annotations: CommandAnnotations, + ) -> Result, MachineError> { + self.shared_state.annotations.override_with(annotations); if matches!( self.state(), ActivityMachineState::Completed(_) @@ -315,6 +324,7 @@ impl WFMachinesAdapter for ActivityMachine { result: Some(ActivityResolution { status: Some(activity_resolution::Status::Failed(ar::Failure { failure: Some(failure), + ..Default::default() })), }), is_local: false, @@ -352,6 +362,7 @@ pub(super) struct SharedState { cancellation_type: ActivityCancellationType, cancelled_before_sent: bool, internal_flags: InternalFlagsRef, + annotations: CommandAnnotations, } #[derive(Default, Clone)] @@ -790,18 +801,14 @@ fn create_request_cancel_activity_task_command( where S: Into, { - let cmd = Command { - command_type: CommandType::RequestCancelActivityTask as i32, - attributes: Some( - command::Attributes::RequestCancelActivityTaskCommandAttributes( - RequestCancelActivityTaskCommandAttributes { - scheduled_event_id: dat.scheduled_event_id, - }, - ), + let cmd = Command::new( + command::Attributes::RequestCancelActivityTaskCommandAttributes( + RequestCancelActivityTaskCommandAttributes { + scheduled_event_id: dat.scheduled_event_id, + }, ), - user_metadata: Default::default(), - event_group_markers: vec![], - }; + dat.annotations.clone(), + ); ActivityMachineTransition::ok( vec![ActivityMachineCommand::RequestCancellation(cmd)], next_state, @@ -937,9 +944,10 @@ mod test { cancellation_type: Default::default(), cancelled_before_sent: false, internal_flags: Rc::new(RefCell::new(InternalFlags::default())), + annotations: Default::default(), }, ); - let cmds = s.cancel().unwrap(); + let cmds = s.cancel(Default::default()).unwrap(); assert_eq!(cmds.len(), 0); assert_eq!(discriminant(&state), discriminant(s.state())); } @@ -956,13 +964,14 @@ mod test { }, Rc::new(RefCell::new(InternalFlags::default())), true, + Default::default(), ); let mut s = if let Machines::ActivityMachine(am) = s.machine { am } else { panic!("Wrong machine type"); }; - let cmds = s.cancel().unwrap(); + let cmds = s.cancel(Default::default()).unwrap(); // We should always be notifying lang that the activity got cancelled, even if it's // abandoned and we aren't telling server assert_matches!( diff --git a/crates/sdk-core/src/worker/workflow/machines/cancel_external_state_machine.rs b/crates/sdk-core/src/worker/workflow/machines/cancel_external_state_machine.rs index 392135176..f1986cced 100644 --- a/crates/sdk-core/src/worker/workflow/machines/cancel_external_state_machine.rs +++ b/crates/sdk-core/src/worker/workflow/machines/cancel_external_state_machine.rs @@ -180,6 +180,7 @@ impl WFMachinesAdapter for CancelExternalMachine { ResolveRequestCancelExternalWorkflow { seq: self.shared_state.seq, failure: None, + cause: CancelExternalWorkflowExecutionFailedCause::Unspecified as i32, } .into(), ] @@ -203,6 +204,7 @@ impl WFMachinesAdapter for CancelExternalMachine { )), ..Default::default() }), + cause: f as i32, } .into(), ] diff --git a/crates/sdk-core/src/worker/workflow/machines/child_workflow_state_machine.rs b/crates/sdk-core/src/worker/workflow/machines/child_workflow_state_machine.rs index 80171659f..609619a81 100644 --- a/crates/sdk-core/src/worker/workflow/machines/child_workflow_state_machine.rs +++ b/crates/sdk-core/src/worker/workflow/machines/child_workflow_state_machine.rs @@ -5,7 +5,9 @@ use super::{ use crate::{ abstractions::dbg_panic, internal_flags::CoreInternalFlags, - worker::workflow::{InternalFlagsRef, fatal, machines::HistEventData, nondeterminism}, + worker::workflow::{ + CommandAnnotations, InternalFlagsRef, fatal, machines::HistEventData, nondeterminism, + }, }; use itertools::Itertools; use std::{ @@ -445,6 +447,7 @@ pub(super) struct SharedState { cancelled_before_sent: bool, cancel_type: ChildWorkflowCancellationType, internal_flags: InternalFlagsRef, + annotations: CommandAnnotations, } impl SharedState { @@ -462,6 +465,7 @@ impl ChildWorkflowMachine { attribs: StartChildWorkflowExecution, internal_flags: InternalFlagsRef, use_compatible_version: bool, + annotations: CommandAnnotations, ) -> NewMachineWithCommand { let mut s = Self::from_parts( Created {}.into(), @@ -476,6 +480,7 @@ impl ChildWorkflowMachine { initiated_event_id: 0, started_event_id: 0, cancelled_before_sent: false, + annotations, }, ); OnEventWrapper::on_event_mut(&mut s, ChildWorkflowMachineEvents::Schedule) @@ -516,7 +521,9 @@ impl ChildWorkflowMachine { pub(super) fn cancel( &mut self, reason: String, + annotations: CommandAnnotations, ) -> Result, MachineError> { + self.shared_state.annotations.override_with(annotations); let event = ChildWorkflowMachineEvents::Cancel(reason); let vec = OnEventWrapper::on_event_mut(self, event)?; let res = vec @@ -733,8 +740,8 @@ impl WFMachinesAdapter for ChildWorkflowMachine { let mut resps = vec![]; if self.shared_state.cancel_type != ChildWorkflowCancellationType::Abandon { #[allow(deprecated)] - resps.push(MachineResponse::NewCoreOriginatedCommand( - RequestCancelExternalWorkflowExecutionCommandAttributes { + resps.push(MachineResponse::NewCoreOriginatedCommand { + attrs: RequestCancelExternalWorkflowExecutionCommandAttributes { namespace: self.shared_state.namespace.clone(), workflow_id: self.shared_state.workflow_id.clone(), run_id: self.shared_state.run_id.clone(), @@ -743,7 +750,8 @@ impl WFMachinesAdapter for ChildWorkflowMachine { ..Default::default() } .into(), - )) + annotations: self.shared_state.annotations.clone(), + }) } if self.shared_state.resolves_immediately_on_cancel() { resps.push(self.resolve_cancelled_msg().into()) @@ -826,9 +834,12 @@ mod test { cancelled_before_sent: false, cancel_type: Default::default(), internal_flags: Rc::new(RefCell::new(InternalFlags::default())), + annotations: Default::default(), }, ); - let cmds = s.cancel("cancel reason".to_string()).unwrap(); + let cmds = s + .cancel("cancel reason".to_string(), Default::default()) + .unwrap(); assert_eq!(cmds.len(), 0); assert_eq!(discriminant(&state), discriminant(s.state())); } @@ -851,6 +862,7 @@ mod test { cancelled_before_sent: false, cancel_type, internal_flags: Rc::new(RefCell::new(InternalFlags::default())), + annotations: Default::default(), }; let state = Cancelled::default(); let res = state.on_child_workflow_execution_completed(&mut shared, None); @@ -925,10 +937,11 @@ mod test { cancelled_before_sent: false, cancel_type, internal_flags: Rc::new(RefCell::new(InternalFlags::default())), + annotations: Default::default(), }, ); let cmds = s - .cancel("parent cancelled".to_string()) + .cancel("parent cancelled".to_string(), Default::default()) .expect("Cancel in StartEventRecorded should not fail"); assert!( !cmds.is_empty(), diff --git a/crates/sdk-core/src/worker/workflow/machines/local_activity_state_machine.rs b/crates/sdk-core/src/worker/workflow/machines/local_activity_state_machine.rs index 33dfeea63..f956b75f9 100644 --- a/crates/sdk-core/src/worker/workflow/machines/local_activity_state_machine.rs +++ b/crates/sdk-core/src/worker/workflow/machines/local_activity_state_machine.rs @@ -32,6 +32,7 @@ use temporalio_common::protos::{ }, temporal::api::{ command::v1::{Command as ProtoCommand, RecordMarkerCommandAttributes, command}, + common::v1::Payloads, enums::v1::{CommandType, EventType, RetryState}, failure::v1::{Failure, failure::FailureInfo}, }, @@ -134,6 +135,7 @@ impl From for ResolveDat { } else { LocalActivityExecutionResult::Failed(ActFail { failure: Some(fail), + ..Default::default() }) } } @@ -616,7 +618,7 @@ impl WFMachinesAdapter for LocalActivityMachine { maybe_failure = fail.failure; } LocalActivityExecutionResult::Cancelled(Cancellation { failure }) - | LocalActivityExecutionResult::TimedOut(ActFail { failure }) => { + | LocalActivityExecutionResult::TimedOut(ActFail { failure, .. }) => { will_not_run_again = true; maybe_failure = failure; } @@ -712,25 +714,35 @@ impl WFMachinesAdapter for LocalActivityMachine { } if record_marker { + let mut details = build_local_activity_marker_details( + LocalActivityMarkerData { + seq: self.shared_state.attrs.seq, + attempt, + activity_id: self.shared_state.attrs.activity_id.clone(), + activity_type: self.shared_state.attrs.activity_type.clone(), + complete_time: complete_time.map(Into::into), + backoff, + original_schedule_time: original_schedule_time.map(Into::into), + }, + maybe_ok_result, + ); + if self.shared_state.attrs.include_arguments_in_marker { + details.insert( + "input".to_string(), + Payloads { + payloads: self.shared_state.attrs.arguments.clone(), + }, + ); + } let marker_data = RecordMarkerCommandAttributes { marker_name: LOCAL_ACTIVITY_MARKER_NAME.to_string(), - details: build_local_activity_marker_details( - LocalActivityMarkerData { - seq: self.shared_state.attrs.seq, - attempt, - activity_id: self.shared_state.attrs.activity_id.clone(), - activity_type: self.shared_state.attrs.activity_type.clone(), - complete_time: complete_time.map(Into::into), - backoff, - original_schedule_time: original_schedule_time.map(Into::into), - }, - maybe_ok_result, - ), + details, header: None, failure: maybe_failure, }; let command = ProtoCommand { user_metadata: self.shared_state.attrs.user_metadata.clone(), + event_group_markers: self.shared_state.attrs.event_group_markers.clone(), ..command::Attributes::RecordMarkerCommandAttributes(marker_data).into() }; responses.push(MachineResponse::IssueNewCommand(command)); diff --git a/crates/sdk-core/src/worker/workflow/machines/nexus_operation_state_machine.rs b/crates/sdk-core/src/worker/workflow/machines/nexus_operation_state_machine.rs index b8c744982..b8ad98ecc 100644 --- a/crates/sdk-core/src/worker/workflow/machines/nexus_operation_state_machine.rs +++ b/crates/sdk-core/src/worker/workflow/machines/nexus_operation_state_machine.rs @@ -1,6 +1,6 @@ use super::{MachineError, StateMachine, TransitionResult, fsm}; use crate::worker::workflow::{ - WFMachinesError, + CommandAnnotations, ProtoCommandExt, WFMachinesError, machines::{ EventInfo, HistEventData, NewMachineWithCommand, OnEventWrapper, WFMachinesAdapter, workflow_machines::MachineResponse, @@ -17,7 +17,7 @@ use temporalio_common::protos::{ workflow_commands::ScheduleNexusOperation, }, temporal::api::{ - command::v1::{RequestCancelNexusOperationCommandAttributes, command}, + command::v1::{Command, RequestCancelNexusOperationCommandAttributes, command}, common::v1::Payload, enums::v1::{CommandType, EventType}, failure::v1::{self as failure, Failure, failure::FailureInfo}, @@ -129,10 +129,14 @@ pub(super) struct SharedState { cancel_sent: bool, cancel_type: NexusOperationCancellationType, operation_token: Option, + annotations: CommandAnnotations, } impl NexusOperationMachine { - pub(super) fn new_scheduled(attribs: ScheduleNexusOperation) -> NewMachineWithCommand { + pub(super) fn new_scheduled( + attribs: ScheduleNexusOperation, + annotations: CommandAnnotations, + ) -> NewMachineWithCommand { let s = Self::from_parts( ScheduleCommandCreated.into(), SharedState { @@ -145,6 +149,7 @@ impl NexusOperationMachine { cancel_sent: false, cancel_type: attribs.cancellation_type(), operation_token: None, + annotations, }, ); NewMachineWithCommand { @@ -153,7 +158,11 @@ impl NexusOperationMachine { } } - pub(super) fn cancel(&mut self) -> Result, MachineError> { + pub(super) fn cancel( + &mut self, + annotations: CommandAnnotations, + ) -> Result, MachineError> { + self.shared_state.annotations.override_with(annotations); let event = NexusOperationMachineEvents::Cancel; let cmds = OnEventWrapper::on_event_mut(self, event)?; let mach_resps = cmds @@ -640,14 +649,14 @@ impl WFMachinesAdapter for NexusOperationMachine { NexusOperationCommand::IssueCancel => { let mut resps = vec![]; if self.shared_state.cancel_type != NexusOperationCancellationType::Abandon { - resps.push(MachineResponse::IssueNewCommand( + resps.push(MachineResponse::IssueNewCommand(Command::new( command::Attributes::RequestCancelNexusOperationCommandAttributes( RequestCancelNexusOperationCommandAttributes { scheduled_event_id: self.shared_state.scheduled_event_id, }, - ) - .into(), - )) + ), + self.shared_state.annotations.clone(), + ))) } // Immediately resolve abandon/trycancel modes if matches!( diff --git a/crates/sdk-core/src/worker/workflow/machines/patch_state_machine.rs b/crates/sdk-core/src/worker/workflow/machines/patch_state_machine.rs index 1e99e6c50..469f53ec1 100644 --- a/crates/sdk-core/src/worker/workflow/machines/patch_state_machine.rs +++ b/crates/sdk-core/src/worker/workflow/machines/patch_state_machine.rs @@ -25,7 +25,7 @@ use crate::{ internal_flags::CoreInternalFlags, protosext::HistoryEventExt, worker::workflow::{ - InternalFlagsRef, fatal, + CommandAnnotations, InternalFlagsRef, fatal, machines::{ HistEventData, upsert_search_attributes_state_machine::MAX_SEARCH_ATTR_PAYLOAD_SIZE, }, @@ -87,6 +87,8 @@ pub(super) enum PatchCommand {} /// are guaranteed to return the same value. /// `replaying_when_invoked`: If the workflow is replaying when this invocation occurs, this needs /// to be set to true. +/// `annotations`: Lang's annotations on the patch command. They will be attached to both the +/// RecordMarker command and the search attribute upsert we synthesize alongside the marker. pub(super) fn has_change<'a>( patch_id: String, replaying_when_invoked: bool, @@ -94,6 +96,7 @@ pub(super) fn has_change<'a>( seen_in_peekahead: bool, existing_patch_ids: impl Iterator, internal_flags: InternalFlagsRef, + annotations: CommandAnnotations, ) -> Result<(NewMachineWithCommand, Vec), WFMachinesError> { let shared_state = SharedState { patch_id }; let initial_state = if replaying_when_invoked { @@ -150,12 +153,13 @@ pub(super) fn has_change<'a>( m.insert(VERSION_SEARCH_ATTR_KEY.to_string(), serialized); m }; - vec![MachineResponse::NewCoreOriginatedCommand( - UpsertWorkflowSearchAttributesCommandAttributes { + vec![MachineResponse::NewCoreOriginatedCommand { + attrs: UpsertWorkflowSearchAttributesCommandAttributes { search_attributes: Some(SearchAttributes { indexed_fields }), } .into(), - )] + annotations, + }] } }; diff --git a/crates/sdk-core/src/worker/workflow/machines/signal_external_state_machine.rs b/crates/sdk-core/src/worker/workflow/machines/signal_external_state_machine.rs index 8b8710cd2..1a93ca2a6 100644 --- a/crates/sdk-core/src/worker/workflow/machines/signal_external_state_machine.rs +++ b/crates/sdk-core/src/worker/workflow/machines/signal_external_state_machine.rs @@ -220,6 +220,7 @@ impl WFMachinesAdapter for SignalExternalMachine { ResolveSignalExternalWorkflow { seq: self.shared_state.seq, failure: None, + cause: SignalExternalWorkflowExecutionFailedCause::Unspecified as i32, } .into(), ] @@ -248,6 +249,7 @@ impl WFMachinesAdapter for SignalExternalMachine { )), ..Default::default() }), + cause: f as i32, } .into(), ] @@ -278,6 +280,7 @@ impl SignalExternalMachine { )), ..Default::default() }), + cause: SignalExternalWorkflowExecutionFailedCause::Unspecified as i32, } .into(), ]; diff --git a/crates/sdk-core/src/worker/workflow/machines/timer_state_machine.rs b/crates/sdk-core/src/worker/workflow/machines/timer_state_machine.rs index f193fd41d..066de1d8f 100644 --- a/crates/sdk-core/src/worker/workflow/machines/timer_state_machine.rs +++ b/crates/sdk-core/src/worker/workflow/machines/timer_state_machine.rs @@ -4,7 +4,10 @@ use super::{ EventInfo, MachineError, NewMachineWithCommand, OnEventWrapper, StateMachine, TransitionResult, WFMachinesAdapter, fsm, workflow_machines::MachineResponse, }; -use crate::worker::workflow::{WFMachinesError, fatal, machines::HistEventData, nondeterminism}; +use crate::worker::workflow::{ + CommandAnnotations, ProtoCommandExt, WFMachinesError, fatal, machines::HistEventData, + nondeterminism, +}; use std::convert::TryFrom; use temporalio_common::protos::{ coresdk::{ @@ -13,7 +16,7 @@ use temporalio_common::protos::{ workflow_commands::{CancelTimer, StartTimer}, }, temporal::api::{ - command::v1::command, + command::v1::{Command, command}, enums::v1::{CommandType, EventType}, history::v1::{TimerFiredEventAttributes, history_event}, }, @@ -58,11 +61,15 @@ pub(super) enum TimerMachineCommand { pub(super) struct SharedState { attrs: StartTimer, cancelled_before_sent: bool, + annotations: CommandAnnotations, } /// Creates a new, scheduled, timer as a [CancellableCommand] -pub(super) fn new_timer(attribs: StartTimer) -> NewMachineWithCommand { - let (timer, add_cmd) = TimerMachine::new_scheduled(attribs); +pub(super) fn new_timer( + attribs: StartTimer, + annotations: CommandAnnotations, +) -> NewMachineWithCommand { + let (timer, add_cmd) = TimerMachine::new_scheduled(attribs, annotations); NewMachineWithCommand { command: add_cmd, machine: timer.into(), @@ -71,29 +78,40 @@ pub(super) fn new_timer(attribs: StartTimer) -> NewMachineWithCommand { impl TimerMachine { /// Create a new timer and immediately schedule it - fn new_scheduled(attribs: StartTimer) -> (Self, command::Attributes) { - let mut s = Self::new(attribs); + fn new_scheduled( + attribs: StartTimer, + annotations: CommandAnnotations, + ) -> (Self, command::Attributes) { + let mut s = Self::new(attribs, annotations); OnEventWrapper::on_event_mut(&mut s, TimerMachineEvents::Schedule) .expect("Scheduling timers doesn't fail"); let cmd = s.shared_state().attrs.into(); (s, cmd) } - fn new(attribs: StartTimer) -> Self { + fn new(attribs: StartTimer, annotations: CommandAnnotations) -> Self { Self::from_parts( Created {}.into(), SharedState { attrs: attribs, cancelled_before_sent: false, + annotations, }, ) } - pub(super) fn cancel(&mut self) -> Result, MachineError> { + pub(super) fn cancel( + &mut self, + annotations: CommandAnnotations, + ) -> Result, MachineError> { + self.shared_state.annotations.override_with(annotations); Ok( match OnEventWrapper::on_event_mut(self, TimerMachineEvents::Cancel)?.pop() { Some(TimerMachineCommand::IssueCancelCmd(cmd)) => { - vec![MachineResponse::IssueNewCommand(cmd.into())] + vec![MachineResponse::IssueNewCommand(Command::new( + cmd, + self.shared_state.annotations.clone(), + ))] } None => vec![], x => panic!("Invalid cancel event response {x:?}"), @@ -257,7 +275,10 @@ impl WFMachinesAdapter for TimerMachine { .into(), ], TimerMachineCommand::IssueCancelCmd(c) => { - vec![MachineResponse::IssueNewCommand(c.into())] + vec![MachineResponse::IssueNewCommand(Command::new( + c, + self.shared_state.annotations.clone(), + ))] } }) } @@ -272,7 +293,7 @@ mod test { fn cancels_ignored_terminal() { for state in [TimerMachineState::Canceled(Canceled {}), Fired {}.into()] { let mut s = TimerMachine::from_parts(state.clone(), Default::default()); - let cmds = s.cancel().unwrap(); + let cmds = s.cancel(Default::default()).unwrap(); assert_eq!(cmds.len(), 0); assert_eq!(discriminant(&state), discriminant(s.state())); } diff --git a/crates/sdk-core/src/worker/workflow/machines/transition_coverage.rs b/crates/sdk-core/src/worker/workflow/machines/transition_coverage.rs index daee6afbf..db12b886f 100644 --- a/crates/sdk-core/src/worker/workflow/machines/transition_coverage.rs +++ b/crates/sdk-core/src/worker/workflow/machines/transition_coverage.rs @@ -75,7 +75,8 @@ mod machine_coverage_report { fail_workflow_state_machine::FailWorkflowMachine, local_activity_state_machine::LocalActivityMachine, modify_workflow_properties_state_machine::ModifyWorkflowPropertiesMachine, - patch_state_machine::PatchMachine, signal_external_state_machine::SignalExternalMachine, + nexus_operation_state_machine::NexusOperationMachine, patch_state_machine::PatchMachine, + signal_external_state_machine::SignalExternalMachine, subscribe_notification_channel_state_machine::SubscribeNotificationChannelMachine, timer_state_machine::TimerMachine, unsubscribe_notification_channel_state_machine::UnsubscribeNotificationChannelMachine, @@ -119,6 +120,7 @@ mod machine_coverage_report { let mut upsert_search_attr = UpsertSearchAttributesMachine::visualizer().to_owned(); let mut modify_wf_props = ModifyWorkflowPropertiesMachine::visualizer().to_owned(); let mut update = UpdateMachine::visualizer().to_owned(); + let mut nexus = NexusOperationMachine::visualizer().to_owned(); let mut external_stream = ExternalStreamMachine::visualizer().to_owned(); let mut subscribe_channel = SubscribeNotificationChannelMachine::visualizer().to_owned(); let mut unsubscribe_channel = @@ -149,6 +151,7 @@ mod machine_coverage_report { cover_transitions(m, &mut modify_wf_props, coverage) } m @ "UpdateMachine" => cover_transitions(m, &mut update, coverage), + m @ "NexusOperationMachine" => cover_transitions(m, &mut nexus, coverage), m @ "ExternalStreamMachine" => cover_transitions(m, &mut external_stream, coverage), m @ "SubscribeNotificationChannelMachine" => { cover_transitions(m, &mut subscribe_channel, coverage) diff --git a/crates/sdk-core/src/worker/workflow/machines/workflow_machines.rs b/crates/sdk-core/src/worker/workflow/machines/workflow_machines.rs index a18f0c2a0..44e94ad58 100644 --- a/crates/sdk-core/src/worker/workflow/machines/workflow_machines.rs +++ b/crates/sdk-core/src/worker/workflow/machines/workflow_machines.rs @@ -27,9 +27,9 @@ use crate::{ worker::{ ExecutingLAId, LocalActRequest, LocalActivityExecutionResult, LocalActivityResolution, workflow::{ - CommandID, DrivenWorkflow, HistoryUpdate, InternalFlagsRef, LocalResolution, - OutgoingJob, RunBasics, WFCommand, WFCommandVariant, WFMachinesError, - WorkflowStartedInfo, fatal, + CommandAnnotations, CommandID, DrivenWorkflow, HistoryUpdate, InternalFlagsRef, + LocalResolution, OutgoingJob, ProtoCommandExt, RunBasics, WFCommand, WFCommandVariant, + WFMachinesError, WorkflowStartedInfo, fatal, history_update::NextWFT, machines::{ HistEventData, @@ -75,15 +75,13 @@ use temporalio_common::{ }, }, temporal::api::{ - command::v1::{ - Command as ProtoCommand, CommandAttributesExt, command::Attributes as ProtoCmdAttrs, - }, + command::v1::{Command as ProtoCommand, command::Attributes as ProtoCmdAttrs}, common::v1::SearchAttributes, enums::v1::EventType, history::v1::{HistoryEvent, history_event}, notification::v1::Notification, protocol::v1::{Message as ProtocolMessage, message::SequencingId}, - sdk::v1::{UserMetadata, WorkflowTaskCompletedMetadata}, + sdk::v1::WorkflowTaskCompletedMetadata, }, }, worker::WorkerDeploymentVersion, @@ -235,9 +233,13 @@ pub(super) enum MachineResponse { IssueNewMessage(ProtocolMessage), /// The machine requests the creation of another *different* machine. This acts as if lang /// had replied to the activation with a command, but we use a special set of IDs to avoid - /// collisions. - #[display("NewCoreOriginatedCommand({_0:?})")] - NewCoreOriginatedCommand(ProtoCmdAttrs), + /// collisions. The requesting machine supplies the annotations, since it is the only thing + /// that knows which lang command this one is being issued on behalf of. + #[display("NewCoreOriginatedCommand({attrs:?})")] + NewCoreOriginatedCommand { + attrs: ProtoCmdAttrs, + annotations: CommandAnnotations, + }, #[display("TriggerWFTaskStarted")] TriggerWFTaskStarted { task_started_event_id: i64, @@ -546,7 +548,7 @@ impl WorkflowMachines { } let key = self.add_cmd_to_wf_task( ExternalStreamMachine::record_marker(data), - None, + Default::default(), CommandIdKind::CoreInternal, ); self.external_stream_marker_machines.push_back(key); @@ -565,7 +567,7 @@ impl WorkflowMachines { pub(crate) fn emit_notification_channel_subscription(&mut self, channel: String) -> Result<()> { self.add_cmd_to_wf_task( subscribe_notification_channel(SubscribeNotificationChannel { channel }), - None, + Default::default(), CommandIdKind::CoreInternal, ); self.prepare_commands() @@ -579,7 +581,7 @@ impl WorkflowMachines { ) -> Result<()> { self.add_cmd_to_wf_task( unsubscribe_notification_channel(UnsubscribeNotificationChannel { channel }), - None, + Default::default(), CommandIdKind::CoreInternal, ); self.prepare_commands() @@ -808,10 +810,10 @@ impl WorkflowMachines { (*$me.observed_internal_flags) .borrow_mut() .add_from_complete($wtc); - let mut combined_ver = WorkerDeploymentVersion { - deployment_name: "".to_string(), - build_id: "".to_string(), - }; + let mut combined_ver = WorkerDeploymentVersion::builder() + .deployment_name("") + .build_id("") + .build(); #[allow(deprecated)] if let Some(bid) = $wtc.worker_version.as_ref().map(|wv| &wv.build_id) { combined_ver.build_id = bid.to_string(); @@ -1311,8 +1313,9 @@ impl WorkflowMachines { debug!("Ignoring an unreadable external stream wake Signal"); } } else { - self.drive_me - .send_job(workflow_activation::SignalWorkflow::from(attrs).into()); + self.drive_me.send_job( + workflow_activation::SignalWorkflow::from((attrs, event_id)).into(), + ); } } else { // err @@ -1469,7 +1472,7 @@ impl WorkflowMachines { self.message_outbox.push_back(pm); } } - MachineResponse::NewCoreOriginatedCommand(attrs) => match attrs { + MachineResponse::NewCoreOriginatedCommand { attrs, annotations } => match attrs { ProtoCmdAttrs::RequestCancelExternalWorkflowExecutionCommandAttributes( attrs, ) => { @@ -1481,7 +1484,7 @@ impl WorkflowMachines { }; self.add_cmd_to_wf_task( new_external_cancel(0, we, attrs.child_workflow_only, attrs.reason), - None, + annotations, CommandIdKind::CoreInternal, ); } @@ -1491,7 +1494,7 @@ impl WorkflowMachines { // workflows by users (but rather, just for them to search with). self.add_cmd_to_wf_task( upsert_search_attrs_internal(attrs), - None, + annotations, CommandIdKind::NeverResolves, ); } @@ -1592,12 +1595,16 @@ impl WorkflowMachines { /// server. fn handle_driven_results(&mut self, results: Vec) -> Result<()> { for cmd in results { - match cmd.variant { + let WFCommand { + variant, + annotations, + } = cmd; + match variant { WFCommandVariant::AddTimer(attrs) => { let seq = attrs.seq; self.add_cmd_to_wf_task( - new_timer(attrs), - cmd.metadata, + new_timer(attrs, annotations.clone()), + annotations, CommandID::Timer(seq).into(), ); } @@ -1610,12 +1617,18 @@ impl WorkflowMachines { self.observed_internal_flags.clone(), self.replaying, ), - cmd.metadata, + annotations, CommandIdKind::NeverResolves, ); } WFCommandVariant::CancelTimer(attrs) => { - cancel_machine!(self, CommandID::Timer(attrs.seq), TimerMachine, cancel); + cancel_machine!( + self, + CommandID::Timer(attrs.seq), + TimerMachine, + cancel, + annotations + ); } WFCommandVariant::AddActivity(attrs) => { let seq = attrs.seq; @@ -1628,17 +1641,22 @@ impl WorkflowMachines { attrs, self.observed_internal_flags.clone(), use_compat, + annotations.clone(), ), - cmd.metadata, + annotations, CommandID::Activity(seq).into(), ); } WFCommandVariant::AddLocalActivity(attrs) => { let seq = attrs.seq; - let attrs: ValidScheduleLA = - ValidScheduleLA::from_schedule_la(attrs, cmd.metadata).map_err(|e| { - fatal!("Invalid schedule local activity request (seq {seq}): {e}") - })?; + let attrs: ValidScheduleLA = ValidScheduleLA::from_schedule_la( + attrs, + annotations.metadata, + annotations.event_group_markers, + ) + .map_err(|e| { + fatal!("Invalid schedule local activity request (seq {seq}): {e}") + })?; let (la, mach_resp) = new_local_activity( attrs, self.replaying, @@ -1655,7 +1673,8 @@ impl WorkflowMachines { self, CommandID::Activity(attrs.seq), ActivityMachine, - cancel + cancel, + annotations ); } WFCommandVariant::RequestCancelLocalActivity(attrs) => { @@ -1667,10 +1686,10 @@ impl WorkflowMachines { ); } WFCommandVariant::CompleteWorkflow(attrs) => { - self.add_terminal_command(complete_workflow(attrs), cmd.metadata); + self.add_terminal_command(complete_workflow(attrs), annotations); } WFCommandVariant::FailWorkflow(attrs) => { - self.add_terminal_command(fail_workflow(attrs), cmd.metadata); + self.add_terminal_command(fail_workflow(attrs), annotations); } WFCommandVariant::ContinueAsNew(attrs) => { let attrs = self.augment_continue_as_new_with_current_values(attrs); @@ -1678,10 +1697,10 @@ impl WorkflowMachines { attrs.versioning_intent(), &attrs.task_queue, ); - self.add_terminal_command(continue_as_new(attrs, use_compat), cmd.metadata); + self.add_terminal_command(continue_as_new(attrs, use_compat), annotations); } WFCommandVariant::CancelWorkflow(attrs) => { - self.add_terminal_command(cancel_workflow(attrs), cmd.metadata); + self.add_terminal_command(cancel_workflow(attrs), annotations); } WFCommandVariant::SetPatchMarker(attrs) => { // Do not create commands for change IDs that we have already created commands @@ -1699,10 +1718,11 @@ impl WorkflowMachines { .iter() .filter_map(|(k, ci)| ci.created_command.then_some(k.as_str())), self.observed_internal_flags.clone(), + annotations.clone(), )?; let mkey = self.add_cmd_to_wf_task( patch_machine, - cmd.metadata, + annotations, CommandIdKind::NeverResolves, ); self.process_machine_responses(mkey, other_cmds)?; @@ -1730,8 +1750,9 @@ impl WorkflowMachines { attrs, self.observed_internal_flags.clone(), use_compat, + annotations.clone(), ), - cmd.metadata, + annotations, CommandID::ChildWorkflowStart(seq).into(), ); } @@ -1741,7 +1762,8 @@ impl WorkflowMachines { CommandID::ChildWorkflowStart(attrs.child_workflow_seq), ChildWorkflowMachine, cancel, - attrs.reason + attrs.reason, + annotations ); } WFCommandVariant::RequestCancelExternalWorkflow(attrs) => { @@ -1758,7 +1780,7 @@ impl WorkflowMachines { self.run_id, attrs.reason ), ), - cmd.metadata, + annotations, CommandID::CancelExternal(attrs.seq).into(), ); } @@ -1766,7 +1788,7 @@ impl WorkflowMachines { let seq = attrs.seq; self.add_cmd_to_wf_task( new_external_signal(attrs, &self.worker_config.namespace)?, - cmd.metadata, + annotations, CommandID::SignalExternal(seq).into(), ); } @@ -1785,7 +1807,7 @@ impl WorkflowMachines { WFCommandVariant::ModifyWorkflowProperties(attrs) => { self.add_cmd_to_wf_task( modify_workflow_properties(attrs), - cmd.metadata, + annotations, CommandIdKind::NeverResolves, ); } @@ -1794,14 +1816,14 @@ impl WorkflowMachines { // nothing back. The notifications arrive later on a scheduled event. self.add_cmd_to_wf_task( subscribe_notification_channel(attrs), - cmd.metadata, + annotations, CommandIdKind::NeverResolves, ); } WFCommandVariant::UnsubscribeNotificationChannel(attrs) => { self.add_cmd_to_wf_task( unsubscribe_notification_channel(attrs), - cmd.metadata, + annotations, CommandIdKind::NeverResolves, ); } @@ -1822,8 +1844,8 @@ impl WorkflowMachines { WFCommandVariant::ScheduleNexusOperation(attrs) => { let seq = attrs.seq; self.add_cmd_to_wf_task( - NexusOperationMachine::new_scheduled(attrs), - cmd.metadata, + NexusOperationMachine::new_scheduled(attrs, annotations.clone()), + annotations, CommandID::NexusOperation(seq).into(), ); } @@ -1832,7 +1854,8 @@ impl WorkflowMachines { self, CommandID::NexusOperation(attrs.seq), NexusOperationMachine, - cancel + cancel, + annotations ); } // External stream commands are consumed above the machine level -- progress @@ -1840,17 +1863,18 @@ impl WorkflowMachines { // runtime-internal activations that `ManagedRun` resolves (C6, C8, C15a). None of // them should reach the machines; the output commit is also intercepted there. // Reaching any one silently would drop replay-visible state on the floor. - WFCommandVariant::ExternalStreamProgress(_) + leaked @ (WFCommandVariant::ExternalStreamProgress(_) | WFCommandVariant::ExternalStreamQuiescent(_) | WFCommandVariant::ExternalStreamParkResult(_) | WFCommandVariant::ExternalStreamFinalized(_) | WFCommandVariant::ExternalOutputStreamCommit(_) | WFCommandVariant::ExternalOutputStreamBuffered(_) - | WFCommandVariant::ExternalStreamChannels(_) => { + | WFCommandVariant::ExternalStreamChannels(_)) => { + // Named, because this is the only diagnostic for a wait-set bug and seven + // commands share the branch. return Err(fatal!( - "External stream command {} reached the state machines; it should have \ - been consumed by the run's external wait set", - cmd.variant + "External stream command {leaked} reached the state machines; it should \ + have been consumed by the run's external wait set" )); } WFCommandVariant::NoCommandsFromLang => (), @@ -1886,9 +1910,9 @@ impl WorkflowMachines { fn add_terminal_command( &mut self, machine: NewMachineWithCommand, - metadata: Option, + annotations: CommandAnnotations, ) { - let cwfm = self.add_new_command_machine(machine, metadata); + let cwfm = self.add_new_command_machine(machine, annotations); self.workflow_end_time = Some(SystemTime::now()); self.current_wf_task_commands.push_back(cwfm); // Wipe out any pending / executing local activity data since we're about to terminate @@ -1900,10 +1924,10 @@ impl WorkflowMachines { fn add_cmd_to_wf_task( &mut self, machine: NewMachineWithCommand, - metadata: Option, + annotations: CommandAnnotations, id: CommandIdKind, ) -> MachineKey { - let mach = self.add_new_command_machine(machine, metadata); + let mach = self.add_new_command_machine(machine, annotations); let key = mach.machine; if let CommandIdKind::LangIssued(id) = id { self.id_to_machine.insert(id, key); @@ -1918,17 +1942,11 @@ impl WorkflowMachines { fn add_new_command_machine( &mut self, machine: NewMachineWithCommand, - metadata: Option, + annotations: CommandAnnotations, ) -> CommandAndMachine { let k = self.all_machines.insert(machine.machine); - let cmd = ProtoCommand { - command_type: machine.command.as_type() as i32, - attributes: Some(machine.command), - user_metadata: metadata, - event_group_markers: vec![], - }; CommandAndMachine { - command: cmd, + command: ProtoCommand::new(machine.command, annotations), machine: k, } } diff --git a/crates/sdk-core/src/worker/workflow/managed_run.rs b/crates/sdk-core/src/worker/workflow/managed_run.rs index d8974cffb..649a14780 100644 --- a/crates/sdk-core/src/worker/workflow/managed_run.rs +++ b/crates/sdk-core/src/worker/workflow/managed_run.rs @@ -3,7 +3,6 @@ use crate::{ abstractions::dbg_panic, internal_flags::CoreInternalFlags, protosext::WorkflowActivationExt, - telemetry::metrics, worker::{ LEGACY_QUERY_ID, LocalActRequest, WorkflowErrorType, workflow::{ @@ -13,7 +12,8 @@ use crate::{ LocalActivityRequestSink, LocalResolution, NextPageReq, OutstandingActivation, OutstandingTask, PermittedWFT, RequestEvictMsg, RunBasics, RunTimerSink, ServerCommandsWithWorkflowInfo, TaskStorageMetrics, WFCommand, WFCommandVariant, - WFMachinesError, WFT_HEARTBEAT_TIMEOUT_FRACTION, WFTReportStatus, WorkflowTaskInfo, + WFMachinesError, WFT_HEARTBEAT_TIMEOUT_FRACTION, WFTReportStatus, WftFailureKind, + WorkflowTaskInfo, external_streams::{ ChannelSubscriptions, ExternalStreamReadyResult, ExternalStreamRunStatus, ExternalWaitSet, ExternalWaitState, ParkResolution, ParkStartOutcome, ParkTrigger, @@ -599,16 +599,19 @@ impl ManagedRun { self.run_id() ); } - let outcome = if let Some((tt, reason)) = self.trying_to_evict.as_mut().and_then(|te| { - te.auto_reply_fail_tt - .take() - .map(|tt| (tt, te.message.clone())) - }) { - ActivationCompleteOutcome::ReportWFTFail(FailedActivationWFTReport::Report( - tt, + let outcome = if let Some((info, reason)) = self + .trying_to_evict + .as_mut() + .and_then(|te| te.auto_reply_fail.take().map(|i| (i, te.message.clone()))) + { + ActivationCompleteOutcome::ReportWFTFail(Box::new(FailedActivationWFTReport::new( + info.task_token, + info.attempt, WorkflowTaskFailedCause::WorkflowWorkerUnhandledFailure, Failure::application_failure(reason, true).into(), - )) + WftFailureKind::Task, + &self.metrics, + ))) } else { ActivationCompleteOutcome::DoNothing }; @@ -766,8 +769,8 @@ impl ManagedRun { self.clear_external_output_buffered(); self.waiting_on_local_work.output_latency_pending = false; self.waiting_on_local_work.park_rollback_resolve_pending = false; - let tt = if let Some(tt) = self.wft.as_ref().map(|t| t.info.task_token.clone()) { - tt + let (tt, attempt) = if let Some(t) = self.wft.as_ref() { + (t.info.task_token.clone(), t.info.attempt) } else { dbg_panic!( "No workflow task for run id {} found when trying to fail activation", @@ -786,72 +789,68 @@ impl ManagedRun { EvictionReason::Unspecified | EvictionReason::PaginationOrHistoryFetch ); - let (should_report, rur) = if is_no_report_query_fail { - (false, None) + let rur = if is_no_report_query_fail { + None } else { // Blow up any cached data associated with the workflow - let evict_req_outcome = self.request_eviction(RequestEvictMsg { + self.request_eviction(RequestEvictMsg { run_id: self.run_id().to_string(), message, reason, - auto_reply_fail_tt: None, - }); - let should_report = match &evict_req_outcome { - EvictionRequestResult::EvictionRequested(Some(attempt), _) - | EvictionRequestResult::EvictionAlreadyRequested(Some(attempt)) => *attempt <= 1, - _ => false, - }; - let rur = evict_req_outcome.into_run_update_resp(); - (should_report, rur) + auto_reply_fail: None, + }) + .into_run_update_resp() }; - let outcome = if self.pending_work_is_legacy_query() { - if is_no_report_query_fail { - ActivationCompleteOutcome::WFTFailedDontReport - } else { - ActivationCompleteOutcome::ReportWFTFail( - FailedActivationWFTReport::ReportLegacyQueryFailure(tt, failure), - ) - } - } else if should_report { - // Check if we should fail the workflow instead of the WFT because of user's preferences - if matches!(cause, WorkflowTaskFailedCause::NonDeterministicError) - && self.config.should_fail_workflow( - &self.wfm.machines.workflow_type, - &WorkflowErrorType::Nondeterminism, - ) - { - warn!(failure=?failure, "Failing workflow due to nondeterminism error"); - return self - .successful_completion( - vec![WFCommand { - variant: WFCommandVariant::FailWorkflow(FailWorkflowExecution { - failure: failure.failure, - }), - metadata: None, - }], - vec![], - VersioningBehavior::Unspecified, // Doesn't matter since we're failing wf - resp_chan, - true, - ) - .unwrap_or_else(|e| { - dbg_panic!("Got next page request when auto-failing workflow: {e:?}"); - None - }); - } else { - ActivationCompleteOutcome::ReportWFTFail(FailedActivationWFTReport::Report( - tt, cause, failure, - )) - } + let kind = if !self.pending_work_is_legacy_query() { + WftFailureKind::Task + } else if is_no_report_query_fail { + WftFailureKind::RetryableLegacyQuery } else { - ActivationCompleteOutcome::WFTFailedDontReport + WftFailureKind::LegacyQuery }; - self.metrics - .with_new_attrs([metrics::failure_reason(cause.into())]) - .wf_task_failed(); - self.reply_to_complete(outcome, resp_chan); + // Check if we should fail the workflow instead of the WFT because of user's preferences. + // Only done on the first attempt: if that attempt's completion didn't reach the server, + // later attempts fall through to the normal task failure path, which won't re-report. + if kind == WftFailureKind::Task + && attempt <= 1 + && matches!(cause, WorkflowTaskFailedCause::NonDeterministicError) + && self.config.should_fail_workflow( + &self.wfm.machines.workflow_type, + &WorkflowErrorType::Nondeterminism, + ) + { + warn!(failure=?failure, "Failing workflow due to nondeterminism error"); + return self + .successful_completion( + vec![WFCommand::new(WFCommandVariant::FailWorkflow( + FailWorkflowExecution { + failure: failure.failure, + }, + ))], + vec![], + VersioningBehavior::Unspecified, // Doesn't matter since we're failing wf + resp_chan, + true, + ) + .unwrap_or_else(|e| { + dbg_panic!("Got next page request when auto-failing workflow: {e:?}"); + None + }); + } + + self.reply_to_complete( + ActivationCompleteOutcome::ReportWFTFail(Box::new(FailedActivationWFTReport::new( + tt, + attempt, + cause, + failure, + kind, + &self.metrics, + ))), + resp_chan, + ); rur } @@ -1347,6 +1346,10 @@ impl ManagedRun { // consumed data, and so must the marker recording it. Emitting after lang's commands were // pushed would put the marker *after* the terminal command in History, and on replay the // command would then be matched before the record it came from was validated. + // + // Completion pagination does not weaken that. It distributes the same command list across + // pages in order, the server only buffers the intermediate ones, and the final page is + // what merges them, so History still sees one commit with the marker where it was put. let terminal = if will_retain { None } else if let Some(ParkApplication::Confirmed(reason)) = park_outcome { @@ -2453,8 +2456,6 @@ impl ManagedRun { } pub(super) fn request_eviction(&mut self, info: RequestEvictMsg) -> EvictionRequestResult { - let attempts = self.wft.as_ref().map(|wt| wt.info.attempt); - // If we were waiting on a page fetch and we're getting evicted because fetching failed, // then make sure we allow the completion to proceed, otherwise we're stuck waiting forever. if self.completion_waiting_on_page_fetch.is_some() @@ -2469,7 +2470,7 @@ impl ManagedRun { true, c.resp_chan, ); - return EvictionRequestResult::EvictionRequested(attempts, run_upd); + return EvictionRequestResult::EvictionRequested(run_upd); } if !self.activation_is_eviction() && self.trying_to_evict.is_none() { @@ -2496,11 +2497,11 @@ impl ManagedRun { } self.trying_to_evict = Some(info); - EvictionRequestResult::EvictionRequested(attempts, self.check_more_activations()) + EvictionRequestResult::EvictionRequested(self.check_more_activations()) } else { // Always store the most recent eviction reason self.trying_to_evict = Some(info); - EvictionRequestResult::EvictionAlreadyRequested(attempts) + EvictionRequestResult::EvictionAlreadyRequested } } @@ -2802,7 +2803,7 @@ impl ManagedRun { run_id: self.run_id().to_string(), message: warnstr, reason: EvictionReason::Fatal, - auto_reply_fail_tt: None, + auto_reply_fail: None, }); } } @@ -3903,39 +3904,29 @@ mod tests { use super::*; pub(crate) fn complete() -> WFCommand { - WFCommand { - variant: WFCommandVariant::CompleteWorkflow(CompleteWorkflowExecution { - result: None, - }), - metadata: None, - } + WFCommand::new(WFCommandVariant::CompleteWorkflow( + CompleteWorkflowExecution { result: None }, + )) } pub(crate) fn cancel() -> WFCommand { - WFCommand { - variant: WFCommandVariant::CancelWorkflow(CancelWorkflowExecution {}), - metadata: None, - } + WFCommand::new(WFCommandVariant::CancelWorkflow( + CancelWorkflowExecution::default(), + )) } pub(crate) fn query_response() -> WFCommand { - WFCommand { - variant: WFCommandVariant::QueryResponse(QueryResult { - query_id: "".into(), - variant: None, - }), - metadata: None, - } + WFCommand::new(WFCommandVariant::QueryResponse(QueryResult { + query_id: "".into(), + variant: None, + })) } pub(crate) fn update_response() -> WFCommand { - WFCommand { - variant: WFCommandVariant::UpdateResponse(UpdateResponse { - protocol_instance_id: "".into(), - response: None, - }), - metadata: None, - } + WFCommand::new(WFCommandVariant::UpdateResponse(UpdateResponse { + protocol_instance_id: "".into(), + response: None, + })) } pub(crate) fn command_types(commands: &[WFCommand]) -> Vec> { diff --git a/crates/sdk-core/src/worker/workflow/mod.rs b/crates/sdk-core/src/worker/workflow/mod.rs index b8618bfd2..0043be83a 100644 --- a/crates/sdk-core/src/worker/workflow/mod.rs +++ b/crates/sdk-core/src/worker/workflow/mod.rs @@ -26,15 +26,16 @@ use crate::{ internal_flags::InternalFlags, pollers::TrackedPermittedTqResp, protosext::{ValidPollWFTQResponse, protocol_messages::IncomingProtocolMessage}, - telemetry::{ - VecDisplayer, - metrics::{self, FailureReason}, - }, + telemetry::{VecDisplayer, metrics}, worker::{ ActivitySlotKind, CompleteWfError, LocalActRequest, LocalActivityExecutionResult, - LocalActivityResolution, PollError, PostActivateHookData, WorkflowSlotKind, + LocalActivityResolution, NamespaceCapabilities, PollError, PostActivateHookData, + WorkflowSlotKind, activities::{ActivitiesFromWFTsHandle, LocalActivityManager}, - client::{LegacyQueryResult, WorkerClient, WorkflowTaskCompletion}, + client::{ + LegacyQueryResult, REQUEST_TOO_LARGE_KEY, WorkerClient, WorkflowTaskCompletion, + payload_limit_violation_from, + }, workflow::{ history_update::HistoryPaginator, machines::MachineError, @@ -66,7 +67,7 @@ use std::{ thread, time::{Duration, Instant}, }; -use temporalio_client::{MESSAGE_TOO_LARGE_KEY, payload_limit_violation_from}; +use temporalio_client::MESSAGE_TOO_LARGE_KEY; use temporalio_common::{ payload_limits::PayloadLimitViolation, protos::{ @@ -83,7 +84,9 @@ use temporalio_common::{ }, }, temporal::api::{ - command::v1::{Command as ProtoCommand, Command, command::Attributes}, + command::v1::{ + Command as ProtoCommand, Command, CommandAttributesExt, command::Attributes, + }, common::v1::{ Memo, MeteringMetadata, RetryPolicy, SearchAttributes, WorkflowExecution, }, @@ -91,7 +94,7 @@ use temporalio_common::{ failure::v1::{ApplicationFailureInfo, failure::FailureInfo}, protocol::v1::Message as ProtocolMessage, query::v1::WorkflowQuery, - sdk::v1::{UserMetadata, WorkflowTaskCompletedMetadata}, + sdk::v1::{EventGroupMarker, UserMetadata, WorkflowTaskCompletedMetadata}, taskqueue::v1::StickyExecutionAttributes, workflowservice::v1::{PollActivityTaskQueueResponse, get_system_info_response}, }, @@ -145,6 +148,8 @@ pub(crate) struct Workflows { local_act_mgr: Option>, ever_polled: AtomicBool, default_versioning_behavior: Option, + namespace_capabilities: Arc, + shutdown_token: CancellationToken, } pub(crate) struct WorkflowBasics { @@ -158,6 +163,7 @@ pub(crate) struct WorkflowBasics { pub(crate) worker_shutdown_token: CancellationToken, pub(crate) metrics: MetricsContext, pub(crate) server_capabilities: get_system_info_response::Capabilities, + pub(crate) namespace_capabilities: Arc, pub(crate) sdk_name: String, pub(crate) sdk_version: String, pub(crate) default_versioning_behavior: Option, @@ -200,6 +206,8 @@ impl Workflows { .worker_config .max_eager_activity_reservations_per_workflow_task; let default_versioning_behavior = basics.default_versioning_behavior; + let namespace_capabilities = basics.namespace_capabilities.clone(); + let shutdown_token = basics.shutdown_token.clone(); let extracted_wft_stream = WFTExtractor::build( client.clone(), basics.worker_config.fetching_concurrency, @@ -292,6 +300,8 @@ impl Workflows { local_act_mgr, ever_polled: AtomicBool::new(false), default_versioning_behavior, + namespace_capabilities, + shutdown_token, } } @@ -354,18 +364,9 @@ impl Workflows { } } }, - WorkflowStreamAction::FailUnstoredWft { - run_id, - task_token, - cause, - failure, - } => { - self.handle_activation_failed( - &run_id, - Instant::now(), - FailedActivationWFTReport::Report(task_token, cause, failure), - ) - .await; + WorkflowStreamAction::FailUnstoredWft { run_id, report } => { + self.handle_activation_failed(&run_id, Instant::now(), *report) + .await; } } } @@ -423,6 +424,12 @@ impl Workflows { nonfirst_local_activity_execution_attempts, }, versioning_behavior, + pagination_enabled: self + .namespace_capabilities + .workflow_task_completion_pagination(), + wft_completion_size_limit: self + .namespace_capabilities + .workflow_task_completion_size_limit(), }; let sticky_attrs = self.sticky_attrs.clone(); // Do not return new WFT if we would not cache, because returned new WFTs are @@ -434,7 +441,11 @@ impl Workflows { let mut reset_last_started_to = None; self.handle_wft_reporting_errs(run_id, || async { - match self.client.complete_workflow_task(completion).await { + match self + .client + .complete_workflow_task(completion, self.shutdown_token.clone()) + .await + { Ok(response) => { if let Some(record) = maybe_record_terminal_metric.take() { record(&run_metrics); @@ -451,37 +462,43 @@ impl Workflows { ); } Err(e) => { - let cause_reason_failure = if e - .metadata() - .contains_key(MESSAGE_TOO_LARGE_KEY) - && attempt < 2 - { - // gRPC message too large from server; skip on nonfirst attempts to - // avoid spamming. - Some(( - WorkflowTaskFailedCause::GrpcMessageTooLarge, - FailureReason::GrpcMessageTooLarge, - make_grpc_message_too_large_failure(), - )) - } else { - // Client layer rejected the completion for exceeding the worker's - // payload error limit. - payload_limit_violation_from(&e).map(|violation| { - ( - WorkflowTaskFailedCause::PayloadsTooLarge, - FailureReason::PayloadsTooLarge, - make_payloads_too_large_failure(violation), - ) - }) - }; - if let Some((cause, reason, failure)) = cause_reason_failure { - let new_outcome = - FailedActivationWFTReport::Report(task_token, cause, failure); - self.handle_activation_failed(run_id, completion_time, new_outcome) - .await; - run_metrics - .with_new_attrs([metrics::failure_reason(reason)]) - .wf_task_failed(); + let cause_and_failure = + if e.metadata().contains_key(REQUEST_TOO_LARGE_KEY) { + // Completion exceeds the namespace's recombined size limit, so the + // worker failed it proactively rather than sending doomed pages. + Some(( + WorkflowTaskFailedCause::RequestTooLarge, + make_request_too_large_failure(), + )) + } else if e.metadata().contains_key(MESSAGE_TOO_LARGE_KEY) { + Some(( + WorkflowTaskFailedCause::GrpcMessageTooLarge, + make_grpc_message_too_large_failure(), + )) + } else { + // Client layer rejected the completion for exceeding the worker's + // payload error limit. + payload_limit_violation_from(&e).map(|violation| { + ( + WorkflowTaskFailedCause::PayloadsTooLarge, + make_payloads_too_large_failure(violation), + ) + }) + }; + if let Some((cause, failure)) = cause_and_failure { + self.handle_activation_failed( + run_id, + completion_time, + FailedActivationWFTReport::new( + task_token, + attempt, + cause, + failure, + WftFailureKind::Task, + &run_metrics, + ), + ) + .await; } return Err(e); } @@ -510,36 +527,53 @@ impl Workflows { } } + /// The single point through which every workflow task failure passes on its way to the + /// server. A task failure is only sent for the first attempt of a task. Later attempts almost + /// always fail the same way, so reporting them would just spam the server with failures it + /// already knows about. Instead they are left to time out. async fn handle_activation_failed( &self, run_id: &str, completion_time: Instant, - outcome: FailedActivationWFTReport, + report: FailedActivationWFTReport, ) -> WFTReportStatus { - match outcome { - FailedActivationWFTReport::Report(tt, cause, failure) => { + let FailedActivationWFTReport { + task_token, + attempt, + cause, + failure, + kind, + } = report; + match kind { + WftFailureKind::LegacyQuery => { + warn!(run_id=%run_id, failure=?failure, "Failing legacy query request"); + self.respond_legacy_query(task_token, LegacyQueryResult::Failed(failure)) + .await; + } + WftFailureKind::RetryableLegacyQuery => { + debug!(run_id=%run_id, failure=?failure, + "Dropping legacy query with retryable failure"); + return WFTReportStatus::DropWft { completion_time }; + } + WftFailureKind::Task if attempt > 1 => { + debug!(run_id=%run_id, attempt, failure=?failure, + "Not reporting workflow task failure on non-first attempt"); + return WFTReportStatus::DropWft { completion_time }; + } + WftFailureKind::Task => { warn!(run_id=%run_id, failure=?failure, "Failing workflow task"); self.handle_wft_reporting_errs(run_id, || async { self.client - .fail_workflow_task(tt, cause, failure.failure) + .fail_workflow_task(task_token, cause, failure.failure) .await }) .await; - WFTReportStatus::Reported { - reset_last_started_to: None, - completion_time, - } - } - FailedActivationWFTReport::ReportLegacyQueryFailure(task_token, failure) => { - warn!(run_id=%run_id, failure=?failure, "Failing legacy query request"); - self.respond_legacy_query(task_token, LegacyQueryResult::Failed(failure)) - .await; - WFTReportStatus::Reported { - reset_last_started_to: None, - completion_time, - } } } + WFTReportStatus::Reported { + reset_last_started_to: None, + completion_time, + } } async fn handle_activation_completed_result( @@ -560,13 +594,10 @@ impl Workflows { ) .await } - ActivationCompleteOutcome::ReportWFTFail(outcome) => { - self.handle_activation_failed(run_id, completion_time, outcome) + ActivationCompleteOutcome::ReportWFTFail(report) => { + self.handle_activation_failed(run_id, completion_time, *report) .await } - ActivationCompleteOutcome::WFTFailedDontReport => { - WFTReportStatus::DropWft { completion_time } - } ActivationCompleteOutcome::DoNothing => WFTReportStatus::NotReported, } } @@ -681,7 +712,7 @@ impl Workflows { run_id: run_id.into(), message: message.into(), reason, - auto_reply_fail_tt: None, + auto_reply_fail: None, }); } @@ -1090,9 +1121,7 @@ enum WorkflowStreamAction { #[display("FailUnstoredWft(run_id={run_id})")] FailUnstoredWft { run_id: String, - task_token: TaskToken, - cause: WorkflowTaskFailedCause, - failure: Failure, + report: Box, }, } @@ -1213,10 +1242,59 @@ struct WorkflowTaskInfo { wf_id: String, } +/// Everything needed to tell the server a workflow task failed. Every path that fails a WFT +/// must produce one of these via [FailedActivationWFTReport::new] and feed it to +/// [Workflows::handle_activation_failed]. #[derive(Debug)] -enum FailedActivationWFTReport { - Report(TaskToken, WorkflowTaskFailedCause, Failure), - ReportLegacyQueryFailure(TaskToken, Failure), +struct FailedActivationWFTReport { + task_token: TaskToken, + attempt: u32, + cause: WorkflowTaskFailedCause, + failure: Failure, + kind: WftFailureKind, +} +impl FailedActivationWFTReport { + /// Records the task-failed metric as part of constructing the report. + fn new( + task_token: TaskToken, + attempt: u32, + cause: WorkflowTaskFailedCause, + failure: Failure, + kind: WftFailureKind, + metrics: &MetricsContext, + ) -> Self { + metrics + .with_new_attrs([metrics::failure_reason(cause.into())]) + .wf_task_failed(); + Self { + task_token, + attempt, + cause, + failure, + kind, + } + } +} + +#[derive(Debug, Copy, Clone, PartialEq, Eq)] +enum WftFailureKind { + /// A normal workflow task, failed via the task failure API. + Task, + /// A legacy query, answered through the query response API rather than by failing the task. + /// The caller is blocked until an answer arrives, so these are always reported. + LegacyQuery, + /// A legacy query whose failure is transient and may succeed if retried. Failing the query + /// would surface that transient error to the caller, so the task is dropped unanswered. + RetryableLegacyQuery, +} + +/// Identifies a WFT whose failure must be reported outside the normal activation completion path, +/// because the run it belongs to was never stored or is being torn down. +#[derive(Debug, Clone)] +struct UnstoredWftFailInfo { + task_token: TaskToken, + attempt: u32, + workflow_type: String, } struct ServerCommandsWithWorkflowInfo { @@ -1252,16 +1330,17 @@ pub(crate) enum ActivationAction { #[derive(Debug)] #[allow(clippy::large_enum_variant)] enum EvictionRequestResult { - EvictionRequested(Option, RunUpdateAct), + EvictionRequested(RunUpdateAct), NotFound, - EvictionAlreadyRequested(Option), + EvictionAlreadyRequested, } impl EvictionRequestResult { fn into_run_update_resp(self) -> RunUpdateAct { match self { - EvictionRequestResult::EvictionRequested(_, resp) => resp, - EvictionRequestResult::NotFound - | EvictionRequestResult::EvictionAlreadyRequested(_) => None, + EvictionRequestResult::EvictionRequested(resp) => resp, + EvictionRequestResult::NotFound | EvictionRequestResult::EvictionAlreadyRequested => { + None + } } } } @@ -1342,9 +1421,9 @@ struct RequestEvictMsg { message: String, reason: EvictionReason, /// If set, we requested eviction because something went wrong processing a brand new poll task, - /// which means we won't have stored the WFT and we need to track the task token separately so - /// we can reply with a failure to server after the evict goes through. - auto_reply_fail_tt: Option, + /// which means we won't have stored the WFT and we need to track it separately so we can + /// reply with a failure to server after the evict goes through. + auto_reply_fail: Option, } #[derive(Debug)] pub(crate) struct HeartbeatTimeoutMsg { @@ -1490,12 +1569,9 @@ enum ActivationCompleteOutcome { /// The WFT must be reported as successful to the server using the contained information. ReportWFTSuccess(ServerCommandsWithWorkflowInfo), /// The WFT must be reported as failed to the server using the contained information. - ReportWFTFail(FailedActivationWFTReport), + ReportWFTFail(Box), /// There's nothing to do right now. EX: The workflow needs to keep replaying. DoNothing, - /// The workflow task failed, but we shouldn't report it. EX: We have failed 2 or more attempts - /// in a row. - WFTFailedDontReport, } /// Did we report, or not, completion of a WFT to server? #[derive(Debug, Copy, Clone)] @@ -1508,7 +1584,7 @@ enum WFTReportStatus { /// work to be done. EX: Running LAs. NotReported, /// We didn't report, but we want to clear the outstanding workflow task anyway. See - /// [ActivationCompleteOutcome::WFTFailedDontReport]. + /// [Workflows::handle_activation_failed] for when this happens. DropWft { completion_time: Instant }, } impl WFTReportStatus { @@ -1686,7 +1762,61 @@ struct EmptyWorkflowCommandErr; #[display("{}", variant)] struct WFCommand { variant: WFCommandVariant, + annotations: CommandAnnotations, +} + +impl WFCommand { + fn new(variant: WFCommandVariant) -> Self { + Self { + variant, + annotations: CommandAnnotations::default(), + } + } +} + +/// The lang-supplied decorations that ride along on a [WFCommand] and end up on the [ProtoCommand] +/// we send to the server. They are kept together because a command machine must remember them in +/// order to repeat them on any further command it issues, most notably a cancellation. +#[derive(Debug, Default, Clone, PartialEq)] +struct CommandAnnotations { metadata: Option, + event_group_markers: Vec, +} + +impl CommandAnnotations { + /// Apply annotations lang attached to a cancellation command on top of the ones the command + /// being cancelled carried. Anything lang set explicitly wins; anything it left out is + /// inherited, which is what makes a cancellation land in the same event group as the command + /// it cancels even when it is issued from somewhere no group is active. + fn override_with(&mut self, other: Self) { + if let Some(other_metadata) = other.metadata { + let metadata = self.metadata.get_or_insert_with(UserMetadata::default); + if let Some(summary) = other_metadata.summary { + metadata.summary = Some(summary); + } + if let Some(details) = other_metadata.details { + metadata.details = Some(details); + } + } + if !other.event_group_markers.is_empty() { + self.event_group_markers = other.event_group_markers; + } + } +} + +trait ProtoCommandExt { + fn new(attributes: Attributes, annotations: CommandAnnotations) -> Self; +} + +impl ProtoCommandExt for ProtoCommand { + fn new(attributes: Attributes, annotations: CommandAnnotations) -> Self { + Self { + command_type: attributes.as_type() as i32, + attributes: Some(attributes), + user_metadata: annotations.metadata, + event_group_markers: annotations.event_group_markers, + } + } } #[derive(Debug, derive_more::From, derive_more::Display)] @@ -1823,7 +1953,10 @@ impl TryFrom for WFCommand { }; Ok(Self { variant, - metadata: c.user_metadata, + annotations: CommandAnnotations { + metadata: c.user_metadata, + event_group_markers: c.event_group_markers, + }, }) } } @@ -2114,6 +2247,25 @@ fn make_grpc_message_too_large_failure() -> Failure { } } +fn make_request_too_large_failure() -> Failure { + Failure { + failure: Some( + temporalio_common::protos::temporal::api::failure::v1::Failure { + message: "Workflow task completion exceeds the namespace size limit".to_string(), + failure_info: Some(FailureInfo::ApplicationFailureInfo( + ApplicationFailureInfo { + r#type: "RequestTooLarge".to_string(), + non_retryable: true, + ..Default::default() + }, + )), + ..Default::default() + }, + ), + force_cause: WorkflowTaskFailedCause::RequestTooLarge as i32, + } +} + fn make_payloads_too_large_failure(violation: &PayloadLimitViolation) -> Failure { Failure { failure: Some( diff --git a/crates/sdk-core/src/worker/workflow/wft_extraction.rs b/crates/sdk-core/src/worker/workflow/wft_extraction.rs index f08e2f6a9..bb97f15b1 100644 --- a/crates/sdk-core/src/worker/workflow/wft_extraction.rs +++ b/crates/sdk-core/src/worker/workflow/wft_extraction.rs @@ -5,14 +5,14 @@ use crate::{ WorkflowSlotKind, client::WorkerClient, workflow::{ - CacheMissFetchReq, HistoryUpdate, NextPageReq, PermittedWFT, + CacheMissFetchReq, HistoryUpdate, NextPageReq, PermittedWFT, UnstoredWftFailInfo, history_update::HistoryPaginator, }, }, }; use futures_util::{FutureExt, Stream, StreamExt, stream, stream::PollNext}; use std::{future, sync::Arc}; -use temporalio_common::protos::{TaskToken, coresdk::WorkflowSlotInfo}; +use temporalio_common::protos::coresdk::WorkflowSlotInfo; use tracing::Span; /// Transforms incoming validated WFTs and history fetching requests into [PermittedWFT]s ready @@ -35,7 +35,7 @@ pub(super) enum WFTExtractorOutput { FailedFetch { run_id: String, err: tonic::Status, - auto_reply_fail_tt: Option, + auto_reply_fail: Option, }, PollerDead, } @@ -71,7 +71,11 @@ impl WFTExtractor { match stream_in { Ok((wft, permit)) => { let run_id = wft.workflow_execution.run_id.clone(); - let tt = wft.task_token.clone(); + let fail_info = UnstoredWftFailInfo { + task_token: wft.task_token.clone(), + attempt: wft.attempt, + workflow_type: wft.workflow_type.clone(), + }; Ok(match HistoryPaginator::from_poll(wft, client).await { Ok((pag, prep)) => WFTExtractorOutput::NewWFT(PermittedWFT { permit: permit.into_used(WorkflowSlotInfo { @@ -84,7 +88,7 @@ impl WFTExtractor { Err(err) => WFTExtractorOutput::FailedFetch { run_id, err, - auto_reply_fail_tt: Some(tt), + auto_reply_fail: Some(fail_info), }, }) } @@ -111,13 +115,17 @@ impl WFTExtractor { // failure. We'll just proceed with shutdown. HistoryFetchReq::Full(req, rc) => { let run_id = req.original_wft.work.execution.run_id.clone(); - let task_token = req.original_wft.work.task_token.clone(); + let fail_info = UnstoredWftFailInfo { + task_token: req.original_wft.work.task_token.clone(), + attempt: req.original_wft.work.attempt, + workflow_type: req.original_wft.work.workflow_type.clone(), + }; match HistoryPaginator::from_fetchreq(req, client).await { Ok(r) => WFTExtractorOutput::FetchResult(r, rc), Err(err) => WFTExtractorOutput::FailedFetch { run_id, err, - auto_reply_fail_tt: Some(task_token), + auto_reply_fail: Some(fail_info), }, } } @@ -132,7 +140,7 @@ impl WFTExtractor { Err(err) => WFTExtractorOutput::FailedFetch { run_id: req.paginator.run_id, err, - auto_reply_fail_tt: None, + auto_reply_fail: None, }, } } diff --git a/crates/sdk-core/src/worker/workflow/wft_poller.rs b/crates/sdk-core/src/worker/workflow/wft_poller.rs index 06f78a6a8..ee6dee2db 100644 --- a/crates/sdk-core/src/worker/workflow/wft_poller.rs +++ b/crates/sdk-core/src/worker/workflow/wft_poller.rs @@ -43,9 +43,13 @@ pub(crate) fn make_wft_poller( &capabilities, ); let wft_poller_shared = if sticky_queue_name.is_some() { - Some(Arc::new(WFTPollerShared::new( - wft_slots.available_permits(), - ))) + // Balance on the limit `acquire_owned` actually enforces (min of slot supplier and cache + // size). Using only the slot supplier lets a small cache starve the non-sticky poller. + let balance_limit = [wft_slots.available_permits(), wft_slots.max_permits()] + .into_iter() + .flatten() + .min(); + Some(Arc::new(WFTPollerShared::new(balance_limit))) } else { None }; @@ -94,6 +98,7 @@ pub(crate) fn make_wft_poller( pub(crate) struct WFTPollerShared { last_seen_sticky_backlog: (watch::Receiver, watch::Sender), sticky_active: OnceLock>, + sticky_target: OnceLock>, non_sticky_active: OnceLock>, max_slots: Option, } @@ -103,15 +108,21 @@ impl WFTPollerShared { Self { last_seen_sticky_backlog: (rx, tx), sticky_active: OnceLock::new(), + sticky_target: OnceLock::new(), non_sticky_active: OnceLock::new(), max_slots, } } - pub(crate) fn set_sticky_active(&self, rx: watch::Receiver) { - let _ = self.sticky_active.set(rx); + pub(crate) fn set_sticky_active( + &self, + active_rx: watch::Receiver, + target_rx: watch::Receiver, + ) { + let _ = self.sticky_active.set(active_rx); + let _ = self.sticky_target.set(target_rx); } - pub(crate) fn set_non_sticky_active(&self, rx: watch::Receiver) { - let _ = self.non_sticky_active.set(rx); + pub(crate) fn set_non_sticky_active(&self, active_rx: watch::Receiver) { + let _ = self.non_sticky_active.set(active_rx); } /// Makes either the sticky or non-sticky poller wait pre-permit-acquisition so that we can /// balance which kind of queue we poll appropriately. @@ -120,17 +131,23 @@ impl WFTPollerShared { // that we won't end up using every available permit with one kind of poller. In practice // this is only ever likely to be an issue with very small numbers of slots. if let Some(max_slots) = self.max_slots - && let Some((sticky_active, non_sticky_active)) = - self.sticky_active.get().zip(self.non_sticky_active.get()) + && let Some(sticky_active) = self.sticky_active.get() + && let Some(sticky_target) = self.sticky_target.get() + && let Some(non_sticky_active) = self.non_sticky_active.get() { let mut sticky_active = sticky_active.clone(); + let mut sticky_target = sticky_target.clone(); let mut non_sticky_active = non_sticky_active.clone(); let mut sticky_backlog = self.last_seen_sticky_backlog.0.clone(); loop { let num_sticky_active = *sticky_active.borrow_and_update(); + let num_sticky_target = *sticky_target.borrow_and_update(); let num_non_sticky_active = *non_sticky_active.borrow_and_update(); let num_sticky_backlog = *sticky_backlog.borrow_and_update(); + let sticky_should_be_chosen = num_sticky_backlog > 1 + && num_sticky_backlog > num_sticky_active + && num_sticky_active < num_sticky_target; let allow = || { if !is_sticky { @@ -145,7 +162,7 @@ impl WFTPollerShared { } // If there's a meaningful sticky backlog, prioritize sticky. - if num_sticky_backlog > 1 && num_sticky_backlog > num_sticky_active { + if sticky_should_be_chosen { return false; } } else { @@ -154,13 +171,18 @@ impl WFTPollerShared { return true; } + // Do not reserve a slot that the sticky scaler cannot use. + if num_sticky_active >= num_sticky_target { + return false; + } + // Do not allow an additional sticky poller to prevent starting a first non-sticky poller. if num_non_sticky_active == 0 && num_sticky_active + 1 >= max_slots { return false; } // If there's a meaningful sticky backlog, prioritize sticky. - if num_sticky_backlog > 1 && num_sticky_backlog > num_sticky_active { + if sticky_should_be_chosen { return true; } } @@ -180,6 +202,7 @@ impl WFTPollerShared { tokio::select! { _ = sticky_active.changed() => (), + _ = sticky_target.changed() => (), _ = non_sticky_active.changed() => (), _ = sticky_backlog.changed() => (), } @@ -260,11 +283,72 @@ pub(crate) fn validate_wft( mod tests { use super::*; use crate::{ - abstractions::tests::fixed_size_permit_dealer, pollers::MockPermittedPollBuffer, - test_help::mock_poller, worker::WorkflowSlotKind, + abstractions::tests::fixed_size_permit_dealer, + pollers::MockPermittedPollBuffer, + replay::TestHistoryBuilder, + test_help::{ResponseType, hist_to_poll_resp, mock_poller, test_worker_cfg}, + worker::{ + PollerBehavior, WorkflowSlotKind, client::mocks::mock_manual_worker_client, + tuner::FixedSizeSlotSupplier, + }, }; - use futures_util::{StreamExt, pin_mut}; - use std::sync::Arc; + use futures_util::{FutureExt, StreamExt, pin_mut}; + use std::{ + future::Future, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + task::{Context, Poll, Wake}, + }; + use temporalio_common::protos::temporal::api::enums::v1::EventType; + + struct WakeCounter(AtomicUsize); + + impl Wake for WakeCounter { + fn wake(self: Arc) { + self.0.fetch_add(1, Ordering::Relaxed); + } + } + + #[tokio::test] + async fn sticky_priority_respects_poller_target() { + let shared = WFTPollerShared::new(Some(4)); + let (_sticky_tx, sticky_rx) = watch::channel(2); + let (_non_sticky_tx, non_sticky_rx) = watch::channel(1); + let (sticky_target_tx, sticky_target_rx) = watch::channel(3); + shared.set_sticky_active(sticky_rx, sticky_target_rx); + shared.set_non_sticky_active(non_sticky_rx); + shared.record_sticky_backlog(10); + + assert!(shared.wait_if_needed(false).now_or_never().is_none()); + + sticky_target_tx.send_replace(2); + assert!(shared.wait_if_needed(false).now_or_never().is_some()); + assert!(shared.wait_if_needed(true).now_or_never().is_none()); + } + + #[tokio::test] + async fn target_change_wakes_balancer() { + let shared = WFTPollerShared::new(Some(4)); + let (_sticky_tx, sticky_rx) = watch::channel(2); + let (_non_sticky_tx, non_sticky_rx) = watch::channel(1); + let (sticky_target_tx, sticky_target_rx) = watch::channel(3); + shared.set_sticky_active(sticky_rx, sticky_target_rx); + shared.set_non_sticky_active(non_sticky_rx); + shared.record_sticky_backlog(10); + + let waiter = shared.wait_if_needed(false); + pin_mut!(waiter); + let wake_counter = Arc::new(WakeCounter(AtomicUsize::new(0))); + let waker = wake_counter.clone().into(); + let mut context = Context::from_waker(&waker); + assert_eq!(waiter.as_mut().poll(&mut context), Poll::Pending); + + sticky_target_tx.send_replace(2); + assert_ne!(wake_counter.0.load(Ordering::Relaxed), 0); + assert_eq!(waiter.as_mut().poll(&mut context), Poll::Ready(())); + } #[tokio::test] async fn poll_timeouts_do_not_produce_responses() { @@ -303,6 +387,72 @@ mod tests { assert_matches!(stream.next().await, None); } + /// Cache (`max_permits`) of 2 under a supplier of 10: the balancer must reserve a poll slot + /// against the cache, not the supplier, or a saturated sticky poller starves the non-sticky one. + /// Sticky long-polls forever (holding permits) while non-sticky repeatedly times out and must + /// re-acquire; without the reservation sticky recaptures every freed permit, so non-sticky never + /// delivers its eventual real task. (A single fresh-start poll won't reproduce this: a free + /// permit is always available at startup.) + #[tokio::test] + async fn small_cache_does_not_starve_nonsticky_poller() { + let mut t = TestHistoryBuilder::default(); + t.add_by_type(EventType::WorkflowExecutionStarted); + t.add_full_wf_task(); + let real_task = hist_to_poll_resp(&t, "wf-id", ResponseType::AllHistory).resp; + + // Non-sticky times out (empty response) this many times before its real task, giving the + // sticky poller ample opportunity to grab every freed permit if nothing is reserved. + let nonsticky_timeouts = Arc::new(AtomicUsize::new(0)); + let mut client = mock_manual_worker_client(); + client + .expect_poll_workflow_task() + .returning(move |_po, wfo| { + if wfo.sticky_queue_name.is_some() { + // Sticky poll never returns -> holds its permit for the test's duration. + std::future::pending().boxed() + } else if nonsticky_timeouts.fetch_add(1, Ordering::SeqCst) < 20 { + // Poll timeout: empty response releases the permit so it must be re-acquired. + async { Ok(PollWorkflowTaskQueueResponse::default()) }.boxed() + } else { + let real_task = real_task.clone(); + async move { Ok(real_task) }.boxed() + } + }); + + let wft_slots = MeteredPermitDealer::::new( + Arc::new(FixedSizeSlotSupplier::new(10)), + MetricsContext::no_op(), + Some(2), + Arc::new(Default::default()), + None, + ); + // Poller max > 1 so the sticky poller alone could otherwise claim every permit. + let cfg = { + let mut cfg = test_worker_cfg().build().unwrap(); + cfg.workflow_task_poller_behavior = Some(PollerBehavior::SimpleMaximum(5_usize)); + cfg + }; + + let client: Arc = Arc::new(client); + let stream = make_wft_poller( + &cfg, + &Some("sticky-q".to_string()), + &client, + &MetricsContext::no_op(), + &CancellationToken::new(), + &wft_slots, + Arc::new(AtomicCell::new(None)), + Arc::new(AtomicCell::new(None)), + Arc::new(NamespaceCapabilities::default()), + ); + pin_mut!(stream); + let got = tokio::time::timeout(std::time::Duration::from_secs(10), stream.next()) + .await + .expect("non-sticky poll must be delivered; a small cache is starving it") + .expect("stream should yield a task"); + assert!(got.is_ok()); + } + #[tokio::test] async fn poll_errors_do_produce_responses() { let mut mock_poller = mock_poller(); diff --git a/crates/sdk-core/src/worker/workflow/workflow_stream.rs b/crates/sdk-core/src/worker/workflow/workflow_stream.rs index 4317766cb..6e96a9817 100644 --- a/crates/sdk-core/src/worker/workflow/workflow_stream.rs +++ b/crates/sdk-core/src/worker/workflow/workflow_stream.rs @@ -2,6 +2,7 @@ use super::external_streams::{ExternalStreamReadyResult, ExternalStreamRunStatus use crate::{ MetricsContext, abstractions::dbg_panic, + telemetry::metrics::workflow_type, worker::workflow::{ managed_run::RunUpdateAct, run_cache::RunCache, @@ -232,17 +233,24 @@ impl WFStream { WFStreamInput::FailedFetch { run_id, err, - auto_reply_fail_tt, + auto_reply_fail, } => { let message = format!("Fetching history failed: {err:?}"); if !state.runs.has_run(&run_id) - && let Some(task_token) = auto_reply_fail_tt.clone() + && let Some(info) = auto_reply_fail.clone() { actions.push(WorkflowStreamAction::FailUnstoredWft { run_id, - task_token, - cause: WorkflowTaskFailedCause::WorkflowWorkerUnhandledFailure, - failure: ApiFailure::application_failure(message, true).into(), + report: Box::new(FailedActivationWFTReport::new( + info.task_token, + info.attempt, + WorkflowTaskFailedCause::WorkflowWorkerUnhandledFailure, + ApiFailure::application_failure(message, true).into(), + WftFailureKind::Task, + &state + .metrics + .with_new_attrs([workflow_type(info.workflow_type)]), + )), }); None } else { @@ -251,7 +259,7 @@ impl WFStream { run_id, message, reason: EvictionReason::PaginationOrHistoryFetch, - auto_reply_fail_tt, + auto_reply_fail, }) .into_run_update_resp() } @@ -529,7 +537,7 @@ impl WFStream { run_id: run_id.to_string(), message: "Workflow completed".to_string(), reason: EvictionReason::WorkflowExecutionEnding, - auto_reply_fail_tt: None, + auto_reply_fail: None, }) .into_run_update_resp() } @@ -619,7 +627,7 @@ impl WFStream { run_id, message: "Workflow cache full".to_string(), reason: EvictionReason::CacheFull, - auto_reply_fail_tt: None, + auto_reply_fail: None, }) } else { // This branch shouldn't really be possible @@ -708,7 +716,7 @@ impl WFStream { run_id, message: "Workflow cache full".to_string(), reason: EvictionReason::CacheFull, - auto_reply_fail_tt: None, + auto_reply_fail: None, }) .into_run_update_resp(), ); @@ -778,7 +786,7 @@ enum WFStreamInput { FailedFetch { run_id: String, err: tonic::Status, - auto_reply_fail_tt: Option, + auto_reply_fail: Option, }, } impl From for WFStreamInput { @@ -884,7 +892,7 @@ enum ExternalPollerInputs { FailedFetch { run_id: String, err: tonic::Status, - auto_reply_fail_tt: Option, + auto_reply_fail: Option, }, } impl From for WFStreamInput { @@ -897,11 +905,11 @@ impl From for WFStreamInput { ExternalPollerInputs::FailedFetch { run_id, err, - auto_reply_fail_tt, + auto_reply_fail, } => WFStreamInput::FailedFetch { run_id, err, - auto_reply_fail_tt, + auto_reply_fail, }, ExternalPollerInputs::NextPage { paginator, @@ -934,11 +942,11 @@ impl From> for ExternalPollerInputs { Ok(WFTExtractorOutput::FailedFetch { run_id, err, - auto_reply_fail_tt, + auto_reply_fail, }) => ExternalPollerInputs::FailedFetch { run_id, err, - auto_reply_fail_tt, + auto_reply_fail, }, Ok(WFTExtractorOutput::PollerDead) => ExternalPollerInputs::PollerDead, Err(e) => ExternalPollerInputs::PollerError(e), diff --git a/crates/sdk-core/tests/activities_procmacro.rs b/crates/sdk-core/tests/activities_procmacro.rs index f51371952..01e73fc7a 100644 --- a/crates/sdk-core/tests/activities_procmacro.rs +++ b/crates/sdk-core/tests/activities_procmacro.rs @@ -1,6 +1,5 @@ #[test] fn activities_procmacro_build_tests() { let t = trybuild::TestCases::new(); - t.pass("tests/activities_trybuild/*_pass.rs"); t.compile_fail("tests/activities_trybuild/*_fail.rs"); } diff --git a/crates/sdk-core/tests/activities_trybuild/basic_pass.rs b/crates/sdk-core/tests/activities_trybuild/basic_pass.rs deleted file mode 100644 index ac37a2ca0..000000000 --- a/crates/sdk-core/tests/activities_trybuild/basic_pass.rs +++ /dev/null @@ -1,54 +0,0 @@ -use std::sync::Arc; -use temporalio_macros::activities; -use temporalio_sdk::activities::{ActivityContext, ActivityError}; - -pub struct MyActivities; - -#[activities] -impl MyActivities { - #[activity] - pub async fn static_activity( - _ctx: ActivityContext, - _in: String, - ) -> Result { - Ok("Can be static".to_string()) - } - - #[activity] - pub async fn activity( - self: Arc, - _ctx: ActivityContext, - _in: bool, - ) -> Result { - Ok("I'm done!".to_string()) - } - - #[activity] - pub async fn activity_arc_fully_qualified( - self: std::sync::Arc, - _ctx: ActivityContext, - _in: bool, - ) -> Result { - Ok("I'm done!".to_string()) - } - - #[activity] - pub fn sync_activity(_ctx: ActivityContext, _in: bool) -> Result { - Ok("Sync activities are supported too".to_string()) - } -} - -pub struct MyActivitiesStatic; - -#[activities] -impl MyActivitiesStatic { - #[activity] - pub async fn static_activity( - _ctx: ActivityContext, - _in: String, - ) -> Result { - Ok("Can be static".to_string()) - } -} - -fn main() {} diff --git a/crates/sdk-core/tests/activities_trybuild/multi_arg_pass.rs b/crates/sdk-core/tests/activities_trybuild/multi_arg_pass.rs deleted file mode 100644 index 659761197..000000000 --- a/crates/sdk-core/tests/activities_trybuild/multi_arg_pass.rs +++ /dev/null @@ -1,48 +0,0 @@ -use std::sync::Arc; -use temporalio_macros::activities; -use temporalio_sdk::activities::{ActivityContext, ActivityError}; - -pub struct MultiArgActivities; - -#[activities] -impl MultiArgActivities { - #[activity] - pub async fn two_args( - _ctx: ActivityContext, - _a: String, - _b: i32, - ) -> Result { - Ok("done".to_string()) - } - - #[activity] - pub async fn three_args( - _ctx: ActivityContext, - _a: String, - _b: i32, - _c: bool, - ) -> Result { - Ok("done".to_string()) - } - - #[activity] - pub async fn instance_two_args( - self: Arc, - _ctx: ActivityContext, - _a: String, - _b: i32, - ) -> Result { - Ok("done".to_string()) - } - - #[activity] - pub fn sync_two_args( - _ctx: ActivityContext, - _a: String, - _b: i32, - ) -> Result { - Ok("done".to_string()) - } -} - -fn main() {} diff --git a/crates/sdk-core/tests/activities_trybuild/no_input_pass.rs b/crates/sdk-core/tests/activities_trybuild/no_input_pass.rs deleted file mode 100644 index f7af64181..000000000 --- a/crates/sdk-core/tests/activities_trybuild/no_input_pass.rs +++ /dev/null @@ -1,14 +0,0 @@ -use temporalio_macros::activities; -use temporalio_sdk::activities::{ActivityContext, ActivityError}; - -pub struct SimpleActivities; - -#[activities] -impl SimpleActivities { - #[activity] - pub async fn no_input_activity(_ctx: ActivityContext) -> Result { - Ok("No input needed".to_string()) - } -} - -fn main() {} diff --git a/crates/sdk-core/tests/activities_trybuild/no_return_type_pass.rs b/crates/sdk-core/tests/activities_trybuild/no_return_type_pass.rs deleted file mode 100644 index b4697aa11..000000000 --- a/crates/sdk-core/tests/activities_trybuild/no_return_type_pass.rs +++ /dev/null @@ -1,19 +0,0 @@ -use temporalio_macros::activities; -use temporalio_sdk::activities::ActivityContext; - -pub struct VoidActivities; - -#[activities] -impl VoidActivities { - #[activity] - pub async fn no_return(_ctx: ActivityContext, _in: String) { - println!("Doing work..."); - } - - #[activity] - pub fn sync_no_return(_ctx: ActivityContext) { - println!("Sync work..."); - } -} - -fn main() {} diff --git a/crates/sdk-core/tests/cloud_namespace/mod.rs b/crates/sdk-core/tests/cloud_namespace/mod.rs new file mode 100644 index 000000000..78590e884 --- /dev/null +++ b/crates/sdk-core/tests/cloud_namespace/mod.rs @@ -0,0 +1,200 @@ +use anyhow::{Context, bail}; +use std::{collections::HashMap, env, fs::OpenOptions, io::Write, time::Duration}; +use temporalio_client::{Connection, ConnectionOptions, TlsOptions, grpc::CloudService}; +use temporalio_common::protos::temporal::api::cloud::{ + cloudservice::v1::{ + CreateNamespaceRequest, DeleteNamespaceRequest, GetAsyncOperationRequest, + GetNamespaceRequest, + }, + namespace::v1::{MtlsAuthSpec, NamespaceSpec, ReplicaSpec}, + operation::v1::{AsyncOperation, async_operation}, +}; +use tokio::time::Instant; +use tonic::IntoRequest; +use url::Url; +use uuid::Uuid; + +const CLOUD_OPS_ADDRESS: &str = "https://saas-api.tmprl.cloud:443"; +const CLOUD_REGION: &str = "aws-ca-central-1"; +const OPERATION_TIMEOUT: Duration = Duration::from_secs(10 * 60); +const DEFAULT_POLL_DELAY: Duration = Duration::from_secs(10); +const MIN_POLL_DELAY: Duration = Duration::from_secs(1); + +pub(crate) async fn create_namespace() -> anyhow::Result<()> { + let namespace_name = format!( + "sdk-rust-ci-{}-{}", + required_env("GITHUB_RUN_ID")?, + required_env("GITHUB_RUN_ATTEMPT")? + ); + let accepted_client_ca = tokio::fs::read(required_env("TEMPORAL_CLOUD_CLIENT_CA_PATH")?) + .await + .context("failed to read the Cloud test CA certificate")?; + let connection = cloud_connection().await?; + let mut client = connection.cloud_service(); + let response = client + .create_namespace( + CreateNamespaceRequest { + spec: Some(NamespaceSpec { + name: namespace_name, + retention_days: 1, + mtls_auth: Some(MtlsAuthSpec { + accepted_client_ca, + enabled: true, + ..Default::default() + }), + replicas: vec![ReplicaSpec { + region: CLOUD_REGION.to_owned(), + }], + ..Default::default() + }), + async_operation_id: Uuid::new_v4().to_string(), + ..Default::default() + } + .into_request(), + ) + .await + .context("failed to create Cloud namespace")? + .into_inner(); + + if response.namespace.is_empty() { + bail!("create namespace response did not include a namespace"); + } + append_github_output("namespace", &response.namespace)?; + wait_for_operation( + client.as_mut(), + response + .async_operation + .context("create namespace response did not include an operation")?, + ) + .await +} + +pub(crate) async fn delete_namespace(namespace: String) -> anyhow::Result<()> { + let connection = cloud_connection().await?; + let mut client = connection.cloud_service(); + let existing = client + .get_namespace( + GetNamespaceRequest { + namespace: namespace.clone(), + } + .into_request(), + ) + .await + .context("failed to read Cloud namespace before deletion")? + .into_inner(); + let resource_version = existing + .namespace + .map(|namespace| namespace.resource_version) + .filter(|version| !version.is_empty()) + .context("Cloud namespace did not include a resource version")?; + let response = client + .delete_namespace( + DeleteNamespaceRequest { + namespace, + resource_version, + async_operation_id: Uuid::new_v4().to_string(), + } + .into_request(), + ) + .await + .context("failed to delete Cloud namespace")? + .into_inner(); + wait_for_operation( + client.as_mut(), + response + .async_operation + .context("delete namespace response did not include an operation")?, + ) + .await +} + +async fn cloud_connection() -> anyhow::Result { + let api_version = required_env("TEMPORAL_CLIENT_CLOUD_API_VERSION")?; + let options = ConnectionOptions::new(Url::parse(CLOUD_OPS_ADDRESS)?) + .api_key(required_env("TEMPORAL_CLIENT_CLOUD_API_KEY")?) + .headers(HashMap::from([( + "temporal-cloud-api-version".to_owned(), + api_version, + )])) + .tls_options(TlsOptions::default()) + // The Cloud Operations endpoint does not expose the Workflow Service probe. + .skip_get_system_info(true) + .build(); + Connection::connect(options) + .await + .context("failed to connect to the Cloud Operations API") +} + +async fn wait_for_operation( + client: &mut dyn CloudService, + operation: AsyncOperation, +) -> anyhow::Result<()> { + if operation.id.is_empty() { + bail!("Cloud operation response did not include an ID"); + } + let operation_id = operation.id; + let deadline = Instant::now() + OPERATION_TIMEOUT; + + loop { + let operation = client + .get_async_operation( + GetAsyncOperationRequest { + async_operation_id: operation_id.clone(), + } + .into_request(), + ) + .await? + .into_inner() + .async_operation + .with_context(|| { + format!("Cloud operation {operation_id} response did not include an operation") + })?; + let state = async_operation::State::try_from(operation.state) + .with_context(|| format!("Cloud operation {operation_id} had an unknown state"))?; + + match state { + async_operation::State::Fulfilled => return Ok(()), + async_operation::State::Failed + | async_operation::State::Cancelled + | async_operation::State::Rejected => { + bail!( + "Cloud operation {operation_id} {}: {}", + state.as_str_name(), + operation.failure_reason + ); + } + async_operation::State::Unspecified + | async_operation::State::Pending + | async_operation::State::InProgress => {} + } + + let remaining = deadline.saturating_duration_since(Instant::now()); + if remaining.is_zero() { + bail!("timed out waiting for Cloud operation {operation_id}"); + } + let delay = operation + .check_duration + .and_then(|duration| Duration::try_from(duration).ok()) + .unwrap_or(DEFAULT_POLL_DELAY) + .max(MIN_POLL_DELAY) + .min(remaining); + tokio::time::sleep(delay).await; + } +} + +fn append_github_output(name: &str, value: &str) -> anyhow::Result<()> { + let output_path = required_env("GITHUB_OUTPUT")?; + let mut output = OpenOptions::new() + .create(true) + .append(true) + .open(output_path) + .context("failed to open GITHUB_OUTPUT")?; + writeln!(output, "{name}={value}").context("failed to write GITHUB_OUTPUT") +} + +fn required_env(name: &str) -> anyhow::Result { + env::var(name) + .ok() + .filter(|value| !value.is_empty()) + .with_context(|| format!("missing required environment variable {name}")) +} diff --git a/crates/sdk-core/tests/common/http_proxy.rs b/crates/sdk-core/tests/common/http_proxy.rs deleted file mode 100644 index c39c3a6c5..000000000 --- a/crates/sdk-core/tests/common/http_proxy.rs +++ /dev/null @@ -1,134 +0,0 @@ -use bytes::Bytes; -use http_body_util::Empty; -use hyper::{ - Request, Response, StatusCode, body::Incoming, server::conn::http1, service::service_fn, -}; -use hyper_util::rt::TokioIo; -use std::{ - io, - sync::{ - Arc, - atomic::{AtomicUsize, Ordering}, - }, -}; -use temporalio_client::proxy::ProxyStream; -#[cfg(unix)] -use tokio::net::UnixListener; -use tokio::{ - net::{TcpListener, TcpStream}, - sync::oneshot, -}; - -pub(crate) struct HttpProxy { - proxy_hits: Arc, - shutdown_tx: oneshot::Sender<()>, -} -impl HttpProxy { - pub(crate) fn spawn_tcp(listener: TcpListener) -> Self { - Self::spawn(ProxyListener::Tcp(listener)) - } - - #[cfg(unix)] - pub(crate) fn spawn_unix(listener: UnixListener) -> Self { - Self::spawn(ProxyListener::Unix(listener)) - } - - fn spawn(listener: ProxyListener) -> Self { - let (shutdown_tx, mut shutdown_rx) = oneshot::channel::<()>(); - let proxy_hits = Arc::new(AtomicUsize::new(0)); - let proxy_hits_cloned = proxy_hits.clone(); - tokio::spawn(async move { - loop { - let proxy_hits_cloned = proxy_hits_cloned.clone(); - tokio::select! { - _ = &mut shutdown_rx => break, - stream = listener.accept() => { - let stream = match stream { - Ok(stream) => stream, - Err(e) => { println!("Proxy accept error: {e}"); continue; } - }; - tokio::spawn(async move { - if let Err(e) = http1::Builder::new() - .serve_connection( - TokioIo::new(stream), - service_fn(move |req| handle_connect(req, proxy_hits_cloned.clone())), - ) - .with_upgrades() - .await - { - println!("Proxy conn error: {e}"); - } - }); - } - } - } - }); - Self { - proxy_hits, - shutdown_tx, - } - } - - pub(crate) fn hit_count(&self) -> usize { - self.proxy_hits.load(Ordering::SeqCst) - } - - /// Returns before shutdown occurs - pub(crate) fn shutdown(self) { - let _ = self.shutdown_tx.send(()); - } -} - -async fn handle_connect( - req: Request, - counter: Arc, -) -> Result>, hyper::Error> { - if req.method() == hyper::Method::CONNECT { - // Increment atomic counter - counter.fetch_add(1, Ordering::SeqCst); - - // Tell the client the tunnel is established - tokio::spawn(async move { - if let Some(addr) = req.uri().authority().map(|a| a.as_str()) { - match TcpStream::connect(addr).await { - Ok(mut server_stream) => match hyper::upgrade::on(req).await { - Ok(upgraded) => { - let mut upgraded = TokioIo::new(upgraded); - let _ = - tokio::io::copy_bidirectional(&mut upgraded, &mut server_stream) - .await; - } - Err(err) => println!("Upgrade failed: {err}"), - }, - Err(e) => println!("Failed to connect to {addr}: {e}"), - } - } - }); - - Ok(Response::builder() - .status(StatusCode::OK) - .body(Empty::new()) - .unwrap()) - } else { - Ok(Response::builder() - .status(StatusCode::METHOD_NOT_ALLOWED) - .body(Empty::new()) - .unwrap()) - } -} - -enum ProxyListener { - Tcp(TcpListener), - #[cfg(unix)] - Unix(UnixListener), -} - -impl ProxyListener { - async fn accept(&self) -> io::Result { - match self { - ProxyListener::Tcp(tcp) => tcp.accept().await.map(|(s, _)| ProxyStream::Tcp(s)), - #[cfg(unix)] - ProxyListener::Unix(unix) => unix.accept().await.map(|(s, _)| ProxyStream::Unix(s)), - } - } -} diff --git a/crates/sdk-core/tests/common/mod.rs b/crates/sdk-core/tests/common/mod.rs index ed8dd68f7..dc2acf108 100644 --- a/crates/sdk-core/tests/common/mod.rs +++ b/crates/sdk-core/tests/common/mod.rs @@ -3,7 +3,6 @@ pub(crate) mod activity_functions; pub(crate) mod fake_grpc_server; -pub(crate) mod http_proxy; pub(crate) mod workflows; use anyhow::bail; @@ -23,7 +22,7 @@ use std::{ path::PathBuf, str::FromStr, sync::{ - Arc, + Arc, LazyLock, atomic::{AtomicBool, Ordering}, }, time::{Duration, Instant}, @@ -32,6 +31,7 @@ use temporalio_client::{ Client, ClientOptions, ClientTlsOptions, Connection, ConnectionOptions, GrpcCompression, NamespacedClient, TlsOptions, UntypedWorkflow, UntypedWorkflowHandle, WorkflowExecutionInfo, WorkflowGetResultOptions, WorkflowHandle, WorkflowStartOptions, + envconfig::LoadClientConfigProfileOptions, errors::{WorkflowGetResultError, WorkflowStartError}, grpc::WorkflowService, }; @@ -40,7 +40,7 @@ use temporalio_common::{ data_converters::{DataConverter, RawValue}, protos::{ coresdk::{ - workflow_activation::WorkflowActivation, + workflow_activation::{WorkflowActivation, remove_from_cache::EvictionReason}, workflow_completion::WorkflowActivationCompletion, }, temporal::api::{ @@ -55,15 +55,13 @@ use temporalio_common::{ }; use temporalio_sdk::{ Worker, WorkerOptions, - interceptors::{ - FailOnNondeterminismInterceptor, ReturnWorkflowExitValueInterceptor, WorkerInterceptor, - }, + interceptors::{ReturnWorkflowExitValueInterceptor, WorkerInterceptor}, }; #[cfg(any(feature = "test-utilities", test))] pub(crate) use temporalio_sdk_core::test_help::NAMESPACE; use temporalio_sdk_core::{ - CoreRuntime, RuntimeOptions, Worker as CoreWorker, WorkerConfig, WorkerVersioningStrategy, - init_replay_worker, init_worker, + CoreRuntime, RuntimeOptions, Worker as CoreWorker, WorkerConfig, + WorkerTuner as CoreWorkerTuner, WorkerVersioningStrategy, init_replay_worker, init_worker, replay::{HistoryForReplay, ReplayWorkerInput}, test_help::{MockPollCfg, build_mock_pollers, mock_worker}, }; @@ -77,6 +75,7 @@ pub(crate) const INTEG_SERVER_TARGET_ENV_VAR: &str = "TEMPORAL_SERVICE_ADDRESS"; pub(crate) const INTEG_NAMESPACE_ENV_VAR: &str = "TEMPORAL_NAMESPACE"; pub(crate) const INTEG_USE_TLS_ENV_VAR: &str = "TEMPORAL_USE_TLS"; pub(crate) const INTEG_API_KEY: &str = "TEMPORAL_API_KEY_PATH"; +pub(crate) const TEST_ENV_CONFIG_SERVER_ENV_VAR: &str = "TEMPORAL_TEST_ENV_CONFIG_SERVER"; pub(crate) static SEARCH_ATTR_TXT: &str = "CustomTextField"; pub(crate) static SEARCH_ATTR_INT: &str = "CustomIntField"; /// If set, turn export traces and metrics to the OTel collector at the given URL @@ -91,6 +90,36 @@ pub(crate) const INTEG_CLIENT_IDENTITY: &str = "integ_tester"; pub(crate) const INTEG_CLIENT_NAME: &str = "temporal-core"; pub(crate) const INTEG_CLIENT_VERSION: &str = "0.1.0"; +// Envconfig can read TOML profiles and TLS credentials from files. Load one immutable snapshot so +// concurrently-created clients and workers cannot observe different configuration during a test +// run. +static ENV_CONFIG_CLIENT_CONFIG: LazyLock<(ConnectionOptions, String)> = LazyLock::new(|| { + let (mut connection_options, client_options) = + ClientOptions::load_from_config(LoadClientConfigProfileOptions::default()) + .unwrap_or_else(|err| panic!("Failed to load integration test envconfig: {err}")); + connection_options.identity = INTEG_CLIENT_IDENTITY.to_string(); + (connection_options, client_options.namespace) +}); + +/// Causes test workers to fail immediately when Core evicts a workflow for nondeterminism. +pub(crate) struct FailOnNondeterminismInterceptor {} + +#[async_trait::async_trait(?Send)] +impl WorkerInterceptor for FailOnNondeterminismInterceptor { + async fn on_workflow_activation( + &self, + activation: &WorkflowActivation, + ) -> Result<(), anyhow::Error> { + if matches!( + activation.eviction_reason(), + Some(EvictionReason::Nondeterminism) + ) { + bail!("Workflow is being evicted because of nondeterminism! {activation}"); + } + Ok(()) + } +} + /// Create a worker instance which will use the provided test name to base the task queue and wf id /// upon. Returns the instance. pub(crate) async fn init_core_and_create_wf(test_name: &str) -> CoreWfStarter { @@ -101,7 +130,12 @@ pub(crate) async fn init_core_and_create_wf(test_name: &str) -> CoreWfStarter { } pub(crate) fn integ_namespace() -> String { - env::var(INTEG_NAMESPACE_ENV_VAR).unwrap_or(NAMESPACE.to_string()) + if env::var_os(TEST_ENV_CONFIG_SERVER_ENV_VAR).is_some() { + let (_, namespace) = &*ENV_CONFIG_CLIENT_CONFIG; + namespace.clone() + } else { + env::var(INTEG_NAMESPACE_ENV_VAR).unwrap_or(NAMESPACE.to_string()) + } } pub(crate) fn integ_worker_config(tq: &str) -> WorkerConfig { @@ -123,10 +157,12 @@ pub(crate) fn integ_worker_config(tq: &str) -> WorkerConfig { pub(crate) fn integ_sdk_config(tq: &str) -> WorkerOptions { WorkerOptions::new(tq) .deployment_options( - WorkerDeploymentOptions::new(WorkerDeploymentVersion { - deployment_name: "".to_owned(), - build_id: "test_build_id".to_owned(), - }) + WorkerDeploymentOptions::new( + WorkerDeploymentVersion::builder() + .deployment_name("".to_owned()) + .build_id("test_build_id".to_owned()) + .build(), + ) .build(), ) .build() @@ -319,6 +355,7 @@ pub(crate) struct CoreWfStarter { /// Run when initializing, allows for altering the config used to init the core worker #[allow(clippy::type_complexity)] // It's not tho core_config_mutator: Option>, + core_tuner_override: Option>, core_task_types: Option, } struct InitializedWorker { @@ -444,6 +481,7 @@ impl CoreWfStarter { client_override, min_local_server_version: None, core_config_mutator: None, + core_tuner_override: None, core_task_types: None, } } @@ -460,6 +498,7 @@ impl CoreWfStarter { min_local_server_version: self.min_local_server_version.clone(), initted_worker: Default::default(), core_config_mutator: self.core_config_mutator.clone(), + core_tuner_override: self.core_tuner_override.clone(), core_task_types: self.core_task_types, } } @@ -475,9 +514,10 @@ impl CoreWfStarter { let interceptor_router = TestWorkerInterceptorRouter::default(); let mut sdk_config = self.sdk_config.clone(); sdk_config.worker_interceptor(interceptor_router.clone()); - let sdk = Worker::new_from_core_options(worker, client.options().clone(), sdk_config) - .expect("SDK worker should initialize from core worker and options"); - let mut w = TestWorker::new_with_interceptor_router(sdk, interceptor_router); + let sdk = + Worker::new_from_core_options(worker.clone(), client.options().clone(), sdk_config) + .expect("SDK worker should initialize from core worker and options"); + let mut w = TestWorker::new_with_interceptor_router(sdk, worker, interceptor_router); w.client = Some(client); w @@ -487,6 +527,10 @@ impl CoreWfStarter { self.core_config_mutator = Some(Arc::new(mutator)) } + pub(crate) fn set_core_tuner(&mut self, tuner: Arc) { + self.core_tuner_override = Some(tuner); + } + pub(crate) fn set_core_task_types(&mut self, task_types: WorkerTaskTypes) { self.core_task_types = Some(task_types); } @@ -577,9 +621,9 @@ impl CoreWfStarter { let events = client .get_workflow_handle::(self.get_wf_id()) .fetch_history(Default::default()) + .into_events() .await - .unwrap() - .into_events(); + .unwrap(); History { events } } @@ -640,6 +684,9 @@ impl CoreWfStarter { if let Some(ref ccm) = self.core_config_mutator { ccm(&mut core_config); } + if let Some(tuner) = &self.core_tuner_override { + core_config.tuner = Some(tuner.clone()); + } let worker = init_worker(rt, core_config, connection).expect("Worker inits cleanly"); InitializedWorker { @@ -654,6 +701,7 @@ impl CoreWfStarter { /// Provides conveniences for running integ tests with the SDK (against real server or mocks) pub(crate) struct TestWorker { inner: Worker, + core_worker: Arc, interceptor_router: Option, client: Option, pub started_workflows: Arc>>, @@ -663,9 +711,10 @@ pub(crate) struct TestWorker { } impl TestWorker { /// Create a new test worker - pub(crate) fn new(sdk: Worker) -> Self { + pub(crate) fn new(sdk: Worker, core_worker: Arc) -> Self { Self { inner: sdk, + core_worker, interceptor_router: None, client: None, started_workflows: Arc::new(Mutex::new(vec![])), @@ -675,11 +724,12 @@ impl TestWorker { fn new_with_interceptor_router( sdk: Worker, + core_worker: Arc, interceptor_router: TestWorkerInterceptorRouter, ) -> Self { Self { interceptor_router: Some(interceptor_router), - ..Self::new(sdk) + ..Self::new(sdk, core_worker) } } @@ -749,12 +799,13 @@ impl TestWorker { } let wfid = options.workflow_id.clone(); let handle = c.start_workflow(workflow, input, options).await?; - self.started_workflows.lock().push(WorkflowExecutionInfo { - namespace: c.namespace(), - workflow_id: wfid, - run_id: handle.info().run_id.clone(), - first_execution_run_id: None, - }); + self.started_workflows.lock().push( + WorkflowExecutionInfo::builder() + .namespace(c.namespace()) + .workflow_id(wfid) + .maybe_run_id(handle.info().run_id.clone()) + .build(), + ); Ok(handle) } @@ -763,16 +814,18 @@ impl TestWorker { wf_id: impl Into, run_id: Option, ) { - self.started_workflows.lock().push(WorkflowExecutionInfo { - namespace: self - .client - .as_ref() - .map(|c| c.namespace()) - .unwrap_or(NAMESPACE.to_owned()), - workflow_id: wf_id.into(), - run_id, - first_execution_run_id: None, - }); + self.started_workflows.lock().push( + WorkflowExecutionInfo::builder() + .namespace( + self.client + .as_ref() + .map(|c| c.namespace()) + .unwrap_or(NAMESPACE.to_owned()), + ) + .workflow_id(wf_id.into()) + .maybe_run_id(run_id) + .build(), + ); } /// Runs until all expected workflows have completed and then shuts down the worker @@ -816,7 +869,7 @@ impl TestWorker { } pub(crate) fn core_worker(&self) -> Arc { - self.inner.core_worker() + self.core_worker.clone() } } @@ -847,12 +900,13 @@ impl TestWorkerSubmitterHandle { ) .await?; let run_id = handle.run_id().unwrap().to_string(); - self.started_workflows.lock().push(WorkflowExecutionInfo { - namespace: self.client.namespace(), - workflow_id: wfid, - run_id: Some(run_id.clone()), - first_execution_run_id: None, - }); + self.started_workflows.lock().push( + WorkflowExecutionInfo::builder() + .namespace(self.client.namespace()) + .workflow_id(wfid) + .maybe_run_id(Some(run_id.clone())) + .build(), + ); Ok(run_id) } } @@ -908,6 +962,11 @@ impl TestWorkerCompletionIceptor { } /// Returns the connection options used to connect to the server used for integration tests. pub(crate) fn get_integ_server_options() -> ConnectionOptions { + if env::var_os(TEST_ENV_CONFIG_SERVER_ENV_VAR).is_some() { + let (connection_options, _) = &*ENV_CONFIG_CLIENT_CONFIG; + return connection_options.clone(); + } + let temporal_server_address = env::var(INTEG_SERVER_TARGET_ENV_VAR) .unwrap_or_else(|_| "http://localhost:7233".to_owned()); let url = Url::try_from(&*temporal_server_address).unwrap(); @@ -1033,7 +1092,7 @@ where worker: &mut TestWorker, ) -> Result, anyhow::Error> { let wf_id = self.info().workflow_id.clone(); - let events = self.fetch_history(Default::default()).await?.into_events(); + let events = self.fetch_history(Default::default()).into_events().await?; let with_id = HistoryForReplay::new(events, wf_id); let replay_worker = init_core_replay_preloaded(worker.inner.task_queue(), [with_id]); worker.inner.with_new_core_worker(Arc::new(replay_worker)); @@ -1166,7 +1225,7 @@ pub(crate) fn mock_sdk_cfg_with_options( poll_cfg.using_rust_sdk = true; let mut mock = build_mock_pollers(poll_cfg); mock.worker_cfg(mutator); - let core = mock_worker(mock); + let core = Arc::new(mock_worker(mock)); let interceptor_router = TestWorkerInterceptorRouter::default(); let client_options = ClientOptions::new(core.get_config().namespace.clone()) .data_converter(DataConverter::default()) @@ -1175,9 +1234,9 @@ pub(crate) fn mock_sdk_cfg_with_options( .worker_interceptor(interceptor_router.clone()) .build(); options_mutator(&mut worker_options); - let sdk = Worker::new_from_core_options(Arc::new(core), client_options, worker_options) + let sdk = Worker::new_from_core_options(core.clone(), client_options, worker_options) .expect("mock worker options are valid"); - TestWorker::new_with_interceptor_router(sdk, interceptor_router) + TestWorker::new_with_interceptor_router(sdk, core, interceptor_router) } #[derive(Default)] @@ -1265,6 +1324,10 @@ pub(crate) fn integ_dev_server_config( "--dynamic-config-value".to_owned(), "system.enableCancelActivityWorkerCommand=true".to_owned(), "--dynamic-config-value".to_owned(), + "history.enableWorkflowTaskCompletionPagination=true".to_owned(), + "--dynamic-config-value".to_owned(), + "system.transactionSizeLimit=33554432".to_owned(), + "--dynamic-config-value".to_owned(), "matching.rps=12000".to_owned(), "--search-attribute".to_string(), format!("{SEARCH_ATTR_TXT}=Text"), diff --git a/crates/sdk-core/tests/fixtures/wasm_patch_activation/src/lib.rs b/crates/sdk-core/tests/fixtures/wasm_patch_activation/src/lib.rs index 7341b4dd8..695cb215f 100644 --- a/crates/sdk-core/tests/fixtures/wasm_patch_activation/src/lib.rs +++ b/crates/sdk-core/tests/fixtures/wasm_patch_activation/src/lib.rs @@ -1,12 +1,6 @@ use std::time::Duration; use temporalio_workflow::{ WorkflowContext, WorkflowResult, - component::{StaticWorkflowComponent, instantiate_component_workflow}, - runtime::{ - guest::WorkflowInstance, - host::WorkflowHost, - types::{WorkflowDefinitionDescriptor, WorkflowFailure, WorkflowInit}, - }, workflow, workflow_methods, }; @@ -24,34 +18,4 @@ impl PatchActivationWorkflow { } } -struct WasmPatchActivationWorkflowModule; - -impl StaticWorkflowComponent for WasmPatchActivationWorkflowModule { - fn list_workflows() -> Vec { - vec![ - ::definition(), - ] - } - - fn instantiate_workflow( - workflow_type: &str, - init: WorkflowInit, - host: std::rc::Rc, - ) -> Result, WorkflowFailure> { - match workflow_type { - name if name - == ::name() => - { - instantiate_component_workflow::(init, host) - } - _ => unreachable!("unexpected workflow type '{workflow_type}'"), - } - } -} - -type WasmPatchActivationWorkflowComponentExport = - temporalio_workflow::component::ExportedComponent; - -temporalio_workflow::__temporalio_export_workflow_component!( - WasmPatchActivationWorkflowComponentExport -); +temporalio_workflow::export_workflow_module!([PatchActivationWorkflow]); diff --git a/crates/sdk-core/tests/fixtures/wasm_task_failure/.gitignore b/crates/sdk-core/tests/fixtures/wasm_task_failure/.gitignore new file mode 100644 index 000000000..b83d22266 --- /dev/null +++ b/crates/sdk-core/tests/fixtures/wasm_task_failure/.gitignore @@ -0,0 +1 @@ +/target/ diff --git a/crates/sdk-core/tests/fixtures/wasm_task_failure/Cargo.toml b/crates/sdk-core/tests/fixtures/wasm_task_failure/Cargo.toml new file mode 100644 index 000000000..aed713fe1 --- /dev/null +++ b/crates/sdk-core/tests/fixtures/wasm_task_failure/Cargo.toml @@ -0,0 +1,16 @@ +[package] +name = "temporal-wasm-task-failure-workflow" +version = "0.1.0" +edition = "2024" +publish = false + +[dependencies] +temporalio-workflow = { path = "../../../../workflow" } + +[lib] +crate-type = ["cdylib"] + +[package.metadata.component] +package = "temporal:task-failure-workflow" + +[workspace] diff --git a/crates/sdk-core/tests/fixtures/wasm_task_failure/src/lib.rs b/crates/sdk-core/tests/fixtures/wasm_task_failure/src/lib.rs new file mode 100644 index 000000000..c99dc09d3 --- /dev/null +++ b/crates/sdk-core/tests/fixtures/wasm_task_failure/src/lib.rs @@ -0,0 +1,102 @@ +use temporalio_workflow::{ + __private::{ + macros::{ExportedComponent, StaticWorkflowComponent}, + sdk::{ + ActivationJobResult, ActivationResult, MAIN_ROUTINE_ID, MainRoutineCompletion, + RoutineCompletion, RoutinePollResult, TaskFailure, WorkflowActivation, + WorkflowFailure, WorkflowHost, WorkflowInit, WorkflowInstance, + }, + }, + common::protos::temporal::api::{ + enums::v1::WorkflowTaskFailedCause, + failure::v1::{ApplicationFailureInfo, Failure, failure::FailureInfo}, + }, + workflows::WorkflowDefinitionDescriptor, +}; + +struct WasmTaskFailureWorkflow; + +impl WorkflowInstance for WasmTaskFailureWorkflow { + fn activate( + &mut self, + activation: WorkflowActivation, + _waker: &std::task::Waker, + ) -> Result { + Ok(ActivationResult { + job_results: activation + .jobs + .iter() + .map(|_| ActivationJobResult::None) + .collect(), + }) + } + + fn poll_routine( + &mut self, + routine_id: u64, + _waker: &std::task::Waker, + ) -> Result { + if routine_id != MAIN_ROUTINE_ID { + return Err(Box::new(Failure { + message: format!("unexpected routine id {routine_id}"), + ..Default::default() + })); + } + + Ok(RoutinePollResult { + completion: Some(RoutineCompletion::Main(MainRoutineCompletion::TaskFailed( + TaskFailure { + failure: Box::new(Failure { + message: "structured wasm workflow task failure".to_string(), + failure_info: Some(FailureInfo::ApplicationFailureInfo( + ApplicationFailureInfo { + r#type: "WasmTaskFailure".to_string(), + non_retryable: true, + ..Default::default() + }, + )), + ..Default::default() + }), + force_cause: Some(WorkflowTaskFailedCause::NonDeterministicError as u32), + }, + ))), + made_progress: true, + pending_state: None, + }) + } +} + +struct WasmTaskFailureWorkflowModule; + +impl StaticWorkflowComponent for WasmTaskFailureWorkflowModule { + fn list_workflows() -> Vec { + vec![WorkflowDefinitionDescriptor { + workflow_type: "WasmTaskFailureWorkflow".to_string(), + has_init: false, + init_takes_input: false, + signals: vec![], + queries: vec![], + updates: vec![], + }] + } + + fn instantiate_workflow( + workflow_type: &str, + _init: WorkflowInit, + _host: std::rc::Rc, + ) -> Result, WorkflowFailure> { + match workflow_type { + "WasmTaskFailureWorkflow" => Ok(Box::new(WasmTaskFailureWorkflow)), + _ => Err(Box::new(Failure { + message: format!("No workflow named '{workflow_type}' exported by this component"), + ..Default::default() + })), + } + } +} + +type WasmTaskFailureWorkflowComponentExport = ExportedComponent; + +temporalio_workflow::__temporalio_export_workflow_component!( + WasmTaskFailureWorkflowComponentExport +); diff --git a/crates/sdk-core/tests/fsm_procmacro.rs b/crates/sdk-core/tests/fsm_procmacro.rs index 0ce83e788..ec126abbc 100644 --- a/crates/sdk-core/tests/fsm_procmacro.rs +++ b/crates/sdk-core/tests/fsm_procmacro.rs @@ -1,6 +1,5 @@ #[test] fn fsm_procmacro_build_tests() { let t = trybuild::TestCases::new(); - t.pass("tests/fsm_trybuild/*_pass.rs"); t.compile_fail("tests/fsm_trybuild/*_fail.rs"); } diff --git a/crates/sdk-core/tests/fsm_trybuild/dynamic_dest_pass.rs b/crates/sdk-core/tests/fsm_trybuild/dynamic_dest_pass.rs deleted file mode 100644 index 60e8d2142..000000000 --- a/crates/sdk-core/tests/fsm_trybuild/dynamic_dest_pass.rs +++ /dev/null @@ -1,39 +0,0 @@ -#![allow(dead_code)] - -use std::convert::Infallible; -use temporalio_common::fsm_trait::TransitionResult; -use temporalio_macros::fsm; - -fsm! { - name SimpleMachine; command SimpleMachineCommand; error Infallible; - - One --(A(String), foo)--> Two; - One --(A(String), foo)--> Three; - - Two --(B(String), bar)--> One; - Two --(B(String), bar)--> Two; - Two --(B(String), bar)--> Three; -} - -#[derive(Default, Clone)] -pub struct One {} -impl One { - fn foo(self, _: String) -> SimpleMachineTransition { - TransitionResult::ok(vec![], Two {}.into()) - } -} - -#[derive(Default, Clone)] -pub struct Two {} -impl Two { - fn bar(self, _: String) -> SimpleMachineTransition { - TransitionResult::ok(vec![], Three {}.into()) - } -} - -#[derive(Default, Clone)] -pub struct Three {} - -pub enum SimpleMachineCommand {} - -fn main() {} diff --git a/crates/sdk-core/tests/fsm_trybuild/handler_arg_pass.rs b/crates/sdk-core/tests/fsm_trybuild/handler_arg_pass.rs deleted file mode 100644 index 16d4ed9c7..000000000 --- a/crates/sdk-core/tests/fsm_trybuild/handler_arg_pass.rs +++ /dev/null @@ -1,30 +0,0 @@ -use std::convert::Infallible; -use temporalio_common::fsm_trait::TransitionResult; -use temporalio_macros::fsm; - -fsm! { - name Simple; command SimpleCommand; error Infallible; - - One --(A(String), on_a)--> Two -} - -#[derive(Default, Clone)] -pub struct One {} -impl One { - fn on_a(self, _: String) -> SimpleTransition { - SimpleTransition::ok(vec![], Two {}) - } -} - -#[derive(Default, Clone)] -pub struct Two {} - -pub enum SimpleCommand {} - -fn main() { - // state enum exists with both states - let _ = SimpleState::One(One {}); - let _ = SimpleState::Two(Two {}); - // Avoid dead code warning - let _ = SimpleEvents::A("yo".to_owned()); -} diff --git a/crates/sdk-core/tests/fsm_trybuild/handler_pass.rs b/crates/sdk-core/tests/fsm_trybuild/handler_pass.rs deleted file mode 100644 index 30d73bf8d..000000000 --- a/crates/sdk-core/tests/fsm_trybuild/handler_pass.rs +++ /dev/null @@ -1,29 +0,0 @@ -use std::convert::Infallible; -use temporalio_common::fsm_trait::TransitionResult; -use temporalio_macros::fsm; - -fsm! { - name Simple; command SimpleCommand; error Infallible; - - One --(A, on_a)--> Two -} - -#[derive(Default, Clone)] -pub struct One {} -impl One { - fn on_a(self) -> SimpleTransition { - SimpleTransition::ok(vec![], Two {}) - } -} - -#[derive(Default, Clone)] -pub struct Two {} - -pub enum SimpleCommand {} - -fn main() { - // state enum exists with both states - let _ = SimpleState::One(One {}); - let _ = SimpleState::Two(Two {}); - let _ = SimpleEvents::A; -} diff --git a/crates/sdk-core/tests/fsm_trybuild/medium_complex_pass.rs b/crates/sdk-core/tests/fsm_trybuild/medium_complex_pass.rs deleted file mode 100644 index 6246ec17e..000000000 --- a/crates/sdk-core/tests/fsm_trybuild/medium_complex_pass.rs +++ /dev/null @@ -1,44 +0,0 @@ -#![allow(dead_code)] - -use std::convert::Infallible; -use temporalio_common::fsm_trait::TransitionResult; -use temporalio_macros::fsm; - -fsm! { - name SimpleMachine; command SimpleMachineCommand; error Infallible; - - One --(A(String), foo)--> Two; - One --(B)--> Two; - Two --(B)--> One; - Two --(C, baz)--> One -} - -#[derive(Default, Clone)] -pub struct One {} -impl One { - fn foo(self, _: String) -> SimpleMachineTransition { - TransitionResult::default() - } -} -impl From for One { - fn from(_: Two) -> Self { - One {} - } -} - -#[derive(Default, Clone)] -pub struct Two {} -impl Two { - fn baz(self) -> SimpleMachineTransition { - TransitionResult::default() - } -} -impl From for Two { - fn from(_: One) -> Self { - Two {} - } -} - -pub enum SimpleMachineCommand {} - -fn main() {} diff --git a/crates/sdk-core/tests/fsm_trybuild/simple_pass.rs b/crates/sdk-core/tests/fsm_trybuild/simple_pass.rs deleted file mode 100644 index f090c390d..000000000 --- a/crates/sdk-core/tests/fsm_trybuild/simple_pass.rs +++ /dev/null @@ -1,30 +0,0 @@ -use std::convert::Infallible; -use temporalio_common::fsm_trait::TransitionResult; -use temporalio_macros::fsm; - -fsm! { - name SimpleMachine; command SimpleMachineCommand; error Infallible; - - One --(A)--> Two -} - -#[derive(Default, Clone)] -pub struct One {} - -#[derive(Default, Clone)] -pub struct Two {} -impl From for Two { - fn from(_: One) -> Self { - Two {} - } -} - -pub enum SimpleMachineCommand {} - -fn main() { - // state enum exists with both states - let _ = SimpleMachineState::One(One {}); - let _ = SimpleMachineState::Two(Two {}); - // Event enum exists - let _ = SimpleMachineEvents::A; -} diff --git a/crates/sdk-core/tests/heavy_tests.rs b/crates/sdk-core/tests/heavy_tests.rs index 7f3821e6c..3d3e5c49f 100644 --- a/crates/sdk-core/tests/heavy_tests.rs +++ b/crates/sdk-core/tests/heavy_tests.rs @@ -5,10 +5,9 @@ pub(crate) mod common; #[path = "heavy_tests/fuzzy_workflow.rs"] mod fuzzy_workflow; -use crate::common::get_integ_runtime_options; use common::{ - CoreWfStarter, activity_functions::StdActivities, init_integ_telem, prom_metrics, rand_6_chars, - workflows::LaProblemWorkflow, + CoreWfStarter, activity_functions::StdActivities, get_integ_runtime_options, init_integ_telem, + prom_metrics, rand_6_chars, workflows::LaProblemWorkflow, }; use futures_util::{ StreamExt, @@ -24,25 +23,22 @@ use std::{ }; use temporalio_client::{ NamespacedClient, UntypedSignal, UntypedWorkflow, WorkflowExecutionInfo, - WorkflowGetResultOptions, WorkflowSignalOptions, WorkflowStartOptions, -}; -use temporalio_common::{ - data_converters::RawValue, protos::temporal::api::enums::v1::WorkflowIdConflictPolicy, + WorkflowGetResultOptions, WorkflowIdConflictPolicy, WorkflowIdReusePolicy, + WorkflowSignalOptions, WorkflowStartOptions, }; +use temporalio_common::data_converters::RawValue; use temporalio_macros::{activities, workflow, workflow_methods}; -use temporalio_common::protos::{ - coresdk::workflow_commands::ActivityCancellationType, - temporal::api::enums::v1::WorkflowIdReusePolicy, +use temporalio_common::{ + ActivityCloseTimeouts, protos::coresdk::workflow_commands::ActivityCancellationType, }; use temporalio_sdk::{ - ActivityCloseTimeouts, ActivityOptions, SyncWorkflowContext, WorkflowContext, WorkflowResult, + ActivityOptions, SyncWorkflowContext, WorkflowContext, WorkflowResult, activities::{ActivityContext, ActivityError}, + runtime::{AutoscalingOptions, PollerBehavior}, workflows, }; -use temporalio_sdk_core::{ - CoreRuntime, PollerBehavior, ResourceBasedTuner, ResourceSlotOptions, TunerHolder, -}; +use temporalio_sdk_core::{CoreRuntime, ResourceBasedTuner, ResourceSlotOptions, TunerHolder}; #[workflow] #[derive(Clone, Default)] @@ -57,10 +53,12 @@ impl ActivityLoadWf { .execute_activity( StdActivities::echo, input_str.clone(), - ActivityOptions::with_close_timeouts(ActivityCloseTimeouts::Both { - start_to_close: Duration::from_secs(8), - schedule_to_close: Duration::from_secs(8), - }) + ActivityOptions::with_close_timeouts( + ActivityCloseTimeouts::ScheduleAndStartToClose { + start_to_close: Duration::from_secs(8), + schedule_to_close: Duration::from_secs(8), + }, + ) .activity_id("act-1".to_string()) .task_queue(tq) .schedule_to_start_timeout(Duration::from_secs(8)) @@ -81,8 +79,12 @@ async fn activity_load() { let mut starter = CoreWfStarter::new("activity_load"); starter.sdk_config.max_cached_workflows = CONCURRENCY; starter.sdk_config.activity_task_poller_behavior = Some(PollerBehavior::SimpleMaximum(10)); - starter.sdk_config.tuner = - Arc::new(TunerHolder::fixed_size(CONCURRENCY, CONCURRENCY, 100, 100)); + starter.set_core_tuner(Arc::new(TunerHolder::fixed_size( + CONCURRENCY, + CONCURRENCY, + 100, + 100, + ))); starter.sdk_config.register_activities(StdActivities); starter .sdk_config @@ -174,7 +176,7 @@ async fn chunky_activities_resource_based() { Duration::from_millis(0), )) .with_activity_slots_options(ResourceSlotOptions::new(5, 1000, Duration::from_millis(50))); - starter.sdk_config.tuner = Arc::new(tuner); + starter.set_core_tuner(Arc::new(tuner)); starter.sdk_config.register_activities(ChunkyActivities); starter @@ -251,7 +253,7 @@ async fn workflow_load() { let mut starter = CoreWfStarter::new_with_runtime("workflow_load", rt); starter.sdk_config.max_cached_workflows = 200; starter.sdk_config.activity_task_poller_behavior = Some(PollerBehavior::SimpleMaximum(10)); - starter.sdk_config.tuner = Arc::new(TunerHolder::fixed_size(5, 100, 100, 100)); + starter.set_core_tuner(Arc::new(TunerHolder::fixed_size(5, 100, 100, 100))); starter.sdk_config.register_activities(StdActivities); let task_queue = starter.get_task_queue().to_owned(); starter @@ -307,7 +309,7 @@ async fn evict_while_la_running_no_interference() { // Though it doesn't make sense to set wft higher than cached workflows, leaving this commented // introduces more instability that can be useful in the test. // starter.max_wft(20); - starter.sdk_config.tuner = Arc::new(TunerHolder::fixed_size(100, 10, 20, 1)); + starter.set_core_tuner(Arc::new(TunerHolder::fixed_size(100, 10, 20, 1))); starter.sdk_config.register_activities(StdActivities); starter .sdk_config @@ -334,20 +336,19 @@ async fn evict_while_la_running_no_interference() { subfs.push(async move { tokio::time::sleep(Duration::from_secs(1)).await; cw.request_workflow_eviction(&run_id); - WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: wf_id, - run_id: Some(run_id), - first_execution_run_id: None, - } - .bind_untyped(client) - .signal( - UntypedSignal::new("whaatever"), - RawValue::empty(), - WorkflowSignalOptions::default(), - ) - .await - .unwrap(); + WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(wf_id) + .maybe_run_id(Some(run_id)) + .build() + .bind_untyped(client) + .signal( + UntypedSignal::new("whaatever"), + RawValue::empty(), + WorkflowSignalOptions::default(), + ) + .await + .unwrap(); }); } let runf = async { @@ -398,13 +399,12 @@ async fn can_paginate_long_history() { let run_id = handle.run_id().unwrap().to_owned(); let client = starter.get_core_client().await; tokio::spawn(async move { - let handle = WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: wf_name.into(), - run_id: Some(run_id), - first_execution_run_id: None, - } - .bind_untyped(client); + let handle = WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(wf_name) + .maybe_run_id(Some(run_id)) + .build() + .bind_untyped(client); loop { for _ in 0..10 { handle @@ -465,17 +465,21 @@ async fn poller_autoscaling_basic_loadtest() { let wf_name = "poller_load"; let mut starter = CoreWfStarter::new("poller_load"); starter.sdk_config.max_cached_workflows = 5000; - starter.sdk_config.tuner = Arc::new(TunerHolder::fixed_size(1000, 1000, 100, 1)); - starter.sdk_config.workflow_task_poller_behavior = Some(PollerBehavior::Autoscaling { - minimum: 1, - maximum: 200, - initial: 5, - }); - starter.sdk_config.activity_task_poller_behavior = Some(PollerBehavior::Autoscaling { - minimum: 1, - maximum: 200, - initial: 5, - }); + starter.set_core_tuner(Arc::new(TunerHolder::fixed_size(1000, 1000, 100, 1))); + starter.sdk_config.workflow_task_poller_behavior = Some(PollerBehavior::Autoscaling( + AutoscalingOptions::builder() + .minimum(1) + .maximum(200) + .initial(5) + .build(), + )); + starter.sdk_config.activity_task_poller_behavior = Some(PollerBehavior::Autoscaling( + AutoscalingOptions::builder() + .minimum(1) + .maximum(200) + .initial(5) + .build(), + )); starter.sdk_config.register_activities(JitteryActivities); starter diff --git a/crates/sdk-core/tests/heavy_tests/fuzzy_workflow.rs b/crates/sdk-core/tests/heavy_tests/fuzzy_workflow.rs index ade295148..41fd2598f 100644 --- a/crates/sdk-core/tests/heavy_tests/fuzzy_workflow.rs +++ b/crates/sdk-core/tests/heavy_tests/fuzzy_workflow.rs @@ -79,7 +79,7 @@ async fn fuzzy_workflow() { let wf_name = "fuzzy_wf"; let mut starter = CoreWfStarter::new("fuzzy_workflow"); starter.sdk_config.max_cached_workflows = 25; - starter.sdk_config.tuner = Arc::new(TunerHolder::fixed_size(25, 25, 100, 100)); + starter.set_core_tuner(Arc::new(TunerHolder::fixed_size(25, 25, 100, 100))); starter.sdk_config.register_activities(StdActivities); starter.sdk_config.register_workflow::().unwrap(); let mut worker = starter.worker().await; diff --git a/crates/sdk-core/tests/integ_tests/async_activity_client_tests.rs b/crates/sdk-core/tests/integ_tests/async_activity_client_tests.rs index b76790a73..55337ce4d 100644 --- a/crates/sdk-core/tests/integ_tests/async_activity_client_tests.rs +++ b/crates/sdk-core/tests/integ_tests/async_activity_client_tests.rs @@ -44,8 +44,8 @@ async fn async_activity_completions( #[derive(Clone)] struct SharedActivityInfo { task_token: Vec, - workflow_id: String, - run_id: String, + workflow_id: Option, + workflow_run_id: Option, activity_id: String, } @@ -76,11 +76,10 @@ async fn async_activity_completions( } let activity_info = ctx.info(); - let wf_exec = activity_info.workflow_execution.as_ref().unwrap(); let info = SharedActivityInfo { task_token: activity_info.task_token.clone(), - workflow_id: wf_exec.workflow_id().to_owned(), - run_id: wf_exec.run_id().to_owned(), + workflow_id: activity_info.workflow_id.clone(), + workflow_run_id: activity_info.workflow_run_id.clone(), activity_id: activity_info.activity_id.clone(), }; let _ = self.info_tx.send(info).await; @@ -161,10 +160,10 @@ async fn async_activity_completions( let info = info_rx.recv().await.expect("should receive activity info"); eprintln!( - "DEBUG: Received activity info - task_token_len={}, workflow_id={}, run_id={}, activity_id={}", + "DEBUG: Received activity info - task_token_len={}, workflow_id={:?}, run_id={:?}, activity_id={}", info.task_token.len(), info.workflow_id, - info.run_id, + info.workflow_run_id, info.activity_id ); @@ -175,7 +174,11 @@ async fn async_activity_completions( } IdentifierType::ById => { eprintln!("DEBUG: Using ById identifier"); - ActivityIdentifier::by_id(info.workflow_id, info.run_id, info.activity_id) + ActivityIdentifier::by_id_workflow( + info.workflow_id.unwrap(), + info.workflow_run_id.unwrap(), + info.activity_id, + ) } }; diff --git a/crates/sdk-core/tests/integ_tests/client_tests.rs b/crates/sdk-core/tests/integ_tests/client_tests.rs index d893b5a2f..776e2253e 100644 --- a/crates/sdk-core/tests/integ_tests/client_tests.rs +++ b/crates/sdk-core/tests/integ_tests/client_tests.rs @@ -1,8 +1,7 @@ use crate::common::{ CoreWfStarter, NAMESPACE, fake_grpc_server::{FakeServer, GenericService, fake_server}, - get_integ_server_options, - http_proxy::HttpProxy, + get_integ_server_options, integ_namespace, }; use assert_matches::assert_matches; use futures_util::{FutureExt, stream}; @@ -22,25 +21,24 @@ use std::{ time::Duration, }; use temporalio_client::{ - Connection, GrpcCompression, RETRYABLE_ERROR_CODES, RetryOptions, UntypedWorkflow, - errors::ClientConnectError, grpc::WorkflowService, proxy::HttpConnectProxyOptions, + Connection, GrpcCompression, RetryOptions, UntypedWorkflow, errors::ClientConnectError, + grpc::WorkflowService, }; use temporalio_common::protos::temporal::api::{ cloud::cloudservice::v1::GetNamespaceRequest, workflowservice::v1::{ DescribeNamespaceRequest, GetSystemInfoResponse, GetWorkflowExecutionHistoryRequest, - ListNamespacesRequest, RespondActivityTaskCanceledResponse, SignalWorkflowExecutionRequest, - SignalWorkflowExecutionResponse, get_system_info_response, + ListNamespacesRequest, SignalWorkflowExecutionRequest, SignalWorkflowExecutionResponse, + get_system_info_response, }, }; -#[cfg(unix)] -use tokio::net::UnixListener; use tokio::{net::TcpListener, sync::oneshot}; use tonic::{ Code, IntoRequest, Request, Status, body::Body, codegen::http::Response, transport::Server, }; use tracing::info; +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::RequiresOssOnlyApis)] #[tokio::test] async fn can_use_retry_client() { // Not terribly interesting by itself but can be useful for manually inspecting metrics etc @@ -64,7 +62,7 @@ async fn can_use_retry_raw_client() { connection .describe_namespace( DescribeNamespaceRequest { - namespace: NAMESPACE.to_string(), + namespace: integ_namespace(), ..Default::default() } .into_request(), @@ -158,6 +156,7 @@ fn compression_test_options( opts } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn gzip_get_system_info_failure_reconnects_without_compression() { let (fs, records) = @@ -187,6 +186,7 @@ async fn gzip_get_system_info_failure_reconnects_without_compression() { fs.shutdown().await; } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn compression_none_does_not_retry_compression_fallback() { let (fs, records) = @@ -205,6 +205,7 @@ async fn compression_none_does_not_retry_compression_fallback() { fs.shutdown().await; } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn generic_gzip_unimplemented_does_not_reconnect_without_compression() { let (fs, records) = @@ -231,6 +232,7 @@ async fn generic_gzip_unimplemented_does_not_reconnect_without_compression() { fs.shutdown().await; } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn unknown_method_unimplemented_does_not_trigger_compression_reconnect() { let (fs, records) = compression_test_server(CompressionTestBehavior::UnknownMethod).await; @@ -286,6 +288,7 @@ async fn per_call_timeout_respected_one_call() { ); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn timeouts_respected_one_call_fake_server() { let mut fs = fake_server(|_| async { Response::new(Body::empty()) }.boxed()).await; @@ -343,6 +346,7 @@ async fn timeouts_respected_one_call_fake_server() { fs.shutdown().await; } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn non_retryable_errors() { for code in [ @@ -383,47 +387,7 @@ async fn non_retryable_errors() { } } -#[tokio::test] -async fn retryable_errors() { - // Take out retry exhausted since it gets a special policy which would make this take ages - for code in RETRYABLE_ERROR_CODES - .iter() - .copied() - .filter(|p| p != &Code::ResourceExhausted) - { - let count = Arc::new(AtomicUsize::new(0)); - let mut fs = fake_server(move |_| { - let prev = count.fetch_add(1, Ordering::Relaxed); - let r = if prev < 3 { - Status::new(code, "bla").into_http() - } else { - make_ok_response(RespondActivityTaskCanceledResponse::default()) - }; - async { r }.boxed() - }) - .await; - - let mut opts = get_integ_server_options(); - opts.target = format!("http://localhost:{}", fs.addr.port()) - .parse::() - .unwrap(); - opts.set_skip_get_system_info(true); - let connection = Connection::connect(opts).await.unwrap(); - let client_opts = temporalio_client::ClientOptions::new("ns").build(); - let client = temporalio_client::Client::new(connection, client_opts).unwrap(); - - let result = client.count_workflows("whatever", Default::default()).await; - - // Expecting successful response after retries - assert!(result.is_ok(), "{:?}", result); - let mut all_calls = vec![]; - fs.header_rx.recv_many(&mut all_calls, 9999).await; - // Should be 4 attempts - assert_eq!(all_calls.len(), 4); - fs.shutdown().await; - } -} - +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn namespace_header_attached_to_relevant_calls() { let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>(); @@ -463,6 +427,7 @@ async fn namespace_header_attached_to_relevant_calls() { let _ = client .get_workflow_handle::("hi") .fetch_history(Default::default()) + .into_events() .await; let val = header_rx.recv().await.unwrap(); assert_eq!(namespace, val); @@ -496,6 +461,10 @@ async fn grpc_compression() { crate::shared_tests::grpc_compression().await } +#[temporalio_macros::cloud_test_exclusion( + crate::CloudTestExclusionReason::RequiresCloudProvisioning, + "Requires separate Cloud Operations API credentials and a preconfigured namespace." +)] #[tokio::test] async fn cloud_ops_test() { let api_key = match env::var("TEMPORAL_CLIENT_CLOUD_API_KEY") { @@ -534,94 +503,7 @@ async fn cloud_ops_test() { assert_eq!(res.into_inner().namespace.unwrap().namespace, namespace); } -#[tokio::test] -async fn http_proxy() { - // Create server - let call_count = Arc::new(AtomicUsize::new(0)); - let call_count_cloned = call_count.clone(); - let server = fake_server(move |_| { - call_count_cloned.fetch_add(1, Ordering::SeqCst); - async { Response::new(Body::empty()) }.boxed() - }) - .await; - - // Create HTTP TCP proxy - let tcp_proxy_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let tcp_proxy_addr = tcp_proxy_listener.local_addr().unwrap(); - let tcp_proxy = HttpProxy::spawn_tcp(tcp_proxy_listener); - - // General client options - let mut opts = get_integ_server_options(); - opts.retry_options = RetryOptions::no_retries(); - opts.set_skip_get_system_info(true); - - // Connect client with no proxy and make call and confirm reached - opts.target = format!("http://[::1]:{}", server.addr.port()) - .parse() - .unwrap(); - let connection = Connection::connect(opts.clone()).await.unwrap(); - let client_opts = temporalio_client::ClientOptions::new("my-namespace").build(); - let client = temporalio_client::Client::new(connection, client_opts).unwrap(); - let _ = WorkflowService::list_namespaces( - &mut client.clone(), - ListNamespacesRequest::default().into_request(), - ) - .await; - assert!(call_count.load(Ordering::SeqCst) == 1); - assert!(tcp_proxy.hit_count() == 0); - - // Connect client to proxy and make call and confirm reached - opts.http_connect_proxy = - Some(HttpConnectProxyOptions::new(tcp_proxy_addr.to_string()).build()); - opts.dns_load_balancing = None; - let connection = Connection::connect(opts.clone()).await.unwrap(); - let client_opts = temporalio_client::ClientOptions::new("my-namespace").build(); - let proxied_client = temporalio_client::Client::new(connection, client_opts).unwrap(); - let _ = WorkflowService::list_namespaces( - &mut proxied_client.clone(), - ListNamespacesRequest::default().into_request(), - ) - .await; - assert!(call_count.load(Ordering::SeqCst) == 2); - assert!(tcp_proxy.hit_count() == 1); - - // Test Unix socket too only in Unix environments - #[cfg(unix)] - { - // Create temp socket path - let mut sock_path = std::env::temp_dir(); - sock_path.push(format!("http-proxy-test-{}.sock", std::process::id())); - // Remove if there just in case - let _ = std::fs::remove_file(&sock_path); - - // Create unix-socket-based proxy - let unix_proxy = HttpProxy::spawn_unix(UnixListener::bind(&sock_path).unwrap()); - - // Connect client to proxy and make call and confirm reached - opts.http_connect_proxy = Some( - HttpConnectProxyOptions::new(format!("unix:{}", sock_path.to_str().unwrap())).build(), - ); - opts.dns_load_balancing = None; - let connection = Connection::connect(opts.clone()).await.unwrap(); - let client_opts = temporalio_client::ClientOptions::new("my-namespace").build(); - let proxied_client = temporalio_client::Client::new(connection, client_opts).unwrap(); - let _ = WorkflowService::list_namespaces( - &mut proxied_client.clone(), - ListNamespacesRequest::default().into_request(), - ) - .await; - assert!(call_count.load(Ordering::SeqCst) == 3); - assert!(unix_proxy.hit_count() == 1); - - // Shutdown unix proxy - unix_proxy.shutdown(); - } - - // Shutdown server and proxy - server.shutdown().await; - tcp_proxy.shutdown(); -} - +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn update_get_result_retries_on_empty_outcome() { use temporalio_common::protos::temporal::api::{ diff --git a/crates/sdk-core/tests/integ_tests/data_converter_tests.rs b/crates/sdk-core/tests/integ_tests/data_converter_tests.rs index 7bf9d1660..8e045ace9 100644 --- a/crates/sdk-core/tests/integ_tests/data_converter_tests.rs +++ b/crates/sdk-core/tests/integ_tests/data_converter_tests.rs @@ -121,7 +121,7 @@ impl FailurePayloadActivities { "codec-heartbeat-details".to_string(), ))) .await?; - tokio::time::sleep(Duration::from_secs(2)).await; + ctx.cancelled().await; Ok(()) } } @@ -159,7 +159,7 @@ impl FailureConverter for FailingFailureConverter { payload_converter: &PayloadConverter, context: &SerializationContextData, ) -> Result { - DefaultFailureConverter.to_error(failure, payload_converter, context) + DefaultFailureConverter::default().to_error(failure, payload_converter, context) } } @@ -441,9 +441,12 @@ async fn custom_failure_converter_fallback_applied_to_activity_panic_failures() worker.run_until_done().await.unwrap(); handle.get_result(Default::default()).await.unwrap(); - let history = handle.fetch_history(Default::default()).await.unwrap(); - let activity_failure = history + let history = handle + .fetch_history(Default::default()) .into_events() + .await + .unwrap(); + let activity_failure = history .into_iter() .find_map(|event| match event.attributes { Some(Attributes::ActivityTaskFailedEventAttributes(attrs)) => attrs.failure, @@ -717,9 +720,9 @@ async fn multi_args_serializes_as_multiple_payloads() { let events = client .get_workflow_handle::(wf_name) .fetch_history(Default::default()) + .into_events() .await - .unwrap() - .into_events(); + .unwrap(); let workflow_started_event = events .iter() @@ -756,6 +759,60 @@ async fn multi_args_serializes_as_multiple_payloads() { assert_eq!(second_payload_data, 42); } +#[workflow] +#[derive(Default)] +struct BinaryNullWorkflow; + +#[workflow_methods] +impl BinaryNullWorkflow { + #[run] + async fn run(_ctx: &mut WorkflowContext, input: Option) -> WorkflowResult<()> { + assert_eq!(input, None); + Ok(()) + } +} + +#[tokio::test] +async fn option_none_workflow_input_is_recorded_as_binary_null() { + let wf_name = BinaryNullWorkflow::name(); + let mut starter = CoreWfStarter::new(wf_name); + starter + .sdk_config + .register_workflow::() + .unwrap(); + let mut worker = starter.worker().await; + let handle = worker + .submit_workflow( + BinaryNullWorkflow::run, + Option::::None, + WorkflowStartOptions::new(starter.get_task_queue(), wf_name).build(), + ) + .await + .unwrap(); + worker.run_until_done().await.unwrap(); + + let events = handle + .fetch_history(Default::default()) + .into_events() + .await + .unwrap(); + let input = events + .iter() + .find_map(|event| match event.attributes.as_ref() { + Some(Attributes::WorkflowExecutionStartedEventAttributes(attributes)) => { + attributes.input.as_ref() + } + _ => None, + }) + .unwrap(); + assert_eq!(input.payloads.len(), 1); + assert_eq!( + input.payloads[0].metadata.get("encoding").unwrap(), + b"binary/null" + ); + assert!(input.payloads[0].data.is_empty()); +} + /// A codec that XORs payload data with a key and tracks encode/decode operations. struct XorCodec { key: u8, @@ -877,10 +934,10 @@ impl PayloadCodec for FailOnceCodec { (self.failure_point, context), ( CodecFailurePoint::WorkflowEncode, - SerializationContextData::Workflow + SerializationContextData::Workflow(_) ) | ( CodecFailurePoint::ActivityEncode, - SerializationContextData::Activity + SerializationContextData::Activity(_) ) ); let marker = if matches!(self.failure_point, CodecFailurePoint::WorkflowEncode) { @@ -920,10 +977,10 @@ impl PayloadCodec for FailOnceCodec { (self.failure_point, context), ( CodecFailurePoint::WorkflowDecode, - SerializationContextData::Workflow + SerializationContextData::Workflow(_) ) | ( CodecFailurePoint::ActivityDecode, - SerializationContextData::Activity + SerializationContextData::Activity(_) ) ); let matches_payload = payloads.iter().any(|payload| { @@ -961,7 +1018,7 @@ async fn codec_errors_fail_tasks_and_retry(#[case] failure_point: CodecFailurePo let connection = get_integ_connection(None).await; let data_converter = DataConverter::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), codec.clone(), ); let client_opts = ClientOptions::new(integ_namespace()) @@ -1009,7 +1066,7 @@ async fn codec_encodes_and_decodes_payloads() { let connection = get_integ_connection(None).await; let data_converter = DataConverter::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), codec.clone(), ); let client_opts = ClientOptions::new(integ_namespace()) @@ -1067,7 +1124,7 @@ async fn describe_decodes_workflow_payload_fields() { let connection = get_integ_connection(None).await; let data_converter = DataConverter::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), codec.clone(), ); let client_opts = ClientOptions::new(integ_namespace()) @@ -1145,7 +1202,7 @@ async fn describe_decodes_user_metadata_with_ungated_xor_codec() { let connection = get_integ_connection(None).await; let data_converter = DataConverter::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), codec.clone(), ); let client_opts = ClientOptions::new(integ_namespace()) @@ -1209,7 +1266,7 @@ async fn codec_roundtrips_activity_cancellation_details() { let connection = get_integ_connection(None).await; let data_converter = DataConverter::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), codec.clone(), ); let client_opts = ClientOptions::new(integ_namespace()) @@ -1258,7 +1315,7 @@ async fn codec_roundtrips_activity_heartbeat_timeout_details() { let connection = get_integ_connection(None).await; let data_converter = DataConverter::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), codec.clone(), ); let client_opts = ClientOptions::new(integ_namespace()) diff --git a/crates/sdk-core/tests/integ_tests/ephemeral_server_tests.rs b/crates/sdk-core/tests/integ_tests/ephemeral_server_tests.rs index 8b5db1fef..56deca926 100644 --- a/crates/sdk-core/tests/integ_tests/ephemeral_server_tests.rs +++ b/crates/sdk-core/tests/integ_tests/ephemeral_server_tests.rs @@ -1,11 +1,18 @@ -use crate::common::{INTEG_CLIENT_IDENTITY, INTEG_CLIENT_NAME, INTEG_CLIENT_VERSION, NAMESPACE}; +use crate::common::{ + INTEG_CLIENT_IDENTITY, INTEG_CLIENT_NAME, INTEG_CLIENT_VERSION, NAMESPACE, rand_6_chars, +}; use futures_util::{TryStreamExt, stream}; use std::time::{SystemTime, UNIX_EPOCH}; use temporalio_client::{ - Connection, ConnectionOptions, + Connection, ConnectionOptions, WorkflowStartOptions, grpc::{TestService, WorkflowService}, }; use temporalio_common::protos::temporal::api::workflowservice::v1::DescribeNamespaceRequest; +use temporalio_macros::{workflow, workflow_methods}; +use temporalio_sdk::{ + Runtime, Worker, WorkerOptions, WorkflowContext, WorkflowResult, + testing::{LocalWorkflowEnvironmentOptions, WorkflowEnvironment}, +}; use temporalio_sdk_core::ephemeral_server::{ EphemeralExe, EphemeralExeVersion, EphemeralServer, TemporalDevServerConfig, default_cached_download, @@ -13,6 +20,51 @@ use temporalio_sdk_core::ephemeral_server::{ use tonic::IntoRequest; use url::Url; +#[workflow] +#[derive(Default)] +struct TestEnvironmentWorkflow; + +#[workflow_methods] +impl TestEnvironmentWorkflow { + #[run] + async fn run(_ctx: &mut WorkflowContext) -> WorkflowResult<()> { + Ok(()) + } +} + +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::RequiresLocalServer)] +#[tokio::test] +async fn test_workflow_environment_local() { + let env = WorkflowEnvironment::start_local(LocalWorkflowEnvironmentOptions::default()) + .await + .unwrap(); + let runtime = Runtime::from_current_tokio(Default::default()).unwrap(); + let worker_options = WorkerOptions::new(format!("test-env-{}", rand_6_chars())) + .register_workflow::() + .unwrap() + .build(); + let task_queue = worker_options.task_queue.clone(); + let mut worker = Worker::new(&runtime, env.client().clone(), worker_options).unwrap(); + let shutdown = worker.shutdown_handle(); + let handle = env + .client() + .start_workflow( + TestEnvironmentWorkflow::run, + (), + WorkflowStartOptions::new(task_queue, format!("test-env-{}", rand_6_chars())).build(), + ) + .await + .unwrap(); + + let (worker_result, ()) = tokio::join!(worker.run(), async move { + handle.get_result(Default::default()).await.unwrap(); + shutdown(); + }); + worker_result.unwrap(); + env.shutdown().await.unwrap(); +} + +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::RequiresLocalServer)] #[tokio::test] async fn temporal_cli_default() { let config = TemporalDevServerConfig::builder() @@ -28,6 +80,7 @@ async fn temporal_cli_default() { assert!(sysinfo::System::new_all().process(pid).is_none()); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::RequiresLocalServer)] #[tokio::test] async fn temporal_cli_fixed() { let config = TemporalDevServerConfig::builder() @@ -38,6 +91,7 @@ async fn temporal_cli_fixed() { server.shutdown().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::RequiresLocalServer)] #[tokio::test] async fn temporal_cli_shutdown_port_reuse() { // Start, test shutdown, do again immediately on same port to ensure we can @@ -88,6 +142,7 @@ mod test_server { use super::*; use temporalio_sdk_core::ephemeral_server::TestServerConfig; + #[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::RequiresLocalServer)] #[tokio::test] async fn test_server_default() { let config = TestServerConfig::builder() @@ -98,6 +153,7 @@ mod test_server { server.shutdown().await.unwrap(); } + #[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::RequiresLocalServer)] #[tokio::test] async fn test_server_fixed() { let config = TestServerConfig::builder() @@ -108,6 +164,7 @@ mod test_server { server.shutdown().await.unwrap(); } + #[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::RequiresLocalServer)] #[tokio::test] async fn test_server_shutdown_port_reuse() { // Start, test shutdown, do again immediately on same port to ensure we can diff --git a/crates/sdk-core/tests/integ_tests/metrics_tests.rs b/crates/sdk-core/tests/integ_tests/metrics_tests.rs index 5da4859da..d6e2996f6 100644 --- a/crates/sdk-core/tests/integ_tests/metrics_tests.rs +++ b/crates/sdk-core/tests/integ_tests/metrics_tests.rs @@ -2,7 +2,7 @@ use crate::{ common::{ ANY_PORT, CoreWfStarter, NAMESPACE, OTEL_URL_ENV_VAR, PROMETHEUS_QUERY_API, eventually, get_integ_client, get_integ_connection, get_integ_runtime_options, - get_integ_server_options, get_integ_telem_options, prom_metrics, + get_integ_server_options, get_integ_telem_options, integ_namespace, prom_metrics, }, integ_tests::mk_nexus_endpoint, }; @@ -20,15 +20,19 @@ use std::{ }; use temporalio_client::{ Connection, MESSAGE_TOO_LARGE_KEY, NamespacedClient, REQUEST_LATENCY_HISTOGRAM_NAME, - UntypedQuery, UntypedWorkflow, WorkflowExecutionInfo, WorkflowQueryOptions, - WorkflowStartOptions, grpc::WorkflowService, + UntypedQuery, UntypedWorkflow, WorkflowExecutionInfo, WorkflowIdConflictPolicy, + WorkflowIdReusePolicy, WorkflowQueryOptions, WorkflowStartOptions, grpc::WorkflowService, }; use temporalio_common::{ data_converters::RawValue, + payload_limits::{LimitClass, LimitSeverity, PayloadLimitViolation}, protos::{ coresdk::{ ActivityTaskCompletion, - activity_result::ActivityExecutionResult, + activity_result::{ + self as activity_result, ActivityExecutionResult, ActivityTaskFailedCause, + activity_execution_result, + }, nexus::{NexusTaskCompletion, nexus_task, nexus_task_completion}, workflow_activation::{WorkflowActivationJob, workflow_activation_job}, workflow_commands::{ @@ -40,10 +44,7 @@ use temporalio_common::{ }, temporal::api::{ common::v1::RetryPolicy, - enums::v1::{ - NexusHandlerErrorRetryBehavior, WorkflowIdConflictPolicy, WorkflowIdReusePolicy, - WorkflowTaskFailedCause, - }, + enums::v1::{NexusHandlerErrorRetryBehavior, WorkflowTaskFailedCause}, failure::v1::Failure, nexus::{ self, @@ -52,7 +53,9 @@ use temporalio_common::{ request::Variant, start_operation_response, }, }, - workflowservice::v1::{DescribeNamespaceRequest, ListNamespacesRequest}, + workflowservice::v1::{ + DescribeNamespaceRequest, ListNamespacesRequest, PollActivityTaskQueueResponse, + }, }, }, telemetry::{ @@ -72,16 +75,17 @@ use temporalio_sdk::{ ActivityOptions, CancellableFuture, LocalActivityOptions, NexusOperationOptions, WorkflowContext, WorkflowResult, activities::{ActivityContext, ActivityError}, + runtime::{AutoscalingOptions, PollerBehavior}, }; use temporalio_sdk_core::{ - ActivitySlotKind, CoreRuntime, FixedSizeSlotSupplier, PollError, PollerBehavior, SlotKind, - SlotMarkUsedContext, SlotReleaseContext, SlotReservationContext, SlotSupplier, - SlotSupplierPermit, TokioRuntimeBuilder, TunerBuilder, WorkerConfig, WorkerVersioningStrategy, - WorkflowSlotKind, init_worker, prost_dur, + ActivitySlotKind, CoreRuntime, FixedSizeSlotSupplier, PollError, + PollerBehavior as CorePollerBehavior, SlotKind, SlotMarkUsedContext, SlotReleaseContext, + SlotReservationContext, SlotSupplier, SlotSupplierPermit, TokioRuntimeBuilder, TunerBuilder, + WorkerConfig, WorkerVersioningStrategy, WorkflowSlotKind, init_worker, prost_dur, replay::TestHistoryBuilder, test_help::{ - MockPollCfg, ResponseType, TemporalMeter, WorkerExt, WorkerTestHelpers, build_mock_pollers, - mock_worker, mock_worker_client, + MockPollCfg, MocksHolder, ResponseType, TemporalMeter, WorkerExt, WorkerTestHelpers, + build_mock_pollers, mock_worker, mock_worker_client, }, }; use tokio::{ @@ -96,15 +100,18 @@ pub(crate) async fn get_text(endpoint: String) -> String { reqwest::get(endpoint).await.unwrap().text().await.unwrap() } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::RequiresOssOnlyApis)] #[rstest::rstest] #[tokio::test] async fn prometheus_metrics_exported( + #[values(true, false)] counters_total_suffix: bool, #[values(true, false)] use_seconds_latency: bool, #[values(true, false)] custom_buckets: bool, ) { let opts = PrometheusExporterOptions::builder() .global_tags(HashMap::from([("global".to_string(), "hi!".to_string())])) .socket_addr(ANY_PORT.parse().unwrap()) + .counters_total_suffix(counters_total_suffix) .use_seconds_for_durations(use_seconds_latency) .histogram_bucket_overrides(if custom_buckets { HistogramBucketOverrides { @@ -153,8 +160,12 @@ async fn prometheus_metrics_exported( operation=\"GetSystemInfo\",service_name=\"temporal-core-sdk\",global=\"hi!\",le=\"50\"}" )); } - // Verify counter names are appropriate (don't end w/ '_total') - assert!(body.contains("temporal_request{")); + let request_metric_name = if counters_total_suffix { + "temporal_request_total" + } else { + "temporal_request" + }; + assert!(body.contains(&format!("{request_metric_name}{{"))); // Verify non-temporal metrics meter does not prefix let mm = rt.telemetry().get_metric_meter().unwrap(); let g = mm.gauge(MetricParameters::from("mygauge")); @@ -166,12 +177,13 @@ async fn prometheus_metrics_exported( #[tokio::test] async fn one_slot_worker_reports_available_slot() { + let namespace = integ_namespace(); let (telemopts, addr, _aborter) = prom_metrics(None); let tq = "one_slot_worker_tq"; let rt = CoreRuntime::new_assume_tokio(get_integ_runtime_options(telemopts)).unwrap(); let worker_cfg = WorkerConfig::builder() - .namespace(NAMESPACE) + .namespace(namespace.clone()) .task_queue(tq) .versioning_strategy(WorkerVersioningStrategy::None { build_id: "test_build_id".to_owned(), @@ -182,7 +194,7 @@ async fn one_slot_worker_reports_available_slot() { // Need to use two for WFTs because there are a minimum of 2 pollers b/c of sticky polling .max_outstanding_workflow_tasks(2_usize) .max_outstanding_nexus_tasks(1_usize) - .workflow_task_poller_behavior(PollerBehavior::SimpleMaximum(2_usize)) + .workflow_task_poller_behavior(CorePollerBehavior::SimpleMaximum(2_usize)) .task_types(WorkerTaskTypes::all()) .build() .unwrap(); @@ -265,22 +277,22 @@ async fn one_slot_worker_reports_available_slot() { tokio::time::sleep(Duration::from_millis(50)).await; let body = get_text(format!("http://{addr}/metrics")).await; assert!(body.contains(&format!( - "temporal_worker_task_slots_available{{namespace=\"{NAMESPACE}\",\ + "temporal_worker_task_slots_available{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"one_slot_worker_tq\",\ worker_type=\"WorkflowWorker\"}} 2" ))); assert!(body.contains(&format!( - "temporal_worker_task_slots_available{{namespace=\"{NAMESPACE}\",\ + "temporal_worker_task_slots_available{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"one_slot_worker_tq\",\ worker_type=\"ActivityWorker\"}} 1" ))); assert!(body.contains(&format!( - "temporal_worker_task_slots_available{{namespace=\"{NAMESPACE}\",\ + "temporal_worker_task_slots_available{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"one_slot_worker_tq\",\ worker_type=\"LocalActivityWorker\"}} 1" ))); assert!(body.contains(&format!( - "temporal_worker_task_slots_available{{namespace=\"{NAMESPACE}\",\ + "temporal_worker_task_slots_available{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"one_slot_worker_tq\",\ worker_type=\"NexusWorker\"}} 1" ))); @@ -304,37 +316,37 @@ async fn one_slot_worker_reports_available_slot() { // At this point the workflow task is outstanding and the activities haven't started let body = get_text(format!("http://{addr}/metrics")).await; assert!(body.contains(&format!( - "temporal_worker_task_slots_available{{namespace=\"{NAMESPACE}\",\ + "temporal_worker_task_slots_available{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"one_slot_worker_tq\",\ worker_type=\"WorkflowWorker\"}} 1" ))); assert!(body.contains(&format!( - "temporal_worker_task_slots_available{{namespace=\"{NAMESPACE}\",\ + "temporal_worker_task_slots_available{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"one_slot_worker_tq\",\ worker_type=\"ActivityWorker\"}} 1" ))); assert!(body.contains(&format!( - "temporal_worker_task_slots_available{{namespace=\"{NAMESPACE}\",\ + "temporal_worker_task_slots_available{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"one_slot_worker_tq\",\ worker_type=\"LocalActivityWorker\"}} 1" ))); assert!(body.contains(&format!( - "temporal_worker_task_slots_used{{namespace=\"{NAMESPACE}\",\ + "temporal_worker_task_slots_used{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"one_slot_worker_tq\",\ worker_type=\"WorkflowWorker\"}} 1" ))); assert!(body.contains(&format!( - "temporal_worker_task_slots_used{{namespace=\"{NAMESPACE}\",\ + "temporal_worker_task_slots_used{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"one_slot_worker_tq\",\ worker_type=\"ActivityWorker\"}} 0" ))); assert!(body.contains(&format!( - "temporal_worker_task_slots_used{{namespace=\"{NAMESPACE}\",\ + "temporal_worker_task_slots_used{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"one_slot_worker_tq\",\ worker_type=\"LocalActivityWorker\"}} 0" ))); assert!(body.contains(&format!( - "temporal_worker_task_slots_used{{namespace=\"{NAMESPACE}\",\ + "temporal_worker_task_slots_used{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"one_slot_worker_tq\",\ worker_type=\"NexusWorker\"}} 0" ))); @@ -347,17 +359,17 @@ async fn one_slot_worker_reports_available_slot() { tokio::time::sleep(Duration::from_millis(100)).await; let body = get_text(format!("http://{addr}/metrics")).await; assert!(body.contains(&format!( - "temporal_worker_task_slots_available{{namespace=\"{NAMESPACE}\",\ + "temporal_worker_task_slots_available{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"one_slot_worker_tq\",\ worker_type=\"WorkflowWorker\"}} 2" ))); assert!(body.contains(&format!( - "temporal_worker_task_slots_available{{namespace=\"{NAMESPACE}\",\ + "temporal_worker_task_slots_available{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"one_slot_worker_tq\",\ worker_type=\"ActivityWorker\"}} 0" ))); assert!(body.contains(&format!( - "temporal_worker_task_slots_used{{namespace=\"{NAMESPACE}\",\ + "temporal_worker_task_slots_used{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"one_slot_worker_tq\",\ worker_type=\"ActivityWorker\"}} 1" ))); @@ -368,7 +380,7 @@ async fn one_slot_worker_reports_available_slot() { act_task_barr.wait().await; let body = get_text(format!("http://{addr}/metrics")).await; assert!(body.contains(&format!( - "temporal_worker_task_slots_available{{namespace=\"{NAMESPACE}\",\ + "temporal_worker_task_slots_available{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"one_slot_worker_tq\",\ worker_type=\"ActivityWorker\"}} 1" ))); @@ -378,12 +390,12 @@ async fn one_slot_worker_reports_available_slot() { // Ensure that, once we have the LA task, slots are 0 let body = get_text(format!("http://{addr}/metrics")).await; assert!(body.contains(&format!( - "temporal_worker_task_slots_available{{namespace=\"{NAMESPACE}\",\ + "temporal_worker_task_slots_available{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"one_slot_worker_tq\",\ worker_type=\"LocalActivityWorker\"}} 0" ))); assert!(body.contains(&format!( - "temporal_worker_task_slots_used{{namespace=\"{NAMESPACE}\",\ + "temporal_worker_task_slots_used{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"one_slot_worker_tq\",\ worker_type=\"LocalActivityWorker\"}} 1" ))); @@ -392,7 +404,7 @@ async fn one_slot_worker_reports_available_slot() { act_task_barr.wait().await; let body = get_text(format!("http://{addr}/metrics")).await; assert!(body.contains(&format!( - "temporal_worker_task_slots_available{{namespace=\"{NAMESPACE}\",\ + "temporal_worker_task_slots_available{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"one_slot_worker_tq\",\ worker_type=\"LocalActivityWorker\"}} 1" ))); @@ -471,15 +483,17 @@ async fn idle_activity_worker_reports_zero_slots_used() { let rt = CoreRuntime::new_assume_tokio(get_integ_runtime_options(telemopts)).unwrap(); let mut starter = CoreWfStarter::new_with_runtime("idle_activity_worker_reports_zero_slots_used", rt); - starter.sdk_config.activity_task_poller_behavior = Some(PollerBehavior::Autoscaling { - minimum: 1, - maximum: 1, - initial: 1, - }); + starter.sdk_config.activity_task_poller_behavior = Some(PollerBehavior::Autoscaling( + AutoscalingOptions::builder() + .minimum(1) + .maximum(1) + .initial(1) + .build(), + )); let activity_slots = Arc::new(ReservationTrackingActivitySlotSupplier::new(3)); let mut tuner = TunerBuilder::default(); tuner.activity_slot_supplier(activity_slots.clone()); - starter.sdk_config.tuner = Arc::new(tuner.build()); + starter.set_core_tuner(Arc::new(tuner.build())); let finish_activity = Arc::new(Barrier::new(2)); struct BlockingActivity { @@ -580,7 +594,7 @@ async fn query_of_closed_workflow_doesnt_tick_terminal_metric( failure: Some(Failure::application_failure("I'm ded".to_string(), false)), }.into(), ContinueAsNewWorkflowExecution::default().into(), - CancelWorkflowExecution { }.into() + CancelWorkflowExecution::default().into() )] completion: workflow_command::Variant, ) { @@ -646,20 +660,19 @@ async fn query_of_closed_workflow_doesnt_tick_terminal_metric( // Query the now-closed workflow let client = starter.get_core_client().await; let queryer = async { - WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: starter.get_wf_id().to_string(), - run_id: Some(run_id), - first_execution_run_id: None, - } - .bind_untyped(client.clone()) - .query( - UntypedQuery::new("fake_query"), - RawValue::empty(), - WorkflowQueryOptions::default(), - ) - .await - .unwrap(); + WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(starter.get_wf_id().to_string()) + .maybe_run_id(Some(run_id)) + .build() + .bind_untyped(client.clone()) + .query( + UntypedQuery::new("fake_query"), + RawValue::empty(), + WorkflowQueryOptions::default(), + ) + .await + .unwrap(); }; let query_reply = async { // Need to re-complete b/c replay @@ -707,6 +720,7 @@ async fn query_of_closed_workflow_doesnt_tick_terminal_metric( assert!(matching_line.ends_with('1')); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::RequiresOssOnlyApis)] #[test] fn runtime_new() { let mut rt = CoreRuntime::new( @@ -840,6 +854,10 @@ async fn latency_metrics( ); } +#[temporalio_macros::cloud_test_exclusion( + crate::CloudTestExclusionReason::NeedsCloudAdaptation, + "Cloud authorizes the malformed request before validation, so it does not return the expected InvalidArgument status." +)] #[tokio::test] async fn request_fail_codes() { let (telemopts, addr, _aborter) = prom_metrics(None); @@ -907,6 +925,7 @@ async fn request_fail_codes_otel() { // Tests that rely on Prometheus running in a docker container need to start // with `docker_` and set the `DOCKER_PROMETHEUS_RUNNING` env variable to run +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::RequiresOssOnlyApis)] #[rstest::rstest] #[tokio::test] async fn docker_metrics_with_prometheus( @@ -1000,6 +1019,7 @@ async fn docker_metrics_with_prometheus( #[tokio::test] async fn activity_metrics() { + let namespace = integ_namespace(); let (telemopts, addr, _aborter) = prom_metrics(None); let rt = CoreRuntime::new_assume_tokio(get_integ_runtime_options(telemopts)).unwrap(); let wf_name = "activity_metrics"; @@ -1013,7 +1033,7 @@ async fn activity_metrics() { async fn pass_fail_act(ctx: ActivityContext, i: String) -> Result { match i.as_str() { "pass" => Ok("pass".to_string()), - "cancel" => { + "cancel" | "timeout" => { ctx.cancelled().await; Err(ActivityError::cancelled()) } @@ -1084,7 +1104,23 @@ async fn activity_metrics() { ) .build(), ); - let _ = join!(local_act_pass, local_act_fail); + // Outlives its start-to-close timeout, so core resolves it as timed out rather than + // as the cancel the activity reports once core stops it. + let local_act_timeout = ctx.execute_local_activity( + PassFailActivities::pass_fail_act, + "timeout".to_string(), + LocalActivityOptions::builder() + .start_to_close_timeout(Duration::from_millis(100)) + .retry_policy( + RetryPolicy { + maximum_attempts: 1, + ..Default::default() + } + .into(), + ) + .build(), + ); + let _ = join!(local_act_pass, local_act_fail, local_act_timeout); // TODO: Currently takes a WFT b/c of https://github.com/temporalio/sdk-core/issues/856 local_act_cancel.cancel(); let _ = local_act_cancel.await; @@ -1108,56 +1144,69 @@ async fn activity_metrics() { let wf_type = ActivityMetricsWf::name(); assert!(body.contains(&format!( "temporal_activity_execution_failed{{activity_type=\"pass_fail_act\",\ - namespace=\"{NAMESPACE}\",service_name=\"temporal-core-sdk\",\ + failure_reason=\"ActivityError\",\ + namespace=\"{namespace}\",service_name=\"temporal-core-sdk\",\ task_queue=\"{task_queue}\",workflow_type=\"{wf_type}\"}} 1" ))); assert!(body.contains(&format!( "temporal_activity_schedule_to_start_latency_count{{\ - namespace=\"{NAMESPACE}\",service_name=\"temporal-core-sdk\",\ + namespace=\"{namespace}\",service_name=\"temporal-core-sdk\",\ task_queue=\"{task_queue}\"}} 2" ))); assert!(body.contains(&format!( "temporal_activity_execution_latency_count{{activity_type=\"pass_fail_act\",\ - namespace=\"{NAMESPACE}\",service_name=\"temporal-core-sdk\",\ + namespace=\"{namespace}\",service_name=\"temporal-core-sdk\",\ task_queue=\"{task_queue}\",workflow_type=\"{wf_type}\"}} 2" ))); assert!(body.contains(&format!( "temporal_activity_succeed_endtoend_latency_count{{activity_type=\"pass_fail_act\",\ - namespace=\"{NAMESPACE}\",service_name=\"temporal-core-sdk\",\ + namespace=\"{namespace}\",service_name=\"temporal-core-sdk\",\ task_queue=\"{task_queue}\",workflow_type=\"{wf_type}\"}} 1" ))); assert!(body.contains(&format!( - "temporal_local_activity_total{{activity_type=\"pass_fail_act\",namespace=\"{NAMESPACE}\",\ + "temporal_local_activity_total{{activity_type=\"pass_fail_act\",namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",task_queue=\"{task_queue}\",\ - workflow_type=\"{wf_type}\"}} 3" + workflow_type=\"{wf_type}\"}} 4" ))); assert!(body.contains(&format!( "temporal_local_activity_execution_failed{{activity_type=\"pass_fail_act\",\ - namespace=\"{NAMESPACE}\",service_name=\"temporal-core-sdk\",\ + failure_reason=\"ActivityError\",\ + namespace=\"{namespace}\",service_name=\"temporal-core-sdk\",\ + task_queue=\"{task_queue}\",\ + workflow_type=\"{wf_type}\"}} 1" + ))); + assert!(body.contains(&format!( + "temporal_local_activity_execution_failed{{activity_type=\"pass_fail_act\",\ + failure_reason=\"timeout\",\ + namespace=\"{namespace}\",service_name=\"temporal-core-sdk\",\ task_queue=\"{task_queue}\",\ workflow_type=\"{wf_type}\"}} 1" ))); assert!(body.contains(&format!( "temporal_local_activity_execution_cancelled{{activity_type=\"pass_fail_act\",\ - namespace=\"{NAMESPACE}\",service_name=\"temporal-core-sdk\",\ + namespace=\"{namespace}\",service_name=\"temporal-core-sdk\",\ task_queue=\"{task_queue}\",\ workflow_type=\"{wf_type}\"}} 1" ))); assert!(body.contains(&format!( "temporal_local_activity_execution_latency_count{{activity_type=\"pass_fail_act\",\ - namespace=\"{NAMESPACE}\",service_name=\"temporal-core-sdk\",\ + namespace=\"{namespace}\",service_name=\"temporal-core-sdk\",\ task_queue=\"{task_queue}\",\ - workflow_type=\"{wf_type}\"}} 3" + workflow_type=\"{wf_type}\"}} 4" ))); assert!(body.contains(&format!( "temporal_local_activity_succeed_endtoend_latency_count{{activity_type=\"pass_fail_act\",\ - namespace=\"{NAMESPACE}\",service_name=\"temporal-core-sdk\",\ + namespace=\"{namespace}\",service_name=\"temporal-core-sdk\",\ task_queue=\"{task_queue}\",\ workflow_type=\"{wf_type}\"}} 1" ))); } +#[temporalio_macros::cloud_test_exclusion( + crate::CloudTestExclusionReason::RequiresCloudProvisioning, + "Creates a Nexus endpoint through an admin API unavailable to the isolated Cloud credential." +)] #[tokio::test] async fn nexus_metrics() { let (telemopts, addr, _aborter) = prom_metrics(None); @@ -1459,7 +1508,7 @@ async fn metrics_available_from_custom_slot_supplier() { inner: FixedSizeSlotSupplier::new(5), metrics: OnceLock::new(), })); - starter.sdk_config.tuner = Arc::new(tb.build()); + starter.set_core_tuner(Arc::new(tb.build())); starter .sdk_config .register_workflow::() @@ -1625,6 +1674,7 @@ async fn sticky_queue_label_strategy( )] strategy: TaskQueueLabelStrategy, ) { + let namespace = integ_namespace(); let (mut telemopts, addr, _aborter) = prom_metrics(Some( PrometheusExporterOptions::builder() .socket_addr(ANY_PORT.parse().unwrap()) @@ -1678,7 +1728,7 @@ async fn sticky_queue_label_strategy( .filter(|l| { l.contains("temporal_long_request") && l.contains("operation=\"PollWorkflowTaskQueue\"") - && l.contains(&format!("namespace=\"{NAMESPACE}\"")) + && l.contains(&format!("namespace=\"{namespace}\"")) }) .collect(); @@ -1724,7 +1774,7 @@ async fn resource_based_tuner_metrics() { let mut starter = CoreWfStarter::new_with_runtime(wf_name, rt); // Create a resource-based tuner with reasonable thresholds let tuner = ResourceBasedTuner::new(0.8, 0.8); - starter.sdk_config.tuner = Arc::new(tuner); + starter.set_core_tuner(Arc::new(tuner)); starter .sdk_config .register_workflow::() @@ -1782,6 +1832,7 @@ async fn resource_based_tuner_metrics() { ); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn terminal_metric_not_recorded_on_rejected_completion() { let prom_info = start_prometheus_metric_exporter( @@ -1865,6 +1916,7 @@ async fn terminal_metric_not_recorded_on_rejected_completion() { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn wf_task_latency_recorded_on_dropped_wft() { let (telemopts, addr, _aborter) = prom_metrics(None); @@ -1939,6 +1991,7 @@ async fn wf_task_latency_recorded_on_dropped_wft() { .unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn wf_task_execution_failed_metric_includes_workflow_type() { let (telemopts, addr, _aborter) = prom_metrics(None); @@ -2013,6 +2066,7 @@ async fn wf_task_execution_failed_metric_includes_workflow_type() { ); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn grpc_message_too_large_wf_task_execution_failed_metric_includes_workflow_type() { let (telemopts, addr, _aborter) = prom_metrics(None); @@ -2095,3 +2149,152 @@ async fn grpc_message_too_large_wf_task_execution_failed_metric_includes_workflo "Expected workflow_type label on metric, got: {metric_line}" ); } + +/// A cause reported by lang must survive to the metric as its own `failure_reason`, rather than +/// being flattened into the catch-all activity reason. +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] +#[tokio::test] +async fn lang_reported_activity_failure_cause_reaches_metric() { + let (telemopts, addr, _aborter) = prom_metrics(None); + let rt = CoreRuntime::new_assume_tokio(get_integ_runtime_options(telemopts)).unwrap(); + let meter = rt.telemetry().get_temporal_metric_meter().unwrap(); + + let mut mock_client = mock_worker_client(); + mock_client + .expect_fail_activity_task() + .times(1) + .returning(|_, cause, _, _| { + assert_eq!(cause, ActivityTaskFailedCause::ExternalStorageFailure); + Ok(Default::default()) + }); + + let mut mock = MocksHolder::from_client_with_activities( + mock_client, + [PollActivityTaskQueueResponse { + task_token: vec![1], + activity_id: "act1".to_string(), + activity_type: Some("act_type".into()), + ..Default::default() + } + .into()], + ); + mock.set_temporal_meter(meter); + let core = mock_worker(mock); + + let act = core.poll_activity_task().await.unwrap(); + core.complete_activity_task(ActivityTaskCompletion { + task_token: act.task_token, + result: Some(ActivityExecutionResult { + status: Some(activity_execution_result::Status::Failed( + activity_result::Failure { + failure: Some(Failure { + message: "storage exploded".to_string(), + ..Default::default() + }), + cause: ActivityTaskFailedCause::ExternalStorageFailure as i32, + }, + )), + }), + }) + .await + .unwrap(); + core.drain_activity_poller_and_shutdown().await; + + let metric_line = eventually( + || { + let endpoint = format!("http://{addr}/metrics"); + async move { + let body = get_text(endpoint).await; + body.lines() + .find(|l| l.starts_with("temporal_activity_execution_failed{")) + .map(ToString::to_string) + .ok_or_else(|| anyhow!("activity_execution_failed metric not found")) + } + }, + Duration::from_secs(5), + ) + .await + .unwrap(); + + assert!( + metric_line.contains("failure_reason=\"ExternalStorageError\""), + "Expected ExternalStorageError failure reason on metric, got: {metric_line}" + ); +} + +/// A payload-limit violation detected while reporting an activity result is core's own doing, so it +/// must reach the metric as its own reason rather than the generic activity one. +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] +#[tokio::test] +async fn payloads_too_large_activity_failure_reaches_metric() { + let (telemopts, addr, _aborter) = prom_metrics(None); + let rt = CoreRuntime::new_assume_tokio(get_integ_runtime_options(telemopts)).unwrap(); + let meter = rt.telemetry().get_temporal_metric_meter().unwrap(); + + let mut mock_client = mock_worker_client(); + mock_client + .expect_complete_activity_task() + .times(1) + .returning(|_, _| { + let violation = PayloadLimitViolation { + path: "result".to_string(), + class: LimitClass::Blob, + severity: LimitSeverity::Error, + size: 1024, + limit: 10, + }; + let mut status = tonic::Status::invalid_argument("Payload size limit exceeded"); + status.set_source(Arc::new(violation)); + Err(status) + }); + mock_client + .expect_fail_activity_task() + .times(1) + .returning(|_, cause, _, _| { + assert_eq!(cause, ActivityTaskFailedCause::PayloadsTooLarge); + Ok(Default::default()) + }); + + let mut mock = MocksHolder::from_client_with_activities( + mock_client, + [PollActivityTaskQueueResponse { + task_token: vec![1], + activity_id: "act1".to_string(), + activity_type: Some("act_type".into()), + ..Default::default() + } + .into()], + ); + mock.set_temporal_meter(meter); + let core = mock_worker(mock); + + let act = core.poll_activity_task().await.unwrap(); + core.complete_activity_task(ActivityTaskCompletion { + task_token: act.task_token, + result: Some(ActivityExecutionResult::ok(vec![0_u8; 1024].into())), + }) + .await + .unwrap(); + core.drain_activity_poller_and_shutdown().await; + + let metric_line = eventually( + || { + let endpoint = format!("http://{addr}/metrics"); + async move { + let body = get_text(endpoint).await; + body.lines() + .find(|l| l.starts_with("temporal_activity_execution_failed{")) + .map(ToString::to_string) + .ok_or_else(|| anyhow!("activity_execution_failed metric not found")) + } + }, + Duration::from_secs(5), + ) + .await + .unwrap(); + + assert!( + metric_line.contains("failure_reason=\"PayloadsTooLarge\""), + "Expected PayloadsTooLarge failure reason on metric, got: {metric_line}" + ); +} diff --git a/crates/sdk-core/tests/integ_tests/pagination_tests.rs b/crates/sdk-core/tests/integ_tests/pagination_tests.rs index 98cb473b8..76cac1e90 100644 --- a/crates/sdk-core/tests/integ_tests/pagination_tests.rs +++ b/crates/sdk-core/tests/integ_tests/pagination_tests.rs @@ -36,6 +36,7 @@ impl WeirdPaginationWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn weird_pagination_doesnt_drop_wft_events() { let wf_id = "fakeid"; @@ -160,6 +161,7 @@ impl ExtremePaginationWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn extreme_pagination_doesnt_drop_wft_events_worker() { let wf_id = "fakeid"; diff --git a/crates/sdk-core/tests/integ_tests/plugin_tests.rs b/crates/sdk-core/tests/integ_tests/plugin_tests.rs index 8325c706f..52ea47a26 100644 --- a/crates/sdk-core/tests/integ_tests/plugin_tests.rs +++ b/crates/sdk-core/tests/integ_tests/plugin_tests.rs @@ -36,7 +36,7 @@ use url::Url; use uuid::Uuid; fn new_sdk_runtime() -> Runtime { - Runtime::new_assume_tokio( + Runtime::from_current_tokio( RuntimeOptions::builder() .telemetry_options(get_integ_telem_options()) .build() @@ -87,6 +87,10 @@ impl WorkerPlugin for IntegrationPlugin { } } +#[temporalio_macros::cloud_test_exclusion( + crate::CloudTestExclusionReason::NeedsCloudAdaptation, + "Retargeting the client discards envconfig TLS options, so the HTTPS Cloud connection cannot be established." +)] #[tokio::test] async fn plugins_configure_client_and_worker() { let runtime = new_sdk_runtime(); @@ -114,13 +118,11 @@ async fn plugins_configure_client_and_worker() { .register_workflow::() .unwrap() .build(); - let worker = Worker::new(&runtime, client, worker_options).unwrap(); + let _worker = Worker::new(&runtime, client, worker_options).unwrap(); assert_eq!(connection_calls.load(Relaxed), 1); assert_eq!(client_calls.load(Relaxed), 1); assert_eq!(worker_calls.load(Relaxed), 1); - assert_eq!(worker.core_worker().get_config().max_cached_workflows, 0); - assert_eq!(worker.core_worker().get_config().plugins.len(), 1); } struct CountingPayloadCodec { @@ -218,7 +220,7 @@ async fn simple_plugin_configures_working_client_and_worker() { let worker_interceptor_calls = Arc::new(AtomicUsize::new(0)); let data_converter = DataConverter::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), CountingPayloadCodec { encode_calls: encode_calls.clone(), decode_calls: decode_calls.clone(), @@ -364,52 +366,3 @@ impl WorkerPlugin for FailingWorkerPlugin { Err(PluginError::new("worker failure")) } } - -struct ClientOnlyMetadataPlugin; - -impl ClientPlugin for ClientOnlyMetadataPlugin { - fn name(&self) -> &str { - "client-only-plugin" - } -} - -struct WorkerOnlyMetadataPlugin; - -impl WorkerPlugin for WorkerOnlyMetadataPlugin { - fn name(&self) -> &str { - "worker-only-plugin" - } -} - -#[tokio::test] -async fn worker_metadata_includes_client_and_worker_plugin_names() { - let runtime = new_sdk_runtime(); - let client = Client::connect( - get_integ_server_options(), - ClientOptions::new(integ_namespace()) - .client_plugin(ClientOnlyMetadataPlugin) - .build(), - ) - .await - .unwrap(); - let worker = Worker::new( - &runtime, - client, - WorkerOptions::new(format!("plugin-metadata-{}", Uuid::new_v4())) - .register_workflow::() - .unwrap() - .worker_plugin(WorkerOnlyMetadataPlugin) - .build(), - ) - .unwrap(); - let core_worker = worker.core_worker(); - let mut names = core_worker - .get_config() - .plugins - .iter() - .map(|plugin| plugin.name.clone()) - .collect::>(); - names.sort_unstable(); - - assert_eq!(names, ["client-only-plugin", "worker-only-plugin"]); -} diff --git a/crates/sdk-core/tests/integ_tests/polling_tests.rs b/crates/sdk-core/tests/integ_tests/polling_tests.rs index 648bf6dee..991555621 100644 --- a/crates/sdk-core/tests/integ_tests/polling_tests.rs +++ b/crates/sdk-core/tests/integ_tests/polling_tests.rs @@ -32,9 +32,12 @@ use temporalio_common::{ telemetry::{CoreLogStreamConsumer, Logger, TelemetryOptions}, }; use temporalio_macros::{workflow, workflow_methods}; -use temporalio_sdk::{ActivityOptions, WorkflowContext, WorkflowResult}; +use temporalio_sdk::{ + ActivityOptions, WorkflowContext, WorkflowResult, + runtime::{AutoscalingOptions, PollerBehavior}, +}; use temporalio_sdk_core::{ - CoreRuntime, PollerBehavior, RuntimeOptions, TunerHolder, + CoreRuntime, RuntimeOptions, TunerHolder, ephemeral_server::{TemporalDevServerConfig, default_cached_download}, init_worker, prost_dur, test_help::{NAMESPACE, WorkerTestHelpers, drain_pollers_and_shutdown, schedule_activity_cmd}, @@ -120,6 +123,7 @@ async fn out_of_order_completion_doesnt_hang() { jh.await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::RequiresLocalServer)] #[tokio::test] async fn switching_worker_client_changes_poll() { // Start two servers @@ -205,16 +209,15 @@ async fn switching_worker_client_changes_poll() { worker.complete_execution(&act1.run_id).await; worker.handle_eviction().await; info!("Waiting on first workflow complete"); - WorkflowExecutionInfo { - namespace: client1.namespace(), - workflow_id: "my-workflow-1".into(), - run_id: Some(wf1_run_id.clone()), - first_execution_run_id: None, - } - .bind_untyped(client1.clone()) - .get_result(Default::default()) - .await - .unwrap(); + WorkflowExecutionInfo::builder() + .namespace(client1.namespace()) + .workflow_id("my-workflow-1") + .maybe_run_id(Some(wf1_run_id.clone())) + .build() + .bind_untyped(client1.clone()) + .get_result(Default::default()) + .await + .unwrap(); // Swap client, poll for next task, confirm it's second wf, and respond w/ empty info!("Replacing client and polling again"); @@ -224,16 +227,15 @@ async fn switching_worker_client_changes_poll() { worker.complete_execution(&act2.run_id).await; worker.handle_eviction().await; info!("Waiting on second workflow complete"); - WorkflowExecutionInfo { - namespace: client2.namespace(), - workflow_id: "my-workflow-2".into(), - run_id: Some(wf2_run_id), - first_execution_run_id: None, - } - .bind_untyped(client2.clone()) - .get_result(Default::default()) - .await - .unwrap(); + WorkflowExecutionInfo::builder() + .namespace(client2.namespace()) + .workflow_id("my-workflow-2") + .maybe_run_id(Some(wf2_run_id)) + .build() + .bind_untyped(client2.clone()) + .get_result(Default::default()) + .await + .unwrap(); // Shutdown workers and servers drain_pollers_and_shutdown(&worker).await; @@ -277,16 +279,18 @@ async fn small_workflow_slots_and_pollers(#[values(false, true)] use_autoscaling let wf_name = "only_one_workflow_slot_and_two_pollers"; let mut starter = CoreWfStarter::new(wf_name); if use_autoscaling { - starter.sdk_config.workflow_task_poller_behavior = Some(PollerBehavior::Autoscaling { - minimum: 1, - maximum: 5, - initial: 1, - }); + starter.sdk_config.workflow_task_poller_behavior = Some(PollerBehavior::Autoscaling( + AutoscalingOptions::builder() + .minimum(1) + .maximum(5) + .initial(1) + .build(), + )); } else { starter.sdk_config.workflow_task_poller_behavior = Some(PollerBehavior::SimpleMaximum(2)); } starter.sdk_config.activity_task_poller_behavior = Some(PollerBehavior::SimpleMaximum(1)); - starter.sdk_config.tuner = Arc::new(TunerHolder::fixed_size(2, 1, 1, 1)); + starter.set_core_tuner(Arc::new(TunerHolder::fixed_size(2, 1, 1, 1))); starter .sdk_config .register_activities(StdActivities) @@ -325,15 +329,16 @@ async fn small_workflow_slots_and_pollers(#[values(false, true)] use_autoscaling .await .get_workflow_handle::(&wf2id) .fetch_history(Default::default()) + .into_events() .await - .unwrap() - .into_events(); + .unwrap(); let any_task_timeouts = events .iter() .any(|e| e.event_type() == EventType::WorkflowTaskTimedOut); assert!(!any_task_timeouts); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::RequiresLocalServer)] #[tokio::test] async fn replace_client_works_after_polling_failure() { let (log_consumer, mut log_rx) = CoreLogStreamConsumer::new(100); diff --git a/crates/sdk-core/tests/integ_tests/queries_tests.rs b/crates/sdk-core/tests/integ_tests/queries_tests.rs index a4f1c337f..5eff39dae 100644 --- a/crates/sdk-core/tests/integ_tests/queries_tests.rs +++ b/crates/sdk-core/tests/integ_tests/queries_tests.rs @@ -49,36 +49,34 @@ async fn simple_query_legacy() { .unwrap(); tokio::time::sleep(Duration::from_secs(1)).await; // Query after timer should have fired and there should be new WFT + let timer_task = core.poll_workflow_activation().await.unwrap(); + assert_matches!( + timer_task.jobs.as_slice(), + [WorkflowActivationJob { + variant: Some(workflow_activation_job::Variant::FireTimer(_)), + }] + ); let query_fut = async { - WorkflowExecutionInfo { - namespace: starter.get_core_client().await.namespace(), - workflow_id, - run_id: Some(task.run_id.to_string()), - first_execution_run_id: None, - } - .bind_untyped(starter.get_core_client().await.clone()) - .query( - UntypedQuery::new("myquery"), - RawValue::empty(), - WorkflowQueryOptions::default(), - ) - .await - .unwrap() + WorkflowExecutionInfo::builder() + .namespace(starter.get_core_client().await.namespace()) + .workflow_id(workflow_id) + .maybe_run_id(Some(task.run_id.to_string())) + .build() + .bind_untyped(starter.get_core_client().await.clone()) + .query( + UntypedQuery::new("myquery"), + RawValue::empty(), + WorkflowQueryOptions::default(), + ) + .await + .unwrap() }; let workflow_completions_future = async { - // Give query a beat to get going + // Let the query reach the server before completing the outstanding timer task so the + // server sends the query activation next. tokio::time::sleep(Duration::from_millis(400)).await; - // This poll *should* have the `queries` field populated, but doesn't, seemingly due to - // a server bug. So, complete the WF task of the first timer firing with empty commands - let task = core.poll_workflow_activation().await.unwrap(); - assert_matches!( - task.jobs.as_slice(), - [WorkflowActivationJob { - variant: Some(workflow_activation_job::Variant::FireTimer(_)), - }] - ); core.complete_workflow_activation(WorkflowActivationCompletion::from_cmds( - task.run_id, + timer_task.run_id, vec![], )) .await @@ -197,20 +195,19 @@ async fn query_after_execution_complete(#[case] do_evict: bool) { for _ in 0..3 { let gw = starter.get_core_client().await.clone(); let query_fut = async move { - let q_resp: RawValue = WorkflowExecutionInfo { - namespace: gw.namespace(), - workflow_id: workflow_id.to_string(), - run_id: Some(run_id.to_string()), - first_execution_run_id: None, - } - .bind_untyped(gw.clone()) - .query( - UntypedQuery::new("myquery"), - RawValue::empty(), - WorkflowQueryOptions::default(), - ) - .await - .unwrap(); + let q_resp: RawValue = WorkflowExecutionInfo::builder() + .namespace(gw.namespace()) + .workflow_id(workflow_id.to_string()) + .maybe_run_id(Some(run_id.to_string())) + .build() + .bind_untyped(gw.clone()) + .query( + UntypedQuery::new("myquery"), + RawValue::empty(), + WorkflowQueryOptions::default(), + ) + .await + .unwrap(); // Ensure query response is as expected assert_eq!(q_resp.payloads[0].data, query_resp); }; @@ -239,20 +236,19 @@ async fn fail_legacy_query(#[case] with_nde: bool) { core.complete_execution(&task.run_id).await; core.handle_eviction().await; let query_fut = async { - WorkflowExecutionInfo { - namespace: starter.get_core_client().await.namespace(), - workflow_id: workflow_id.to_string(), - run_id: Some(task.run_id.to_string()), - first_execution_run_id: None, - } - .bind_untyped(starter.get_core_client().await.clone()) - .query( - UntypedQuery::new("myquery"), - RawValue::empty(), - WorkflowQueryOptions::default(), - ) - .await - .unwrap_err() + WorkflowExecutionInfo::builder() + .namespace(starter.get_core_client().await.namespace()) + .workflow_id(workflow_id.to_string()) + .maybe_run_id(Some(task.run_id.to_string())) + .build() + .bind_untyped(starter.get_core_client().await.clone()) + .query( + UntypedQuery::new("myquery"), + RawValue::empty(), + WorkflowQueryOptions::default(), + ) + .await + .unwrap_err() }; let query_responder = async { // Have to replay first since we've evicted @@ -311,20 +307,19 @@ async fn multiple_concurrent_queries_no_new_history() { let client = starter.get_core_client().await; let num_queries = 10; let query_futs = (1..=num_queries).map(|_| async { - WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: workflow_id.to_string(), - run_id: Some(task.run_id.to_string()), - first_execution_run_id: None, - } - .bind_untyped(client.clone()) - .query( - UntypedQuery::new("myquery"), - RawValue::empty(), - WorkflowQueryOptions::default(), - ) - .await - .unwrap(); + WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(workflow_id.to_string()) + .maybe_run_id(Some(task.run_id.to_string())) + .build() + .bind_untyped(client.clone()) + .query( + UntypedQuery::new("myquery"), + RawValue::empty(), + WorkflowQueryOptions::default(), + ) + .await + .unwrap(); }); let complete_fut = async { for _ in 1..=num_queries { @@ -381,20 +376,19 @@ async fn queries_handled_before_next_wft() { let client = starter.get_core_client().await; // Send two queries so that one of them is buffered let query_futs = (1..=2).map(|_| async { - WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: workflow_id.to_string(), - run_id: Some(task.run_id.to_string()), - first_execution_run_id: None, - } - .bind_untyped(client.clone()) - .query( - UntypedQuery::new("myquery"), - RawValue::empty(), - WorkflowQueryOptions::default(), - ) - .await - .unwrap(); + WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(workflow_id.to_string()) + .maybe_run_id(Some(task.run_id.to_string())) + .build() + .bind_untyped(client.clone()) + .query( + UntypedQuery::new("myquery"), + RawValue::empty(), + WorkflowQueryOptions::default(), + ) + .await + .unwrap(); }); let complete_fut = async { let task = core.poll_workflow_activation().await.unwrap(); @@ -406,20 +400,19 @@ async fn queries_handled_before_next_wft() { ); // While handling the first query, signal the workflow so a new WFT is generated and the // second query is still in the buffer - WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: workflow_id.to_string(), - run_id: Some(task.run_id.to_string()), - first_execution_run_id: None, - } - .bind_untyped(client.clone()) - .signal( - UntypedSignal::new("blah"), - RawValue::empty(), - WorkflowSignalOptions::default(), - ) - .await - .unwrap(); + WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(workflow_id.to_string()) + .maybe_run_id(Some(task.run_id.to_string())) + .build() + .bind_untyped(client.clone()) + .signal( + UntypedSignal::new("blah"), + RawValue::empty(), + WorkflowSignalOptions::default(), + ) + .await + .unwrap(); tokio::time::sleep(Duration::from_millis(500)).await; core.complete_workflow_activation(WorkflowActivationCompletion::from_cmd( task.run_id, diff --git a/crates/sdk-core/tests/integ_tests/schedule_tests.rs b/crates/sdk-core/tests/integ_tests/schedule_tests.rs index 68777db0e..83bd97dc8 100644 --- a/crates/sdk-core/tests/integ_tests/schedule_tests.rs +++ b/crates/sdk-core/tests/integ_tests/schedule_tests.rs @@ -1,4 +1,4 @@ -use crate::common::{NAMESPACE, eventually, get_integ_client, rand_6_chars}; +use crate::common::{eventually, get_integ_client, integ_namespace, rand_6_chars}; use futures::TryStreamExt; use std::time::{Duration, SystemTime}; use temporalio_client::{ @@ -13,7 +13,7 @@ use temporalio_macros::{workflow, workflow_methods}; use temporalio_sdk::{WorkflowContext, WorkflowResult}; async fn test_client() -> temporalio_client::Client { - get_integ_client(NAMESPACE.to_string(), None).await + get_integ_client(integ_namespace(), None).await } #[workflow] diff --git a/crates/sdk-core/tests/integ_tests/standalone_activity_tests.rs b/crates/sdk-core/tests/integ_tests/standalone_activity_tests.rs new file mode 100644 index 000000000..6f19127ef --- /dev/null +++ b/crates/sdk-core/tests/integ_tests/standalone_activity_tests.rs @@ -0,0 +1,259 @@ +use crate::common::CoreWfStarter; +use futures_util::{FutureExt, StreamExt, pin_mut, stream}; +use std::{ + collections::HashSet, + panic, + panic::{AssertUnwindSafe, resume_unwind}, + sync::Arc, + time::Duration, +}; +use temporalio_client::{ + ActivityCancelOptions, ActivityDescribeOptions, ActivityExecutionInfoLike, + ActivityExecutionStatus, ActivityStartOptions, ActivityStartOptionsBuilder, + ActivityTerminateOptions, Client, NamespacedClient, errors::ActivityResultError, +}; +use temporalio_common::ActivityError; +use temporalio_macros::activities; +use temporalio_sdk::activities::ActivityContext; +use uuid::Uuid; + +const TASK_QUEUE_PREFIX: &str = "standalone_activity_tests"; + +struct Activities; + +#[activities] +impl Activities { + #[activity] + async fn echo(_ctx: ActivityContext, e: String) -> Result { + Ok(e) + } + + #[activity] + async fn wait_for_cancel(self: Arc, ctx: ActivityContext) -> Result<(), ActivityError> { + let mut ticker = tokio::time::interval(Duration::from_millis(100)); + loop { + tokio::select! { biased; + _ = ctx.cancelled() => return Err(ActivityError::Cancelled {details: None}), + _ = ticker.tick() => { let _ = ctx.record_heartbeat(()).await; }, + } + } + } +} + +async fn run_test(test: impl AsyncFnOnce(Client, String)) { + let mut starter = CoreWfStarter::new(TASK_QUEUE_PREFIX); + starter.sdk_config.register_activities(Activities); + let mut worker = starter.worker().await; + let client = starter.get_core_client().await; + let shutdown_handle = worker.inner_mut().shutdown_handle(); + + let worker_fut = worker.inner_mut().run(); + let test_fut = async { + let result = AssertUnwindSafe(test(client, starter.sdk_config.task_queue.clone())) + .catch_unwind() + .await; + shutdown_handle(); + result + }; + pin_mut!(worker_fut); + pin_mut!(test_fut); + + tokio::select! { + test_result = &mut test_fut => { + let worker_result = worker_fut.await; + if let Err(panic) = test_result { + resume_unwind(panic); + } + worker_result.unwrap(); + }, + worker_result = &mut worker_fut => { + worker_result.unwrap(); + if let Err(panic) = test_fut.await { + resume_unwind(panic); + } + } + } +} + +fn test_options(task_queue: String) -> ActivityStartOptionsBuilder { + ActivityStartOptions::with_schedule_to_close_timeout( + task_queue, + Uuid::new_v4(), + Duration::from_secs(60), + ) +} + +#[tokio::test] +async fn get_result() { + run_test(async |client, tq| { + let options = test_options(tq).build(); + let arg = "Hello"; + + let handle = client + .start_activity(Activities::echo, arg.into(), options.clone()) + .await + .unwrap(); + assert_eq!(handle.activity_id(), options.id); + assert!(handle.run_id().is_some()); + assert_eq!(handle.result().await.unwrap(), arg); + + let new_handle = client.get_activity_handle( + Activities::echo, + handle.activity_id(), + handle.run_id().map(Into::into), + ); + assert_eq!(new_handle.result().await.unwrap(), arg); + + let untyped_handle = client + .get_untyped_activity_handle(handle.activity_id(), handle.run_id().map(Into::into)); + assert_eq!( + untyped_handle + .result() + .await + .unwrap() + .to_value::(client.data_converter().payload_converter()), + arg + ); + + let wrong_run_id = loop { + let uuid = Some(Uuid::new_v4().to_string()); + if uuid.as_deref() != handle.run_id() { + break uuid; + } + }; + + let handle_wrong_run_id = + client.get_activity_handle(Activities::echo, handle.activity_id(), wrong_run_id); + assert_matches!( + handle_wrong_run_id.result().await, + Err(ActivityResultError::NotFound(_)) + ); + + let handle_no_run_id = + client.get_activity_handle(Activities::echo, handle.activity_id(), None); + assert_eq!(handle_no_run_id.result().await.unwrap(), arg); + }) + .await; +} + +#[tokio::test] +async fn describe() { + run_test(async |client, tq| { + let options = test_options(tq).build(); + let arg = "Hello"; + + let handle = client + .start_activity(Activities::echo, arg.into(), options.clone()) + .await + .unwrap(); + let result = handle.result().await.unwrap(); + + let desc = handle + .describe( + ActivityDescribeOptions::builder() + .include_input(true) + .include_outcome(true) + .build(), + ) + .await + .unwrap(); + + assert_eq!(desc.activity_id(), options.id); + assert_eq!(Some(desc.activity_run_id()), handle.run_id()); + assert_eq!(desc.status(), ActivityExecutionStatus::Completed); + assert_eq!(desc.input().await.unwrap(), Some(arg.to_string())); + assert_eq!(desc.outcome().await.unwrap().unwrap().unwrap(), result); + }) + .await; +} + +#[tokio::test] +async fn cancel() { + run_test(async |client, tq| { + let reason = "test cancel"; + let handle = client + .start_activity(Activities::wait_for_cancel, (), test_options(tq).build()) + .await + .unwrap(); + handle + .cancel(ActivityCancelOptions::builder().reason(reason).build()) + .await + .unwrap(); + + assert_matches!( + handle.result().await, + Err(ActivityResultError::Cancelled { .. }) + ); + let desc = handle.describe(Default::default()).await.unwrap(); + assert_eq!(desc.status(), ActivityExecutionStatus::Canceled); + assert_eq!(desc.canceled_reason(), Some(reason)); + }) + .await; +} + +#[tokio::test] +async fn terminate() { + run_test(async |client, tq| { + let reason = "test terminate"; + let handle = client + .start_activity(Activities::wait_for_cancel, (), test_options(tq).build()) + .await + .unwrap(); + handle + .terminate(ActivityTerminateOptions::builder().reason(reason).build()) + .await + .unwrap(); + + assert_matches!(handle.result().await, Err(ActivityResultError::Terminated)); + let desc = handle.describe(Default::default()).await.unwrap(); + assert_eq!(desc.status(), ActivityExecutionStatus::Terminated); + }) + .await; +} + +#[temporalio_macros::cloud_test_exclusion( + crate::CloudTestExclusionReason::NeedsCloudAdaptation, + "Cloud list visibility can lag behind count visibility, but the test lists activities only once." +)] +#[tokio::test] +async fn list_and_count() { + run_test(async |client, tq| { + let query = format!("TaskQueue='{tq}'"); + + let started_activity_ids: HashSet<_> = stream::iter(0..3) + .then(async |_| { + client + .start_activity( + Activities::echo, + "Hello".into(), + test_options(tq.clone()).build(), + ) + .await + .unwrap() + .activity_id() + .to_string() + }) + .collect() + .await; + + // in loop because of eventual consistency + loop { + let count = client + .count_activities(query.clone(), Default::default()) + .await + .unwrap(); + if count.count() == started_activity_ids.len() { + break; + } + tokio::time::sleep(Duration::from_millis(100)).await; + } + + let list_activity_ids: HashSet<_> = client + .list_activities(query.clone(), Default::default()) + .map(|a| a.unwrap().activity_id().to_string()) + .collect() + .await; + assert_eq!(list_activity_ids, started_activity_ids); + }) + .await; +} diff --git a/crates/sdk-core/tests/integ_tests/update_tests.rs b/crates/sdk-core/tests/integ_tests/update_tests.rs index 5f4f2491e..448cc1445 100644 --- a/crates/sdk-core/tests/integ_tests/update_tests.rs +++ b/crates/sdk-core/tests/integ_tests/update_tests.rs @@ -14,8 +14,10 @@ use std::{ }; use temporalio_client::{ Client, NamespacedClient, UntypedSignal, UntypedUpdate, UntypedWorkflow, - WorkflowExecuteUpdateOptions, WorkflowExecutionInfo, WorkflowSignalOptions, - WorkflowStartOptions, errors::WorkflowUpdateError, grpc::WorkflowService, + WorkflowExecuteUpdateOptions, WorkflowExecutionInfo, WorkflowIdConflictPolicy, + WorkflowSignalOptions, WorkflowStartOptions, WorkflowUpdateWithStartOptions, + errors::{WorkflowStartError, WorkflowUpdateError, WorkflowUpdateWithStartError}, + grpc::WorkflowService, }; use temporalio_common::{ data_converters::RawValue, @@ -37,6 +39,7 @@ use temporalio_common::{ workflowservice::v1::{ResetStickyTaskQueueRequest, ResetWorkflowExecutionRequest}, }, }, + worker::WorkerTaskTypes, }; use temporalio_macros::{activities, workflow, workflow_methods}; use temporalio_sdk::{ @@ -87,9 +90,9 @@ async fn update_workflow(#[values(FailUpdate::Yes, FailUpdate::No)] will_fail: F let events = client .get_workflow_handle::(workflow_id) .fetch_history(Default::default()) + .into_events() .await - .unwrap() - .into_events(); + .unwrap(); let with_id = HistoryForReplay::new(events, workflow_id.to_string()); let replay_worker = init_core_replay_preloaded(workflow_id, [with_id]); // Init workflow comes by itself @@ -160,17 +163,16 @@ async fn reapplied_updates_due_to_reset() { assert_eq!(post_reset_run_id, reset_response.run_id); // Make sure replay works - let events = WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: workflow_id.to_string(), - run_id: Some(post_reset_run_id.clone()), - first_execution_run_id: None, - } - .bind_untyped(client.clone()) - .fetch_history(Default::default()) - .await - .unwrap() - .into_events(); + let events = WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(workflow_id.to_string()) + .maybe_run_id(Some(post_reset_run_id.clone())) + .build() + .bind_untyped(client.clone()) + .fetch_history(Default::default()) + .into_events() + .await + .unwrap(); let with_id = HistoryForReplay::new(events, workflow_id.to_string()); let replay_worker = init_core_replay_preloaded(workflow_id, [with_id]); @@ -205,13 +207,12 @@ async fn send_and_handle_update( .await .unwrap(); - let handle = WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: workflow_id.to_string(), - run_id: Some(act.run_id.clone()), - first_execution_run_id: None, - } - .bind_untyped(client.clone()); + let handle = WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(workflow_id.to_string()) + .maybe_run_id(Some(act.run_id.clone())) + .build() + .bind_untyped(client.clone()); // Send the update to the server let update_task = async { @@ -314,13 +315,12 @@ async fn update_rejection() { .await .unwrap(); - let handle = WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: workflow_id.clone(), - run_id: Some(res.run_id.clone()), - first_execution_run_id: None, - } - .bind_untyped(client.clone()); + let handle = WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(workflow_id.clone()) + .maybe_run_id(Some(res.run_id.clone())) + .build() + .bind_untyped(client.clone()); // Send the update to the server let update_task = async { @@ -363,9 +363,9 @@ async fn update_rejection() { let events = client .get_workflow_handle::(&workflow_id) .fetch_history(Default::default()) + .into_events() .await - .unwrap() - .into_events(); + .unwrap(); let has_update_event = events.iter().any(|e| { matches!( e.event_type(), @@ -393,13 +393,12 @@ async fn update_insta_complete(#[values(true, false)] accept_first: bool) { .await .unwrap(); - let handle = WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id, - run_id: Some(res.run_id.clone()), - first_execution_run_id: None, - } - .bind_untyped(client.clone()); + let handle = WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(workflow_id) + .maybe_run_id(Some(res.run_id.clone())) + .build() + .bind_untyped(client.clone()); // Send the update to the server let (update_task, stop_wait_update) = future::abortable(async { @@ -486,13 +485,12 @@ async fn update_complete_after_accept_without_new_task() { .await .unwrap(); - let handle = WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id, - run_id: Some(res.run_id.clone()), - first_execution_run_id: None, - } - .bind_untyped(client.clone()); + let handle = WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(workflow_id) + .maybe_run_id(Some(res.run_id.clone())) + .build() + .bind_untyped(client.clone()); // Send the update to the server let update_task = async { @@ -1527,3 +1525,168 @@ async fn update_lost_on_activity_mismatch() { join!(update, runner); handle.fetch_history_and_replay(&mut worker).await.unwrap(); } + +#[workflow] +#[derive(Default)] +struct UpdateWithStartWf { + done: bool, +} + +#[workflow_methods] +impl UpdateWithStartWf { + #[run] + async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { + ctx.wait_condition(|s| s.done).await?; + Ok(()) + } + + #[update] + async fn do_update( + _ctx: &mut WorkflowContext, + arg: String, + ) -> Result> { + if arg == "reject" { + return Err(anyhow!("update rejected").into()); + } + Ok(format!("hello {arg}")) + } + + #[signal] + fn done_signal(&mut self, _ctx: &mut SyncWorkflowContext, _: ()) { + self.done = true; + } +} + +#[derive(Clone, Copy)] +enum UpdateWithStartScenario { + StartAndGetHandle, + ExecuteOnExisting, + UpdateFailure, + StartConflict, +} + +#[rstest::rstest] +#[case::start_and_get_handle(UpdateWithStartScenario::StartAndGetHandle)] +#[case::execute_on_existing(UpdateWithStartScenario::ExecuteOnExisting)] +#[case::update_failure(UpdateWithStartScenario::UpdateFailure)] +#[case::start_conflict(UpdateWithStartScenario::StartConflict)] +#[tokio::test] +async fn update_with_start(#[case] scenario: UpdateWithStartScenario) { + let mut starter = CoreWfStarter::new("update_with_start"); + starter + .sdk_config + .register_workflow::() + .unwrap(); + starter.set_core_task_types(WorkerTaskTypes::workflow_only()); + let mut worker = starter.worker().await; + let client = starter.get_core_client().await; + let task_queue = starter.get_task_queue().to_owned(); + let wf_id = starter.get_wf_id().to_owned(); + + let existing_run_id = if matches!(scenario, UpdateWithStartScenario::StartAndGetHandle) { + None + } else { + let handle = worker + .submit_workflow( + UpdateWithStartWf::run, + (), + WorkflowStartOptions::new(task_queue.clone(), wf_id.clone()).build(), + ) + .await + .unwrap(); + Some(handle.run_id().unwrap().to_owned()) + }; + + let core_worker = worker.core_worker(); + let interactions = async { + let options = |conflict_policy| { + WorkflowUpdateWithStartOptions::new(task_queue.clone(), wf_id.clone(), conflict_policy) + .execution_timeout(Duration::from_secs(60 * 5)) + .build() + }; + match scenario { + UpdateWithStartScenario::StartAndGetHandle => { + let update_handle = client + .start_update_with_start_workflow( + UpdateWithStartWf::run, + (), + UpdateWithStartWf::do_update, + "world".to_owned(), + options(WorkflowIdConflictPolicy::Fail), + ) + .await + .unwrap(); + assert!(update_handle.workflow_run_id().is_some()); + assert_eq!( + update_handle.get_result(Default::default()).await.unwrap(), + "hello world" + ); + } + UpdateWithStartScenario::ExecuteOnExisting => { + let result = client + .execute_update_with_start_workflow( + UpdateWithStartWf::run, + (), + UpdateWithStartWf::do_update, + "again".to_owned(), + options(WorkflowIdConflictPolicy::UseExisting), + ) + .await + .unwrap(); + assert_eq!(result, "hello again"); + } + UpdateWithStartScenario::UpdateFailure => { + let error = client + .execute_update_with_start_workflow( + UpdateWithStartWf::run, + (), + UpdateWithStartWf::do_update, + "reject".to_owned(), + options(WorkflowIdConflictPolicy::UseExisting), + ) + .await + .expect_err("rejected update must be returned as an update failure"); + assert_matches!( + error, + WorkflowUpdateWithStartError::Update(WorkflowUpdateError::Failed(failure)) + if failure.message.contains("update rejected") + ); + } + UpdateWithStartScenario::StartConflict => { + let error = client + .execute_update_with_start_workflow( + UpdateWithStartWf::run, + (), + UpdateWithStartWf::do_update, + "unused".to_owned(), + options(WorkflowIdConflictPolicy::Fail), + ) + .await + .expect_err("update-with-start must fail against a running workflow"); + assert_matches!( + error, + WorkflowUpdateWithStartError::Start(WorkflowStartError::AlreadyStarted { + run_id: Some(run_id), + .. + }) if existing_run_id.as_deref() == Some(run_id.as_str()) + ); + } + } + + let wf_handle = client.get_workflow_handle::(wf_id); + wf_handle + .signal( + UpdateWithStartWf::done_signal, + (), + WorkflowSignalOptions::default(), + ) + .await + .unwrap(); + wf_handle.get_result(Default::default()).await.unwrap(); + core_worker.initiate_shutdown(); + }; + let run = async { + worker.inner_mut().run().await.unwrap(); + }; + join!(interactions, run); +} diff --git a/crates/sdk-core/tests/integ_tests/visibility_tests.rs b/crates/sdk-core/tests/integ_tests/visibility_tests.rs index e151bbf11..e27c6f808 100644 --- a/crates/sdk-core/tests/integ_tests/visibility_tests.rs +++ b/crates/sdk-core/tests/integ_tests/visibility_tests.rs @@ -1,4 +1,4 @@ -use crate::common::{CoreWfStarter, NAMESPACE, eventually, get_integ_client}; +use crate::common::{CoreWfStarter, NAMESPACE, eventually, get_integ_client, integ_namespace}; use assert_matches::assert_matches; use std::{sync::Arc, time::Duration}; use temporalio_client::{NamespacedClient, RegisterNamespaceOptions, grpc::WorkflowService}; @@ -17,6 +17,10 @@ use temporalio_sdk_core::test_help::{WorkerTestHelpers, drain_pollers_and_shutdo use tokio::time::sleep; use tonic::IntoRequest; +#[temporalio_macros::cloud_test_exclusion( + crate::CloudTestExclusionReason::NeedsCloudAdaptation, + "Cloud visibility indexing can exceed the test's fixed 500 ms retry window." +)] #[tokio::test] async fn client_list_open_closed_workflow_executions() { let wf_name = "client_list_open_closed_workflow_executions".to_owned(); @@ -128,6 +132,7 @@ async fn client_list_open_closed_workflow_executions() { assert!(passed); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::RequiresOssOnlyApis)] #[tokio::test] async fn client_create_namespace() { let client = Arc::new(get_integ_client(NAMESPACE.to_string(), None).await); @@ -177,12 +182,13 @@ async fn client_create_namespace() { #[tokio::test] async fn client_describe_namespace() { - let client = Arc::new(get_integ_client(NAMESPACE.to_string(), None).await); + let client = Arc::new(get_integ_client(integ_namespace(), None).await); + let namespace = client.namespace(); let namespace_result = WorkflowService::describe_namespace( &mut client.as_ref().clone(), DescribeNamespaceRequest { - namespace: NAMESPACE.to_owned(), + namespace: namespace.clone(), ..Default::default() } .into_request(), @@ -190,5 +196,5 @@ async fn client_describe_namespace() { .await .unwrap() .into_inner(); - assert_eq!(namespace_result.namespace_info.unwrap().name, NAMESPACE); + assert_eq!(namespace_result.namespace_info.unwrap().name, namespace); } diff --git a/crates/sdk-core/tests/integ_tests/worker_heartbeat_tests.rs b/crates/sdk-core/tests/integ_tests/worker_heartbeat_tests.rs index 72dad9ad9..370c4f3cb 100644 --- a/crates/sdk-core/tests/integ_tests/worker_heartbeat_tests.rs +++ b/crates/sdk-core/tests/integ_tests/worker_heartbeat_tests.rs @@ -42,10 +42,10 @@ use temporalio_macros::{activities, workflow, workflow_methods}; use temporalio_sdk::{ ActivityOptions, SyncWorkflowContext, WorkflowContext, WorkflowResult, activities::{ActivityContext, ActivityError}, + runtime::{AutoscalingOptions, PollerBehavior}, }; use temporalio_sdk_core::{ - CoreRuntime, PollerBehavior, ResourceBasedTuner, ResourceSlotOptions, RuntimeOptions, - TunerHolder, prost_dur, + CoreRuntime, ResourceBasedTuner, ResourceSlotOptions, RuntimeOptions, TunerHolder, prost_dur, }; use tokio::{sync::Notify, time::sleep}; use tonic::IntoRequest; @@ -188,7 +188,7 @@ async fn docker_worker_heartbeat_basic(#[values("otel", "prom", "no_metrics")] b let wf_name = format!("worker_heartbeat_basic_{backing}"); let mut starter = CoreWfStarter::new_with_runtime(&wf_name, rt); starter.sdk_config.max_cached_workflows = 5_usize; - starter.sdk_config.tuner = Arc::new(TunerHolder::fixed_size(5, 5, 100, 0)); + starter.set_core_tuner(Arc::new(TunerHolder::fixed_size(5, 5, 100, 0))); starter.set_core_cfg_mutator(|c| { c.plugins = vec![ PluginInfo { @@ -274,43 +274,46 @@ async fn docker_worker_heartbeat_basic(#[values("otel", "prom", "no_metrics")] b let heartbeat_time = AtomicCell::new(None); let test_fut = async { - // Give enough time to ensure heartbeat interval has been hit - tokio::time::sleep(Duration::from_millis(1500)).await; acts_started.notified().await; let client = starter.get_core_client().await; - let mut raw_client = client.clone(); - let workers_list = WorkflowService::list_workers( - &mut raw_client, - ListWorkersRequest { - namespace: client.namespace().to_owned(), - page_size: 100, - next_page_token: Vec::new(), - query: String::new(), - include_system_workers: false, - } - .into_request(), + let raw_client = client.clone(); + let heartbeat = eventually( + || { + let client = client.clone(); + async move { + let heartbeat = list_worker_heartbeats(&client, String::new()) + .await + .into_iter() + .find(|heartbeat| { + heartbeat.worker_instance_key == worker_instance_key.to_string() + }) + .ok_or_else(|| anyhow!("worker heartbeat has not been recorded"))?; + let workflow_tasks = heartbeat + .workflow_task_slots_info + .as_ref() + .map_or(0, |slots| slots.total_processed_tasks); + let activities = heartbeat + .activity_task_slots_info + .as_ref() + .map_or(0, |slots| slots.current_used_slots); + if workflow_tasks == 1 && activities == 1 { + Ok(heartbeat) + } else { + Err(anyhow!( + "Heartbeat not ready: workflow tasks={workflow_tasks}, activities={activities}" + )) + } + } + }, + Duration::from_secs(5), ) .await - .unwrap() - .into_inner(); - #[allow(deprecated)] - let worker_info = workers_list - .workers_info - .iter() - .find(|worker_info| { - if let Some(hb) = worker_info.worker_heartbeat.as_ref() { - hb.worker_instance_key == worker_instance_key.to_string() - } else { - false - } - }) - .unwrap(); - let heartbeat = worker_info.worker_heartbeat.as_ref().unwrap(); + .unwrap(); assert_eq!( heartbeat.worker_instance_key, worker_instance_key.to_string() ); - in_activity_checks(heartbeat, &start_time, &heartbeat_time); + in_activity_checks(&heartbeat, &start_time, &heartbeat_time); acts_done.notify_one(); // Poll until the heartbeat reflects shutdown with the second WFT processed. @@ -404,17 +407,21 @@ async fn docker_worker_heartbeat_tuner() { tuner .with_workflow_slots_options(ResourceSlotOptions::new(2, 10, Duration::from_millis(0))) .with_activity_slots_options(ResourceSlotOptions::new(5, 10, Duration::from_millis(50))); - starter.sdk_config.workflow_task_poller_behavior = Some(PollerBehavior::Autoscaling { - minimum: 1, - maximum: 200, - initial: 5, - }); - starter.sdk_config.nexus_task_poller_behavior = Some(PollerBehavior::Autoscaling { - minimum: 1, - maximum: 200, - initial: 5, - }); - starter.sdk_config.tuner = Arc::new(tuner); + starter.sdk_config.workflow_task_poller_behavior = Some(PollerBehavior::Autoscaling( + AutoscalingOptions::builder() + .minimum(1) + .maximum(200) + .initial(5) + .build(), + )); + starter.sdk_config.nexus_task_poller_behavior = Some(PollerBehavior::Autoscaling( + AutoscalingOptions::builder() + .minimum(1) + .maximum(200) + .initial(5) + .build(), + )); + starter.set_core_tuner(Arc::new(tuner)); starter.sdk_config.register_activities(StdActivities); #[workflow] @@ -670,7 +677,7 @@ async fn worker_heartbeat_sticky_cache_miss() { let wf_name = "worker_heartbeat_cache_miss"; let mut starter = new_no_metrics_starter(wf_name); starter.sdk_config.max_cached_workflows = 1_usize; - starter.sdk_config.tuner = Arc::new(TunerHolder::fixed_size(2, 10, 10, 10)); + starter.set_core_tuner(Arc::new(TunerHolder::fixed_size(2, 10, 10, 10))); struct StickyCacheActivities; #[activities] @@ -764,26 +771,24 @@ async fn worker_heartbeat_sticky_cache_miss() { HISTORY_WF2_ACTIVITY_STARTED.notified().await; HISTORY_WF1_ACTIVITY_FINISH.notify_one(); - let handle1 = WorkflowExecutionInfo { - namespace: client_for_orchestrator.namespace(), - workflow_id: wf1_id, - run_id: Some(wf1_run), - first_execution_run_id: None, - } - .bind_untyped(client_for_orchestrator.clone()); + let handle1 = WorkflowExecutionInfo::builder() + .namespace(client_for_orchestrator.namespace()) + .workflow_id(wf1_id) + .maybe_run_id(Some(wf1_run)) + .build() + .bind_untyped(client_for_orchestrator.clone()); handle1 .get_result(Default::default()) .await .expect("wf1 result"); HISTORY_WF2_ACTIVITY_FINISH.notify_one(); - let handle2 = WorkflowExecutionInfo { - namespace: client_for_orchestrator.namespace(), - workflow_id: wf2_id, - run_id: Some(wf2_run), - first_execution_run_id: None, - } - .bind_untyped(client_for_orchestrator.clone()); + let handle2 = WorkflowExecutionInfo::builder() + .namespace(client_for_orchestrator.namespace()) + .workflow_id(wf2_id) + .maybe_run_id(Some(wf2_run)) + .build() + .bind_untyped(client_for_orchestrator.clone()); handle2 .get_result(Default::default()) .await @@ -820,7 +825,7 @@ async fn worker_heartbeat_multiple_workers() { let runtime = CoreRuntime::new_assume_tokio(runtime_options).unwrap(); let mut starter = CoreWfStarter::new_with_runtime(wf_name, runtime); starter.sdk_config.max_cached_workflows = 5_usize; - starter.sdk_config.tuner = Arc::new(TunerHolder::fixed_size(5, 10, 10, 10)); + starter.set_core_tuner(Arc::new(TunerHolder::fixed_size(5, 10, 10, 10))); starter.sdk_config.register_activities(StdActivities); starter .sdk_config @@ -990,7 +995,7 @@ static WF_FAIL: Notify = Notify::const_new(); async fn worker_heartbeat_failure_metrics() { let wf_name = "worker_heartbeat_failure_metrics"; let mut starter = new_no_metrics_starter(wf_name); - starter.sdk_config.tuner = Arc::new(TunerHolder::fixed_size(10, 5, 10, 10)); + starter.set_core_tuner(Arc::new(TunerHolder::fixed_size(10, 5, 10, 10))); // This test uses tokio::sync::Notify from workflow code for test coordination. starter.sdk_config.detect_nondeterministic_futures = false; diff --git a/crates/sdk-core/tests/integ_tests/worker_tests.rs b/crates/sdk-core/tests/integ_tests/worker_tests.rs index 39dd07609..ac387a85f 100644 --- a/crates/sdk-core/tests/integ_tests/worker_tests.rs +++ b/crates/sdk-core/tests/integ_tests/worker_tests.rs @@ -2,7 +2,7 @@ use crate::{ common::{ CoreWfStarter, activity_functions::StdActivities, fake_grpc_server::fake_server, get_integ_runtime_options, get_integ_server_options, get_integ_telem_options, - integ_namespace, + integ_namespace, prom_metrics, }, shared_tests::{self, is_oversize_grpc_event}, }; @@ -61,13 +61,13 @@ use temporalio_sdk::{ ActivityOptions, LocalActivityOptions, WorkerOptions, WorkflowContext, WorkflowResult, activities::{ActivityContext, ActivityError}, interceptors::WorkerInterceptor, + runtime::PollerBehavior, }; use temporalio_sdk_core::{ - ActivitySlotKind, CoreRuntime, LocalActivitySlotKind, PollError, PollerBehavior, - ResourceBasedTuner, ResourceSlotOptions, SlotInfo, SlotInfoTrait, SlotMarkUsedContext, - SlotReleaseContext, SlotReservationContext, SlotSupplier, SlotSupplierPermit, TunerBuilder, - WorkerConfig, WorkerValidationError, WorkerVersioningStrategy, WorkflowSlotKind, init_worker, - prost_dur, + ActivitySlotKind, CoreRuntime, LocalActivitySlotKind, PollError, ResourceBasedTuner, + ResourceSlotOptions, SlotInfo, SlotInfoTrait, SlotMarkUsedContext, SlotReleaseContext, + SlotReservationContext, SlotSupplier, SlotSupplierPermit, TunerBuilder, WorkerConfig, + WorkerValidationError, WorkerVersioningStrategy, WorkflowSlotKind, init_worker, prost_dur, replay::{DEFAULT_WORKFLOW_TYPE, TestHistoryBuilder, canned_histories}, test_help::{ FakeWfResponses, MockPollCfg, ResponseType, build_mock_pollers, drain_pollers_and_shutdown, @@ -207,7 +207,7 @@ async fn resource_based_few_pollers_guarantees_non_sticky_poll() { // Set the limits to zero so it's essentially unwilling to hand out slots let mut tuner = ResourceBasedTuner::new(0.0, 0.0); tuner.with_workflow_slots_options(ResourceSlotOptions::new(2, 10, Duration::from_millis(0))); - starter.sdk_config.tuner = Arc::new(tuner); + starter.set_core_tuner(Arc::new(tuner)); starter .sdk_config .register_workflow::() @@ -232,7 +232,7 @@ async fn resource_based_few_pollers_guarantees_non_sticky_poll() { #[tokio::test] async fn oversize_grpc_message() { - use crate::common::{NAMESPACE, prom_metrics}; + let namespace = integ_namespace(); let wf_name = "oversize_grpc_message"; // Enable Prometheus metrics for this test and capture the address let (telemopts, addr, _aborter) = prom_metrics(None); @@ -292,7 +292,7 @@ async fn oversize_grpc_message() { if body.lines().any(|line| { line.starts_with("temporal_workflow_task_execution_failed{") && line.contains("failure_reason=\"GrpcMessageTooLarge\"") - && line.contains(&format!("namespace=\"{NAMESPACE}\"")) + && line.contains(&format!("namespace=\"{namespace}\"")) && line.contains("service_name=\"temporal-core-sdk\"") && line.contains(&format!("task_queue=\"{tq}\"")) && line.ends_with(" 1") @@ -313,6 +313,63 @@ async fn grpc_message_too_large_test() { shared_tests::grpc_message_too_large().await } +#[workflow] +#[derive(Default)] +struct PaginatedCompletionWf; + +#[workflow_methods] +impl PaginatedCompletionWf { + #[run] + async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { + // Schedule many activities in a single workflow task so the completion (~5 MiB across the + // commands) exceeds the per-page limit and must be paginated. Each input is well under the + // per-blob size limit, so it's the aggregate completion size that drives pagination. + let input = "a".repeat(400 * 1024); + let mut futs = vec![]; + for _ in 0..13 { + futs.push(ctx.execute_activity( + StdActivities::echo, + input.clone(), + ActivityOptions::start_to_close_timeout(Duration::from_secs(30)), + )); + } + temporalio_sdk::workflows::join_all(futs).await; + Ok(()) + } +} + +/// A workflow task completion too large for a single gRPC request is split into pages that the +/// server buffers and reassembles; the workflow then completes normally. Local-lane only: it needs +/// the dev server's `history.enableWorkflowTaskCompletionPagination` and a raised +/// `system.transactionSizeLimit`. +#[temporalio_macros::cloud_test_exclusion( + crate::CloudTestExclusionReason::RequiresLocalServer, + "Requires dev-server workflow-task pagination and transaction-size dynamic configuration." +)] +#[tokio::test] +async fn workflow_task_completion_pagination_test() { + let wf_name = "wft_completion_pagination"; + let mut starter = CoreWfStarter::new_cloud_or_local(wf_name, "") + .await + .unwrap(); + starter + .sdk_config + .register_workflow::() + .unwrap(); + starter.sdk_config.register_activities(StdActivities); + let mut worker = starter.worker().await; + let handle = worker + .submit_workflow( + PaginatedCompletionWf::run, + (), + starter.workflow_options.clone(), + ) + .await + .unwrap(); + worker.run_until_done().await.unwrap(); + handle.get_result(Default::default()).await.unwrap(); +} + // Serializes to between the default blob error limit (2 MiB) and the gRPC transport limit (4 MiB). const OVERSIZE_PAYLOAD_BYTES: usize = 3 * 1024 * 1024; @@ -682,6 +739,7 @@ async fn disabled_error_limit_lets_server_hard_fail() { ); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn activity_tasks_from_completion_reserve_slots() { let wf_id = "fake_wf_id"; @@ -787,6 +845,7 @@ async fn activity_tasks_from_completion_reserve_slots() { let mut worker = crate::common::TestWorker::new( temporalio_sdk::Worker::new_from_core_options(core.clone(), client_options, worker_options) .unwrap(), + core.clone(), ); // First poll for activities twice, occupying both slots @@ -862,6 +921,7 @@ async fn activity_tasks_from_completion_reserve_slots() { tokio::join!(run_fut, act_completer); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn max_wft_respected() { let total_wfs = 100; @@ -908,6 +968,7 @@ async fn max_wft_respected() { worker.run_until_done().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest] #[tokio::test] async fn history_length_with_fail_and_timeout( @@ -1043,6 +1104,7 @@ async fn history_length_with_fail_and_timeout( } #[allow(deprecated)] +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn sets_build_id_from_wft_complete() { let wfid = "fake_wf_id"; @@ -1220,7 +1282,7 @@ async fn test_custom_slot_supplier_simple() { tb.workflow_slot_supplier(wf_supplier.clone()); tb.activity_slot_supplier(activity_supplier.clone()); tb.local_activity_slot_supplier(local_activity_supplier.clone()); - starter.sdk_config.tuner = Arc::new(tb.build()); + starter.set_core_tuner(Arc::new(tb.build())); starter .sdk_config .register_workflow::() @@ -1400,6 +1462,7 @@ async fn test_custom_slot_supplier_simple() { ); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn shutdown_worker_not_retried() { let shutdown_call_count = Arc::new(AtomicU8::new(0)); @@ -1429,6 +1492,7 @@ async fn shutdown_worker_not_retried() { assert_eq!(shutdown_call_count.load(Ordering::Relaxed), 1); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[test] fn test_default_build_id() { let o = WorkerOptions::new("task_queue").build(); @@ -1436,6 +1500,10 @@ fn test_default_build_id() { assert_ne!(o.deployment_options.version.build_id, "undetermined"); } +#[temporalio_macros::cloud_test_exclusion( + crate::CloudTestExclusionReason::NeedsCloudAdaptation, + "Uses new_cloud_or_local, which treats envconfig as local and calls a cluster-info RPC unavailable to Cloud namespace credentials." +)] #[tokio::test] async fn shutdown_during_active_timer_activity_workflows() { shared_tests::shutdown_during_active_timer_activity_workflows().await diff --git a/crates/sdk-core/tests/integ_tests/worker_versioning_tests.rs b/crates/sdk-core/tests/integ_tests/worker_versioning_tests.rs index 735ee6348..f7cc90114 100644 --- a/crates/sdk-core/tests/integ_tests/worker_versioning_tests.rs +++ b/crates/sdk-core/tests/integ_tests/worker_versioning_tests.rs @@ -38,10 +38,10 @@ async fn sets_deployment_info_on_task_responses(#[values(true, false)] use_defau let wf_type = "sets_deployment_info_on_task_responses"; let mut starter = CoreWfStarter::new(wf_type); let deploy_name = format!("deployment-{}", starter.get_task_queue()); - let version = WorkerDeploymentVersion { - deployment_name: deploy_name.clone(), - build_id: "1.0".to_string(), - }; + let version = WorkerDeploymentVersion::builder() + .deployment_name(deploy_name.clone()) + .build_id("1.0".to_string()) + .build(); starter.sdk_config.deployment_options = WorkerDeploymentOptions::new(version.clone()) .use_worker_versioning(true) .default_versioning_behavior(VersioningBehavior::AutoUpgrade) @@ -173,10 +173,12 @@ async fn activity_has_deployment_stamp() { let wf_name = "activity_has_deployment_stamp"; let mut starter = CoreWfStarter::new(wf_name); let deploy_name = format!("deployment-{}", starter.get_task_queue()); - starter.sdk_config.deployment_options = WorkerDeploymentOptions::new(WorkerDeploymentVersion { - deployment_name: deploy_name.clone(), - build_id: "1.0".to_string(), - }) + starter.sdk_config.deployment_options = WorkerDeploymentOptions::new( + WorkerDeploymentVersion::builder() + .deployment_name(deploy_name.clone()) + .build_id("1.0".to_string()) + .build(), + ) .use_worker_versioning(true) .default_versioning_behavior(VersioningBehavior::AutoUpgrade) .build(); @@ -268,10 +270,12 @@ async fn versioning_off_with_custom_build_id() { let wf_type = "versioning_off_with_custom_build_id"; let mut starter = CoreWfStarter::new(wf_type); let build_id = "my-custom-build-id-1.0"; - starter.sdk_config.deployment_options = WorkerDeploymentOptions::new(WorkerDeploymentVersion { - deployment_name: format!("deployment-{}", starter.get_task_queue()), - build_id: build_id.to_string(), - }) + starter.sdk_config.deployment_options = WorkerDeploymentOptions::new( + WorkerDeploymentVersion::builder() + .deployment_name(format!("deployment-{}", starter.get_task_queue())) + .build_id(build_id.to_string()) + .build(), + ) .build(); starter.set_core_task_types(WorkerTaskTypes::workflow_only()); let core = starter.get_core_worker().await; @@ -360,14 +364,14 @@ async fn continue_as_new_auto_upgrade_uses_current_deployment_version() { let wf_type = "continue_as_new_auto_upgrade_uses_current_deployment_version"; let mut starter = CoreWfStarter::new(wf_type); let deploy_name = format!("deployment-{}", starter.get_task_queue()); - let v1 = WorkerDeploymentVersion { - deployment_name: deploy_name.clone(), - build_id: "1.0".to_string(), - }; - let v2 = WorkerDeploymentVersion { - deployment_name: deploy_name.clone(), - build_id: "2.0".to_string(), - }; + let v1 = WorkerDeploymentVersion::builder() + .deployment_name(deploy_name.clone()) + .build_id("1.0".to_string()) + .build(); + let v2 = WorkerDeploymentVersion::builder() + .deployment_name(deploy_name.clone()) + .build_id("2.0".to_string()) + .build(); starter.sdk_config.deployment_options = versioned_worker_options(v1.clone()); let mut starter2 = starter.clone_no_worker(); starter2.sdk_config.deployment_options = versioned_worker_options(v2.clone()); @@ -487,14 +491,14 @@ async fn continue_as_new_use_ramping_version_uses_ramping_deployment_version() { let wf_type = "continue_as_new_use_ramping_version_uses_ramping_deployment_version"; let mut starter = CoreWfStarter::new(wf_type); let deploy_name = format!("deployment-{}", starter.get_task_queue()); - let v1 = WorkerDeploymentVersion { - deployment_name: deploy_name.clone(), - build_id: "1.0".to_string(), - }; - let v2 = WorkerDeploymentVersion { - deployment_name: deploy_name.clone(), - build_id: "2.0".to_string(), - }; + let v1 = WorkerDeploymentVersion::builder() + .deployment_name(deploy_name.clone()) + .build_id("1.0".to_string()) + .build(); + let v2 = WorkerDeploymentVersion::builder() + .deployment_name(deploy_name.clone()) + .build_id("2.0".to_string()) + .build(); starter.sdk_config.deployment_options = versioned_worker_options(v1.clone()); let mut starter2 = starter.clone_no_worker(); starter2.sdk_config.deployment_options = versioned_worker_options(v2.clone()); diff --git a/crates/sdk-core/tests/integ_tests/workflow_client_tests.rs b/crates/sdk-core/tests/integ_tests/workflow_client_tests.rs index f16aabf98..82a23e660 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_client_tests.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_client_tests.rs @@ -13,7 +13,7 @@ use temporalio_client::{ WorkflowCountOptions, WorkflowListOptions, WorkflowStartOptions, WorkflowTerminateOptions, errors::WorkflowStartError, }; -use temporalio_common::data_converters::RawValue; +use temporalio_common::{MemoValues, data_converters::RawValue}; use temporalio_macros::{workflow, workflow_methods}; use temporalio_sdk::{WorkflowContext, WorkflowResult}; @@ -250,3 +250,40 @@ async fn already_started_error_contains_run_id() { .await .unwrap(); } + +#[tokio::test] +async fn start_workflow_with_memo() { + let test_name = "start_workflow_with_memo"; + let mut starter = CoreWfStarter::new(test_name); + let client = starter.get_core_client().await; + let task_queue = starter.get_task_queue().to_owned(); + let wf_id = format!("{test_name}_{}", rand_6_chars()); + + let mut memo = MemoValues::new(); + memo.insert("memo-key", "memo-value".to_string()) + .insert("other-key", 42_u32); + + let handle = client + .start_workflow( + UntypedWorkflow::new(test_name), + RawValue::empty(), + WorkflowStartOptions::new(task_queue, wf_id) + .memo(memo) + .build(), + ) + .await + .unwrap(); + + let desc = handle.describe(Default::default()).await.unwrap(); + let memo = desc.memo(); + assert_eq!( + memo.get::("memo-key").unwrap(), + Some("memo-value".to_string()) + ); + assert_eq!(memo.get::("other-key").unwrap(), Some(42)); + + handle + .terminate(WorkflowTerminateOptions::default()) + .await + .unwrap(); +} diff --git a/crates/sdk-core/tests/integ_tests/workflow_replayer_tests.rs b/crates/sdk-core/tests/integ_tests/workflow_replayer_tests.rs index c50a81527..a1432eaaa 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_replayer_tests.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_replayer_tests.rs @@ -11,7 +11,7 @@ use temporalio_client::{ PluginError, WorkflowHistory, WorkflowQueryOptions, WorkflowStartOptions, WorkflowTerminateOptions, errors::WorkflowGetResultError, }; -use temporalio_common::protos::temporal::api::enums::v1::EventType; +use temporalio_common::protos::temporal::api::{enums::v1::EventType, history::v1::History}; use temporalio_macros::{activities, workflow, workflow_methods}; use temporalio_sdk::{ ActivityOptions, ApplicationFailure, SimplePlugin, WorkerPlugin, WorkerRunError, @@ -137,21 +137,29 @@ async fn workflow_replayer_replays_completed_workflow() { handle.get_result(Default::default()).await.unwrap(), "Hello, Temporal!" ); - let history = handle.fetch_history(Default::default()).await.unwrap(); - let history_from_json = WorkflowHistory::from_json(&history.to_json().unwrap()).unwrap(); - assert_eq!(history_from_json.workflow_id(), history.workflow_id()); + let history = handle.fetch_history(Default::default()); + let workflow_id = history.workflow_id().map(str::to_owned); + let history_json = history.to_json().await.unwrap(); + let history_from_json = WorkflowHistory::from_json(&history_json).unwrap(); + assert_eq!(history_from_json.workflow_id(), workflow_id.as_deref()); let replayer = replayer(); - replayer - .replay_workflow(history_from_json.clone()) - .await - .unwrap(); + replayer.replay_workflow(history_from_json).await.unwrap(); let results = replayer - .replay_workflows([history_from_json.clone(), history_from_json]) + .replay_workflows([ + WorkflowHistory::from_json(&history_json).unwrap(), + WorkflowHistory::from_json(&history_json).unwrap(), + ]) .await .unwrap(); assert_eq!(results.len(), 2); assert!(results.iter().all(|result| result.replay_failure.is_none())); + let cloned_result = results[0].clone(); + assert_eq!( + cloned_result.history.workflow_id(), + Some(starter.get_wf_id()) + ); + assert!(!cloned_result.history.events().is_empty()); } #[tokio::test] @@ -192,7 +200,14 @@ async fn workflow_replayer_replays_incomplete_workflow() { ) .await .unwrap(); - let history = handle.fetch_history(Default::default()).await.unwrap(); + let history: WorkflowHistory = History { + events: handle + .fetch_history(Default::default()) + .into_events() + .await + .unwrap(), + } + .into(); handle .terminate(WorkflowTerminateOptions::default()) .await @@ -230,7 +245,7 @@ async fn workflow_replayer_replays_failed_workflow() { handle.get_result(Default::default()).await, Err(WorkflowGetResultError::Failed(_)) )); - let history = handle.fetch_history(Default::default()).await.unwrap(); + let history = handle.fetch_history(Default::default()); replayer().replay_workflow(history).await.unwrap(); } @@ -256,16 +271,30 @@ async fn workflow_replayer_reports_nondeterminism() { .unwrap(); worker.run_until_done().await.unwrap(); - let history = handle.fetch_history(Default::default()).await.unwrap(); + let events = handle + .fetch_history(Default::default()) + .into_events() + .await + .unwrap(); let replayer = replayer(); assert!(matches!( - replayer.replay_workflow(history.clone()).await, + replayer + .replay_workflow( + History { + events: events.clone(), + } + .into() + ) + .await, Err(WorkflowReplayError::Replay( WorkflowReplayFailure::Nondeterminism { .. } )) )); - let results = replayer.replay_workflows([history]).await.unwrap(); + let results = replayer + .replay_workflows([History { events }.into()]) + .await + .unwrap(); assert!(matches!( results[0].replay_failure, Some(WorkflowReplayFailure::Nondeterminism { .. }) @@ -295,12 +324,15 @@ async fn workflow_replayer_replays_history_with_workflow_task_failure() { let fetch_failed_history = async { let history = eventually( || async { - let history = handle.fetch_history(Default::default()).await.unwrap(); - history - .events() + let events = handle + .fetch_history(Default::default()) + .into_events() + .await + .unwrap(); + events .iter() .any(|event| event.event_type() == EventType::WorkflowTaskFailed) - .then_some(history) + .then_some(History { events }.into()) .ok_or("workflow task failure not yet recorded") }, Duration::from_secs(10), @@ -353,14 +385,8 @@ async fn workflow_replayer_returns_ordered_results_for_multiple_histories() { .unwrap(); worker.run_until_done().await.unwrap(); - let successful_history = successful_handle - .fetch_history(Default::default()) - .await - .unwrap(); - let nondeterministic_history = nondeterministic_handle - .fetch_history(Default::default()) - .await - .unwrap(); + let successful_history = successful_handle.fetch_history(Default::default()); + let nondeterministic_history = nondeterministic_handle.fetch_history(Default::default()); let results = replayer() .replay_workflows([successful_history, nondeterministic_history]) @@ -441,7 +467,11 @@ async fn workflow_replayer_applies_plugins() { .await .unwrap(); worker.run_until_done().await.unwrap(); - let history = handle.fetch_history(Default::default()).await.unwrap(); + let events = handle + .fetch_history(Default::default()) + .into_events() + .await + .unwrap(); let configure_calls = Arc::new(AtomicUsize::new(0)); let replayer = WorkflowReplayer::new( @@ -460,7 +490,15 @@ async fn workflow_replayer_applies_plugins() { .count(), 1 ); - replayer.replay_workflow(history.clone()).await.unwrap(); + replayer + .replay_workflow( + History { + events: events.clone(), + } + .into(), + ) + .await + .unwrap(); assert_eq!(configure_calls.load(Ordering::Relaxed), 1); let run_calls = Arc::new(AtomicUsize::new(0)); @@ -477,7 +515,16 @@ async fn workflow_replayer_applies_plugins() { ) .unwrap(); replayer - .replay_workflows([history.clone(), history.clone()]) + .replay_workflows([ + History { + events: events.clone(), + } + .into(), + History { + events: events.clone(), + } + .into(), + ]) .await .unwrap(); assert_eq!(run_calls.load(Ordering::Relaxed), 0); @@ -495,5 +542,8 @@ async fn workflow_replayer_applies_plugins() { .build(), ) .unwrap(); - replayer.replay_workflow(history).await.unwrap(); + replayer + .replay_workflow(History { events }.into()) + .await + .unwrap(); } diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests.rs b/crates/sdk-core/tests/integ_tests/workflow_tests.rs index f2cefa7b4..f3247699a 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests.rs @@ -6,6 +6,7 @@ mod client_interactions; mod continue_as_new; mod determinism; mod eager; +mod event_groups; mod interceptors; mod local_activities; mod modify_wf_properties; @@ -67,9 +68,10 @@ use temporalio_macros::{workflow, workflow_methods}; use temporalio_sdk::{ ActivityOptions, LocalActivityOptions, TimerOptions, WorkflowContext, WorkflowResult, interceptors::WorkerInterceptor, + runtime::{PollerBehavior, WorkflowErrorType}, }; use temporalio_sdk_core::{ - CoreRuntime, PollError, PollerBehavior, TunerHolder, WorkflowErrorType, prost_dur, + CoreRuntime, PollError, TunerHolder, prost_dur, replay::{DEFAULT_WORKFLOW_TYPE, HistoryForReplay, canned_histories}, test_help::{ MockPollCfg, WorkerTestHelpers, drain_pollers_and_shutdown, schedule_activity_cmd, @@ -240,13 +242,12 @@ async fn signal_workflow() { .unwrap(); // Send the signals to the server - let handle = WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: workflow_id.clone(), - run_id: Some(res.run_id.clone()), - first_execution_run_id: None, - } - .bind_untyped(client.clone()); + let handle = WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(workflow_id.clone()) + .maybe_run_id(Some(res.run_id.clone())) + .build() + .bind_untyped(client.clone()); handle .signal( UntypedSignal::new(signal_id_1), @@ -341,20 +342,19 @@ async fn signal_workflow_signal_not_handled_on_workflow_completion() { // Send the signal to the server let sig_client = starter.get_core_client().await; - WorkflowExecutionInfo { - namespace: sig_client.namespace(), - workflow_id: workflow_id.clone(), - run_id: Some(res.run_id.clone()), - first_execution_run_id: None, - } - .bind_untyped(sig_client.clone()) - .signal( - UntypedSignal::new(signal_id_1), - RawValue::empty(), - WorkflowSignalOptions::default(), - ) - .await - .unwrap(); + WorkflowExecutionInfo::builder() + .namespace(sig_client.namespace()) + .workflow_id(workflow_id.clone()) + .maybe_run_id(Some(res.run_id.clone())) + .build() + .bind_untyped(sig_client.clone()) + .signal( + UntypedSignal::new(signal_id_1), + RawValue::empty(), + WorkflowSignalOptions::default(), + ) + .await + .unwrap(); // Send completion - not having seen a poll response with a signal in it yet (unhandled // command error will be logged as a warning and an eviction will be issued) @@ -389,7 +389,7 @@ async fn wft_timeout_doesnt_create_unsolvable_autocomplete() { let mut wf_starter = CoreWfStarter::new("wft_timeout_doesnt_create_unsolvable_autocomplete"); // Test needs eviction on and a short timeout wf_starter.sdk_config.max_cached_workflows = 0_usize; - wf_starter.sdk_config.tuner = Arc::new(TunerHolder::fixed_size(1, 1, 1, 1)); + wf_starter.set_core_tuner(Arc::new(TunerHolder::fixed_size(1, 1, 1, 1))); wf_starter.sdk_config.workflow_task_poller_behavior = Some(PollerBehavior::SimpleMaximum(1_usize)); wf_starter.workflow_options.task_timeout = Some(Duration::from_secs(1)); @@ -423,13 +423,12 @@ async fn wft_timeout_doesnt_create_unsolvable_autocomplete() { // Before polling for a task again, we start and complete the activity and send the // corresponding signals. let ac_task = core.poll_activity_task().await.unwrap(); - let handle = WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: wf_id.to_string(), - run_id: Some(wf_task.run_id.clone()), - first_execution_run_id: None, - } - .bind_untyped(client.clone()); + let handle = WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(wf_id.to_string()) + .maybe_run_id(Some(wf_task.run_id.clone())) + .build() + .bind_untyped(client.clone()); // Send the signals to the server & resolve activity -- sometimes this happens too fast sleep(Duration::from_millis(200)).await; handle @@ -519,7 +518,7 @@ impl SlowCompletesWf { async fn slow_completes_with_small_cache() { let wf_name = "slow_completes_with_small_cache"; let mut starter = CoreWfStarter::new(wf_name); - starter.sdk_config.tuner = Arc::new(TunerHolder::fixed_size(5, 10, 1, 1)); + starter.set_core_tuner(Arc::new(TunerHolder::fixed_size(5, 10, 1, 1))); starter.sdk_config.max_cached_workflows = 5_usize; starter.sdk_config.register_activities(StdActivities); starter @@ -560,16 +559,20 @@ async fn deployment_version_correct_in_wf_info(#[values(true, false)] use_only_b let wf_type = "deployment_version_correct_in_wf_info"; let mut starter = CoreWfStarter::new(wf_type); starter.sdk_config.deployment_options = if use_only_build_id { - WorkerDeploymentOptions::new(WorkerDeploymentVersion { - deployment_name: "".to_string(), - build_id: "1.0".to_string(), - }) + WorkerDeploymentOptions::new( + WorkerDeploymentVersion::builder() + .deployment_name("".to_string()) + .build_id("1.0".to_string()) + .build(), + ) .build() } else { - WorkerDeploymentOptions::new(WorkerDeploymentVersion { - deployment_name: "deployment-1".to_string(), - build_id: "1.0".to_string(), - }) + WorkerDeploymentOptions::new( + WorkerDeploymentVersion::builder() + .deployment_name("deployment-1".to_string()) + .build_id("1.0".to_string()) + .build(), + ) .build() }; starter.set_core_task_types(WorkerTaskTypes::workflow_only()); @@ -602,13 +605,12 @@ async fn deployment_version_correct_in_wf_info(#[values(true, false)] use_only_b .unwrap(); // Ensure a query on first wft also sees the correct id - let query_handle = WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: workflow_id.clone(), - run_id: Some(res.run_id.clone()), - first_execution_run_id: None, - } - .bind_untyped(client.clone()); + let query_handle = WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(workflow_id.clone()) + .maybe_run_id(Some(res.run_id.clone())) + .build() + .bind_untyped(client.clone()); let query_fut = async { query_handle .query( @@ -676,16 +678,20 @@ async fn deployment_version_correct_in_wf_info(#[values(true, false)] use_only_b let mut starter = starter.clone_no_worker(); starter.sdk_config.deployment_options = if use_only_build_id { - WorkerDeploymentOptions::new(WorkerDeploymentVersion { - deployment_name: "".to_string(), - build_id: "2.0".to_string(), - }) + WorkerDeploymentOptions::new( + WorkerDeploymentVersion::builder() + .deployment_name("".to_string()) + .build_id("2.0".to_string()) + .build(), + ) .build() } else { - WorkerDeploymentOptions::new(WorkerDeploymentVersion { - deployment_name: "deployment-1".to_string(), - build_id: "2.0".to_string(), - }) + WorkerDeploymentOptions::new( + WorkerDeploymentVersion::builder() + .deployment_name("deployment-1".to_string()) + .build_id("2.0".to_string()) + .build(), + ) .build() }; @@ -912,8 +918,9 @@ async fn nondeterminism_errors_fail_workflow_when_configured_to( worker.run_until_done().await.unwrap(); let body = metrics_tests::get_text(format!("http://{addr}/metrics")).await; + let namespace = client.namespace(); let match_this = format!( - "temporal_workflow_failed{{namespace=\"default\",\ + "temporal_workflow_failed{{namespace=\"{namespace}\",\ service_name=\"temporal-core-sdk\",\ task_queue=\"{wf_id}\",workflow_type=\"{wf_name}\"}} 1" ); @@ -1037,6 +1044,7 @@ async fn history_out_of_order_on_restart() { assert_matches!(res, Err(WorkflowGetResultError::Failed(_))); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn pass_timer_summary_to_metadata() { let t = canned_histories::single_timer("1"); diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/activities.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/activities.rs index 00668a797..879311e85 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/activities.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/activities.rs @@ -58,9 +58,10 @@ use temporalio_sdk::{ ActivityInboundInterceptor, ExecuteActivityInput, ExecuteActivityOutput, ExecuteActivityResult, Next, }, + runtime::{AutoscalingOptions, PollerBehavior}, }; use temporalio_sdk_core::{ - PollerBehavior, prost_dur, + prost_dur, replay::{DEFAULT_ACTIVITY_TYPE, DEFAULT_WORKFLOW_TYPE, TestHistoryBuilder, canned_histories}, test_help::{ MockPollCfg, ResponseType, WorkerTestHelpers, drain_pollers_and_shutdown, @@ -135,7 +136,7 @@ struct ActivityInterceptorRecord { interceptor: &'static str, phase: ActivityInterceptorPhase, activity_type: String, - workflow_type: String, + workflow_type: Option, is_local: bool, input: Option, output: Option, @@ -208,6 +209,7 @@ fn activity_execution_result_status(result: &ExecuteActivityResult) -> &'static Err(ActivityError::Application(_)) => "failed", Err(ActivityError::Cancelled { .. }) => "cancelled", Err(ActivityError::WillCompleteAsync) => "will_complete_async", + Err(_) => "unknown", } } @@ -352,9 +354,12 @@ async fn propagated_activity_input_conversion_failure_fails_workflow_task() { worker.run_until_done().await.unwrap(); handle.get_result(Default::default()).await.unwrap(); - let history = handle.fetch_history(Default::default()).await.unwrap(); + let history = handle + .fetch_history(Default::default()) + .into_events() + .await + .unwrap(); let workflow_task_failures: Vec<_> = history - .events() .iter() .filter(|event| event.event_type() == EventType::WorkflowTaskFailed) .collect(); @@ -495,7 +500,7 @@ async fn activity_interceptor_wraps_activity_execution() { interceptor: "outer", phase: ActivityInterceptorPhase::Before, activity_type: "StdActivities::echo".to_owned(), - workflow_type: wf_name.to_owned(), + workflow_type: Some(wf_name.to_owned()), is_local: false, input: Some(input.clone()), output: None, @@ -505,7 +510,7 @@ async fn activity_interceptor_wraps_activity_execution() { interceptor: "inner", phase: ActivityInterceptorPhase::Before, activity_type: "StdActivities::echo".to_owned(), - workflow_type: wf_name.to_owned(), + workflow_type: Some(wf_name.to_owned()), is_local: false, input: Some(input.clone()), output: None, @@ -515,7 +520,7 @@ async fn activity_interceptor_wraps_activity_execution() { interceptor: "inner", phase: ActivityInterceptorPhase::After, activity_type: "StdActivities::echo".to_owned(), - workflow_type: wf_name.to_owned(), + workflow_type: Some(wf_name.to_owned()), is_local: false, input: None, output: Some(input.clone()), @@ -525,7 +530,7 @@ async fn activity_interceptor_wraps_activity_execution() { interceptor: "outer", phase: ActivityInterceptorPhase::After, activity_type: "StdActivities::echo".to_owned(), - workflow_type: wf_name.to_owned(), + workflow_type: Some(wf_name.to_owned()), is_local: false, input: None, output: Some(input.clone()), @@ -574,7 +579,7 @@ async fn activity_interceptor_wraps_local_activity_execution() { interceptor: "local", phase: ActivityInterceptorPhase::Before, activity_type: "StdActivities::echo".to_owned(), - workflow_type: wf_name.to_owned(), + workflow_type: Some(wf_name.to_owned()), is_local: true, input: Some(input.clone()), output: None, @@ -584,7 +589,7 @@ async fn activity_interceptor_wraps_local_activity_execution() { interceptor: "local", phase: ActivityInterceptorPhase::After, activity_type: "StdActivities::echo".to_owned(), - workflow_type: wf_name.to_owned(), + workflow_type: Some(wf_name.to_owned()), is_local: true, input: None, output: Some(input.clone()), @@ -710,7 +715,7 @@ async fn activity_interceptor_observes_activity_error() { interceptor: "failure", phase: ActivityInterceptorPhase::Before, activity_type: "FailingActivities::fail".to_owned(), - workflow_type: wf_name.to_owned(), + workflow_type: Some(wf_name.to_owned()), is_local: false, input: Some(input), output: None, @@ -720,7 +725,7 @@ async fn activity_interceptor_observes_activity_error() { interceptor: "failure", phase: ActivityInterceptorPhase::After, activity_type: "FailingActivities::fail".to_owned(), - workflow_type: wf_name.to_owned(), + workflow_type: Some(wf_name.to_owned()), is_local: false, input: None, output: None, @@ -806,7 +811,7 @@ async fn activity_interceptor_observes_activity_panic() { interceptor: "panic", phase: ActivityInterceptorPhase::Before, activity_type: "PanickingActivities::panic_activity".to_owned(), - workflow_type: wf_name.to_owned(), + workflow_type: Some(wf_name.to_owned()), is_local: false, input: Some(input), output: None, @@ -816,7 +821,7 @@ async fn activity_interceptor_observes_activity_panic() { interceptor: "panic", phase: ActivityInterceptorPhase::After, activity_type: "PanickingActivities::panic_activity".to_owned(), - workflow_type: wf_name.to_owned(), + workflow_type: Some(wf_name.to_owned()), is_local: false, input: None, output: None, @@ -1076,7 +1081,7 @@ async fn activity_non_retryable_failure() { variant: Some(workflow_activation_job::Variant::ResolveActivity( ResolveActivity {seq, result: Some(ActivityResolution{ status: Some(act_res::Status::Failed(activity_result::Failure{ - failure: Some(f), + failure: Some(f), .. }))}),..} )), }, @@ -1143,7 +1148,7 @@ async fn activity_non_retryable_failure_with_error() { variant: Some(workflow_activation_job::Variant::ResolveActivity( ResolveActivity {seq, result: Some(ActivityResolution{ status: Some(act_res::Status::Failed(activity_result::Failure{ - failure: Some(f), + failure: Some(f), .. }))}),..} )), }, @@ -1499,7 +1504,7 @@ async fn started_activity_timeout() { result: Some(ActivityResolution{ status: Some( act_res::Status::Failed( - activity_result::Failure{failure: Some(_)} + activity_result::Failure{failure: Some(_), ..} ) ), .. @@ -2130,16 +2135,20 @@ async fn activity_can_be_cancelled_by_local_timeout() { async fn long_activity_timeout_repro() { let wf_name = "long_activity_timeout_repro"; let mut starter = CoreWfStarter::new(wf_name); - starter.sdk_config.workflow_task_poller_behavior = Some(PollerBehavior::Autoscaling { - minimum: 1, - maximum: 10, - initial: 5, - }); - starter.sdk_config.activity_task_poller_behavior = Some(PollerBehavior::Autoscaling { - minimum: 1, - maximum: 10, - initial: 5, - }); + starter.sdk_config.workflow_task_poller_behavior = Some(PollerBehavior::Autoscaling( + AutoscalingOptions::builder() + .minimum(1) + .maximum(10) + .initial(5) + .build(), + )); + starter.sdk_config.activity_task_poller_behavior = Some(PollerBehavior::Autoscaling( + AutoscalingOptions::builder() + .minimum(1) + .maximum(10) + .initial(5) + .build(), + )); starter .set_core_cfg_mutator(|m| m.local_timeout_buffer_for_activities = Duration::from_secs(0)); starter.sdk_config.register_activities(StdActivities); @@ -2185,6 +2194,7 @@ async fn long_activity_timeout_repro() { worker.run_until_done().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn pass_activity_summary_to_metadata() { let t = canned_histories::single_activity("1"); @@ -2255,6 +2265,7 @@ async fn pass_activity_summary_to_metadata() { worker.run_until_done().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest(hist_batches, case::incremental(&[1, 2, 3, 4]), case::replay(&[4]))] #[tokio::test] async fn abandoned_activities_ignore_start_and_complete(hist_batches: &'static [usize]) { @@ -2358,6 +2369,7 @@ impl ImmediateActivityCancelationWorkflow { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn immediate_activity_cancelation() { let mut t = TestHistoryBuilder::default(); diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/cancel_external.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/cancel_external.rs index 0c70d31ea..8ade47ff4 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/cancel_external.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/cancel_external.rs @@ -5,7 +5,9 @@ use temporalio_common::protos::{ temporal::api::enums::v1::{CommandType, EventType}, }; use temporalio_macros::{workflow, workflow_methods}; -use temporalio_sdk::{ApplicationFailure, WorkflowContext, WorkflowResult}; +use temporalio_sdk::{ + ApplicationFailure, CancelExternalWorkflowError, WorkflowContext, WorkflowResult, +}; use temporalio_sdk_core::{ replay::{DEFAULT_WORKFLOW_TYPE, TestHistoryBuilder}, test_help::MockPollCfg, @@ -20,10 +22,15 @@ impl CancelSender { #[run] async fn run( ctx: &mut WorkflowContext, - (run_id, workflow_id): (String, String), + (run_id, workflow_id, expect_not_found): (String, String, bool), ) -> WorkflowResult<()> { let handle = ctx.external_workflow(workflow_id, Some(run_id)); - handle.cancel(Some("cancel-reason".into())).await.unwrap(); + let result = handle.cancel(Some("cancel-reason".into())).await; + if expect_not_found { + assert_matches!(result, Err(CancelExternalWorkflowError::NotFound(_))); + } else { + result.unwrap(); + } Ok(()) } } @@ -68,7 +75,7 @@ async fn sends_cancel_to_other_wf() { let sender_handle = worker .submit_workflow( CancelSender::run, - (receiver_run_id.to_owned(), receiver_wfid.to_owned()), + (receiver_run_id.to_owned(), receiver_wfid.to_owned(), false), WorkflowStartOptions::new(task_queue, "sends-cancel-sender").build(), ) .await @@ -93,6 +100,32 @@ async fn sends_cancel_to_other_wf() { ); } +#[tokio::test] +async fn cancel_missing_external_wf_returns_not_found() { + let wf_name = "cancel_missing_external_wf_returns_not_found"; + let mut starter = CoreWfStarter::new(wf_name); + starter + .sdk_config + .register_workflow::() + .unwrap(); + let mut worker = starter.worker().await; + + let task_queue = starter.get_task_queue().to_owned(); + let handle = worker + .submit_workflow( + CancelSender::run, + (uuid::Uuid::new_v4().to_string(), wf_name.to_owned(), true), + WorkflowStartOptions::new(task_queue, wf_name).build(), + ) + .await + .unwrap(); + worker.run_until_done().await.unwrap(); + handle + .get_result(Default::default()) + .await + .expect("workflow should observe a not-found cancellation error"); +} + #[workflow] #[derive(Default)] struct CancelSenderCanned; @@ -103,14 +136,17 @@ impl CancelSenderCanned { async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { let handle = ctx.external_workflow("fake_wid", Some("fake_rid".into())); let res = handle.cancel(None).await; - if res.is_err() { - Err(ApplicationFailure::new("Cancel fail!").into()) - } else { - Ok(()) + match res { + Err(CancelExternalWorkflowError::Failed(_)) => { + Err(ApplicationFailure::new("Cancel fail!").into()) + } + Err(error) => panic!("unexpected external cancellation error: {error}"), + Ok(()) => Ok(()), } } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest::rstest] #[case::succeeds(false)] #[case::fails(true)] diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/cancel_wf.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/cancel_wf.rs index 9e0f40cf9..f14d3b464 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/cancel_wf.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/cancel_wf.rs @@ -38,6 +38,67 @@ impl CancelledWf { } } +#[derive(Debug, PartialEq, serde::Deserialize, serde::Serialize)] +struct CancellationDetails { + reason: String, +} + +#[workflow] +#[derive(Default)] +struct CancelledWithDetailsWf; + +#[workflow_methods] +impl CancelledWithDetailsWf { + #[run] + async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { + ctx.cancelled().await; + Err(WorkflowTermination::cancelled_with_details( + CancellationDetails { + reason: "contract expired".to_owned(), + }, + )) + } +} + +#[tokio::test] +async fn workflow_cancellation_records_details() { + let wf_name = "workflow_cancellation_records_details"; + let mut starter = CoreWfStarter::new(wf_name); + starter + .sdk_config + .register_workflow::() + .unwrap(); + let mut worker = starter.worker().await; + let task_queue = starter.get_task_queue().to_owned(); + let wf_handle = worker + .submit_workflow( + CancelledWithDetailsWf::run, + (), + WorkflowStartOptions::new(task_queue, wf_name).build(), + ) + .await + .unwrap(); + + let (cancel_result, worker_result) = tokio::join!( + wf_handle.cancel(WorkflowCancelOptions::default()), + worker.run_until_done() + ); + cancel_result.unwrap(); + worker_result.unwrap(); + + let Err(WorkflowGetResultError::Cancelled { details }) = + wf_handle.get_result(Default::default()).await + else { + panic!("workflow should complete as cancelled"); + }; + assert_eq!( + details.deserialize::().unwrap(), + CancellationDetails { + reason: "contract expired".to_owned(), + } + ); +} + #[tokio::test] async fn cancel_during_timer() { let wf_name = "cancel_during_timer"; @@ -155,13 +216,88 @@ impl CancellationPropagationActivities { let mut heartbeat = tokio::time::interval(Duration::from_millis(100)); loop { tokio::select! { - _ = ctx.cancelled() => return Err(ActivityError::cancelled()), + _ = ctx.cancelled() => return Err(ActivityError::cancelled_with_details( + CancellationDetails { + reason: "operation cancelled".to_owned(), + }, + )), _ = heartbeat.tick() => ctx.record_heartbeat(()).await?, } } } } +#[workflow] +#[derive(Default)] +struct FailedCancellationDetailsWf; + +#[workflow_methods] +impl FailedCancellationDetailsWf { + #[run] + async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { + let err = ctx + .execute_activity( + CancellationPropagationActivities::wait_for_cancellation, + (), + ActivityOptions::with_start_to_close_timeout(Duration::from_secs(30)) + .heartbeat_timeout(Duration::from_secs(1)) + .cancellation_type(ActivityCancellationType::WaitCancellationCompleted) + .build(), + ) + .await + .expect_err("activity should be cancelled"); + Err(err.into()) + } +} + +#[tokio::test] +async fn workflow_failed_cancellation_propagates_details() { + let wf_name = "workflow_failed_cancellation_propagates_details"; + let mut starter = CoreWfStarter::new(wf_name); + let started = Arc::new(Semaphore::new(0)); + starter + .sdk_config + .register_activities(CancellationPropagationActivities { + started: started.clone(), + }); + starter + .sdk_config + .register_workflow::() + .unwrap(); + let mut worker = starter.worker().await; + let task_queue = starter.get_task_queue().to_owned(); + let wf_handle = worker + .submit_workflow( + FailedCancellationDetailsWf::run, + (), + WorkflowStartOptions::new(task_queue, wf_name).build(), + ) + .await + .unwrap(); + + let canceller = async { + let _started = started.acquire().await.unwrap(); + wf_handle + .cancel(WorkflowCancelOptions::default()) + .await + .unwrap(); + }; + let (_, worker_result) = tokio::join!(canceller, worker.run_until_done()); + worker_result.unwrap(); + + let Err(WorkflowGetResultError::Cancelled { details }) = + wf_handle.get_result(Default::default()).await + else { + panic!("workflow should complete as cancelled"); + }; + assert_eq!( + details.deserialize::().unwrap(), + CancellationDetails { + reason: "operation cancelled".to_owned(), + } + ); +} + #[workflow] struct CancellationPropagationChild { started: Arc, @@ -173,7 +309,7 @@ impl CancellationPropagationChild { async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { ctx.state(|wf| wf.started.add_permits(1)); ctx.cancelled().await; - Err(WorkflowTermination::Cancelled) + Err(WorkflowTermination::cancelled()) } #[signal] @@ -304,9 +440,16 @@ async fn workflow_cancellation_propagates_to_operations() { let (_, worker_result) = tokio::join!(canceller, worker.run_until_done()); worker_result.unwrap(); - assert_matches!( - wf_handle.get_result(Default::default()).await, - Err(WorkflowGetResultError::Cancelled { .. }) + let Err(WorkflowGetResultError::Cancelled { details }) = + wf_handle.get_result(Default::default()).await + else { + panic!("workflow should complete as cancelled"); + }; + assert_eq!( + details.deserialize::().unwrap(), + CancellationDetails { + reason: "operation cancelled".to_owned(), + } ); } @@ -319,10 +462,11 @@ impl WfWithTimer { #[run(name = DEFAULT_WORKFLOW_TYPE)] async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { ctx.timer(Duration::from_millis(500)).await; - Err(WorkflowTermination::Cancelled) + Err(WorkflowTermination::cancelled()) } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn wf_completing_with_cancelled() { let t = canned_histories::timer_wf_cancel_req_cancelled("1"); diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/child_workflows.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/child_workflows.rs index 8a37f5b04..38e770953 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/child_workflows.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/child_workflows.rs @@ -146,7 +146,7 @@ impl AbandonedChildBugReproChild { #[run(name = "child_wf")] async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { ctx.cancelled().await; - Err(WorkflowTermination::Cancelled) + Err(WorkflowTermination::cancelled()) } } @@ -198,7 +198,7 @@ async fn abandoned_child_bug_repro() { #[workflow] struct AbandonedChildResolvesPostCancelParent { - barr: Arc, + ready: Arc, } #[workflow_methods(factory_only)] @@ -217,8 +217,7 @@ impl AbandonedChildResolvesPostCancelParent { ) .await .expect("Child should start OK"); - let barr = ctx.state(|wf| wf.barr.clone()); - barr.wait().await; + ctx.state(|wf| wf.ready.notify_one()); ctx.cancelled().await; started.cancel("Die reason".to_string()); ctx.timer(Duration::from_secs(1)).await; @@ -242,12 +241,12 @@ impl AbandonedChildResolvesPostCancelChild { #[tokio::test] async fn abandoned_child_resolves_post_cancel() { let mut starter = CoreWfStarter::new("child-workflow-resolves-post-cancel"); - let barr = Arc::new(Barrier::new(2)); - let barr_clone = barr.clone(); + let ready = Arc::new(Notify::new()); + let ready_clone = ready.clone(); starter .sdk_config .register_workflow_with_factory(move || AbandonedChildResolvesPostCancelParent { - barr: barr_clone.clone(), + ready: ready_clone.clone(), }) .unwrap(); starter @@ -267,7 +266,7 @@ async fn abandoned_child_resolves_post_cancel() { .unwrap(); let client = starter.get_core_client().await; let canceller = async { - barr.wait().await; + ready.notified().await; handle .cancel(WorkflowCancelOptions::builder().reason("die").build()) .await @@ -277,6 +276,7 @@ async fn abandoned_child_resolves_post_cancel() { worker.run_until_done().await.unwrap(); }; tokio::join!(canceller, runner); + handle.get_result(Default::default()).await.unwrap(); // Verify no WFT failures on the child workflow. A failure here indicates // the child couldn't deserialize its input (e.g., sending a payload when none expected). @@ -284,10 +284,10 @@ async fn abandoned_child_resolves_post_cancel() { client.get_workflow_handle::("abandoned-child-resolve-post-cancel"); let history = child_handle .fetch_history(Default::default()) + .into_events() .await .unwrap(); let wft_failures: Vec<_> = history - .events() .iter() .filter(|e| { matches!( @@ -420,6 +420,7 @@ impl UnusedChildWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest::rstest] #[case::signal_then_result(true)] #[case::signal_and_result_concurrent(false)] @@ -485,6 +486,7 @@ impl ParentCancelsChildWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn cancel_child_workflow() { let t = canned_histories::single_child_workflow_cancelled("child-id-1"); @@ -567,7 +569,7 @@ impl GrandchildCancelled { #[run(name = "grandchild_wf")] async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { ctx.cancelled().await; - Err(WorkflowTermination::Cancelled) + Err(WorkflowTermination::cancelled()) } } @@ -813,6 +815,7 @@ impl PassChildWorkflowSummaryToMetadata { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn pass_child_workflow_summary_to_metadata() { let wf_id = "1"; @@ -938,6 +941,7 @@ impl ParentWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest( mock_cfg, case::success(child_workflow_happy_hist()), @@ -972,6 +976,7 @@ async fn single_child_workflow_until_completion(mut mock_cfg: MockPollCfg) { worker.run().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn single_child_workflow_start_fail() { let child_wf_id = "child-id-1"; @@ -1040,6 +1045,7 @@ impl CancelBeforeSendWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn single_child_workflow_cancel_before_sent() { let mut t = TestHistoryBuilder::default(); @@ -1100,7 +1106,7 @@ async fn cancel_child_before_started_event() { reason: "parent cancelled".to_string(), } .into(), - CancelWorkflowExecution {}.into(), + CancelWorkflowExecution::default().into(), ], )) .await @@ -1110,7 +1116,7 @@ async fn cancel_child_before_started_event() { let act = core.poll_workflow_activation().await.unwrap(); core.complete_workflow_activation(WorkflowActivationCompletion::from_cmd( act.run_id, - CancelWorkflowExecution {}.into(), + CancelWorkflowExecution::default().into(), )) .await .unwrap(); @@ -1149,10 +1155,11 @@ impl CancelChildBeforeStartedCannedWf { }; assert!(cancelled.raw_details().is_none()); assert!(cancelled.cause().is_none()); - Err(WorkflowTermination::Cancelled) + Err(WorkflowTermination::cancelled()) } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn cancel_child_before_started_event_exposes_cancelled_error() { let t = canned_histories::cancel_child_workflow_before_started_event("child-id-1"); @@ -1189,7 +1196,7 @@ impl CancelChildBeforeStartedParent { // Wait for parent cancellation ctx.cancelled().await; started.cancel(); - Err(WorkflowTermination::Cancelled) + Err(WorkflowTermination::cancelled()) } } @@ -1235,9 +1242,12 @@ async fn cancel_child_wf_before_started_event_real_server() { // Verify no unexpected workflow task failures in history. The bug manifests as a WFT failure // with a nondeterminism error. UnhandledCommand failures are acceptable since the server // may reject a cancel command if it races with the child workflow start. - let history = handle.fetch_history(Default::default()).await.unwrap(); + let history = handle + .fetch_history(Default::default()) + .into_events() + .await + .unwrap(); let unexpected_wft_failures: Vec<_> = history - .events() .iter() .filter(|e| { if let Some(history_event::Attributes::WorkflowTaskFailedEventAttributes(attrs)) = @@ -1344,7 +1354,7 @@ impl UnserializableSignalChild { #[run] async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { ctx.cancelled().await; - Err(WorkflowTermination::Cancelled) + Err(WorkflowTermination::cancelled()) } #[signal] @@ -1488,6 +1498,7 @@ impl UnitChildParentWf { /// Parent that starts a typed child returning () and awaits its result. /// With a canned history whose completion is missing a result (simulating a /// non-Rust child workflow that might not have a result payload). +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn child_workflow_unit_result_none_payload() { // single_child_workflow produces a completion with result: None diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/continue_as_new.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/continue_as_new.rs index 9d9fc315b..adc132768 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/continue_as_new.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/continue_as_new.rs @@ -12,13 +12,14 @@ use temporalio_common::{ search_attributes::{SearchAttributeKey, SearchAttributes}, }; use temporalio_macros::{workflow, workflow_methods}; -use temporalio_sdk::{ContinueAsNewOptions, WorkflowContext, WorkflowResult, WorkflowTermination}; +use temporalio_sdk::{ + ContinueAsNewOptions, ContinueAsNewVersioningBehavior, WorkflowContext, WorkflowResult, +}; use temporalio_sdk_core::{ TunerHolder, replay::{DEFAULT_WORKFLOW_TYPE, canned_histories}, test_help::MockPollCfg, }; -use temporalio_workflow::runtime::types::ContinueAsNewRequest; const SA_TXT: SearchAttributeKey = SearchAttributeKey::text(SEARCH_ATTR_TXT); @@ -60,12 +61,61 @@ async fn continue_as_new_happy_path() { worker.run_until_done().await.unwrap(); } +#[workflow] +#[derive(Default)] +struct ContinueAsNewRandomWf; + +#[workflow_methods] +impl ContinueAsNewRandomWf { + #[run] + async fn run( + ctx: &mut WorkflowContext, + previous_value: Option, + ) -> WorkflowResult<(u64, u64)> { + let value = ctx.random_stream("continue-as-new-test").random::(); + if ctx.info().continued_from_run_id().is_none() { + ctx.continue_as_new(Some(value), ContinueAsNewOptions::default())?; + } + Ok(( + previous_value.expect("first run should pass its stream value"), + value, + )) + } +} + +#[tokio::test] +async fn continue_as_new_reseeds_named_random_streams() { + let wf_name = "continue_as_new_reseeds_named_random_streams"; + let mut starter = CoreWfStarter::new(wf_name); + starter + .sdk_config + .register_workflow::() + .unwrap(); + let mut worker = starter.worker().await; + + let task_queue = starter.get_task_queue().to_owned(); + let handle = worker + .submit_workflow( + ContinueAsNewRandomWf::run, + None, + WorkflowStartOptions::new(task_queue, wf_name).build(), + ) + .await + .unwrap(); + worker.run_until_done().await.unwrap(); + let (first_value, continued_value) = handle.get_result(Default::default()).await.unwrap(); + assert_ne!( + first_value, continued_value, + "continue-as-new should independently seed named streams" + ); +} + #[tokio::test] async fn continue_as_new_multiple_concurrent() { let wf_name = "continue_as_new_multiple_concurrent"; let mut starter = CoreWfStarter::new(wf_name); starter.sdk_config.max_cached_workflows = 5_usize; - starter.sdk_config.tuner = Arc::new(TunerHolder::fixed_size(5, 1, 1, 1)); + starter.set_core_tuner(Arc::new(TunerHolder::fixed_size(5, 1, 1, 1))); starter .sdk_config .register_workflow::() @@ -96,14 +146,17 @@ impl WfWithTimer { #[run(name = DEFAULT_WORKFLOW_TYPE)] async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { ctx.timer(Duration::from_millis(500)).await; - Err(WorkflowTermination::continue_as_new(ContinueAsNewRequest { - arguments: vec![[1].into()], - initial_versioning_behavior: ProtoContinueAsNewVersioningBehavior::AutoUpgrade.into(), - ..Default::default() - })) + ctx.continue_as_new( + (), + ContinueAsNewOptions::builder() + .initial_versioning_behavior(ContinueAsNewVersioningBehavior::AutoUpgrade) + .build(), + )?; + Ok(()) } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn wf_completing_with_continue_as_new() { let t = canned_histories::timer_then_continue_as_new("1"); @@ -154,6 +207,7 @@ impl ContinueAsNewSuggestedWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn continue_as_new_suggested_flag_exposed() { let mut t = canned_histories::timer_then_continue_as_new("1"); @@ -195,6 +249,10 @@ impl ClearSearchAttrsOnContinueAsNewWf { } } +#[temporalio_macros::cloud_test_exclusion( + crate::CloudTestExclusionReason::RequiresCloudProvisioning, + "Uses a custom search attribute that isolated Cloud CI does not provision." +)] #[tokio::test] async fn clear_search_attributes_on_continue_as_new() { let wf_name = "clear_search_attrs_on_continue_as_new"; diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/determinism.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/determinism.rs index 760ee7868..7922b4892 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/determinism.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/determinism.rs @@ -152,7 +152,15 @@ struct RandomReplayWf; impl RandomReplayWf { #[run] async fn run(ctx: &mut WorkflowContext) -> WorkflowResult { - Ok(format!("{}:{}", ctx.random::(), ctx.uuid4())) + let orders = ctx.random_stream("example.com/orders"); + let first_order = orders.random::(); + let _ = ctx.random_stream("example.com/telemetry").random::(); + let second_order = ctx.random_stream("example.com/orders").random::(); + Ok(format!( + "{}:{}:{first_order}:{second_order}", + ctx.random::(), + ctx.uuid4() + )) } } @@ -212,6 +220,7 @@ impl TimerWfFailsOnce { /// Verifies that workflow panics (which in this case the Rust SDK turns into workflow activation /// failures) are turned into unspecified WFT failures. +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn test_panic_wf_task_rejected_properly() { let wf_id = "fakeid"; @@ -271,6 +280,7 @@ impl NondeterministicTimerWf { /// Verifies nondeterministic behavior in workflows results in automatic WFT failure with the /// appropriate nondeterminism cause. +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest::rstest] #[case::with_cache(true)] #[case::without_cache(false)] @@ -370,6 +380,7 @@ impl ActivityIdOrTypeChangeWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest::rstest] #[tokio::test] async fn activity_id_or_type_change_is_nondeterministic( @@ -459,6 +470,7 @@ impl ChildWfIdOrTypeChangeWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest::rstest] #[tokio::test] async fn child_wf_id_or_type_change_is_nondeterministic( @@ -617,6 +629,7 @@ impl ReproChannelMissingWf { /// us to want to auto-fail the workflow task while there is also an outstanding eviction, the wf /// would get evicted but then try to send some info down the completion channel afterward, causing /// a panic. +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn repro_channel_missing_because_nondeterminism() { for _ in 1..50 { diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/eager.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/eager.rs index 403c50e1b..2724e2f6b 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/eager.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/eager.rs @@ -1,4 +1,4 @@ -use crate::common::{CoreWfStarter, NAMESPACE, get_integ_connection}; +use crate::common::{CoreWfStarter, get_integ_connection}; use std::time::Duration; use temporalio_client::{Client, NamespacedClient, WorkflowStartOptions, grpc::WorkflowService}; use temporalio_common::protos::temporal::api::{ @@ -58,7 +58,8 @@ async fn eager_wf_start_different_clients() { let mut worker = starter.worker().await; let connection = get_integ_connection(None).await; - let client_opts = temporalio_client::ClientOptions::new(NAMESPACE).build(); + let client_opts = + temporalio_client::ClientOptions::new(starter.get_core_client().await.namespace()).build(); let client = Client::new(connection, client_opts).unwrap(); let task_queue = starter.get_task_queue().to_string(); let res = eager_start( diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/event_groups.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/event_groups.rs new file mode 100644 index 000000000..39ab55ae7 --- /dev/null +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/event_groups.rs @@ -0,0 +1,430 @@ +//! Verify that `EventGroupMarker`s attached to lang-side options propagate all the +//! way down to the server-side `Command`s issued by Core. One mocked test per command +//! kind we currently expose `event_group_markers` on: activity, child workflow, timer, +//! local activity. +//! +//! Plus one end-to-end test against a real server, verifying that the markers also +//! land on the resulting `HistoryEvent` (i.e. the server persists what we send). +//! +//! Event Groups are not implemented in the Rust SDK, so these tests build markers as raw +//! protos and set them through the `#[doc(hidden)]` `event_group_markers` option fields, +//! which exist for that purpose only. + +use std::time::Duration; + +use crate::common::{ + CoreWfStarter, activity_functions::StdActivities, build_fake_sdk_with_options, + mock_sdk_cfg_with_options, +}; +use temporalio_client::{UntypedWorkflow, WorkflowStartOptions}; +use temporalio_common::{ + data_converters::RawValue, + protos::{ + coresdk::AsJsonPayloadExt, + temporal::api::{ + enums::v1::{CommandType, EventType}, + sdk::v1::{ + EventGroupMarker, + event_group_marker::{Label, Variant}, + }, + }, + }, +}; +use temporalio_macros::{workflow, workflow_methods}; +use temporalio_sdk::{ + ActivityOptions, ChildWorkflowOptions, LocalActivityOptions, TimerOptions, WorkflowContext, + WorkflowResult, +}; +use temporalio_sdk_core::{ + replay::{DEFAULT_WORKFLOW_TYPE, canned_histories}, + test_help::MockPollCfg, +}; + +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] +#[tokio::test] +async fn pass_event_group_markers_on_schedule_activity() { + let t = canned_histories::single_activity("1"); + let mut mock_cfg = MockPollCfg::from_hist_builder(t); + let wf_id = mock_cfg.hists[0].wf_id.clone(); + let wf_type = DEFAULT_WORKFLOW_TYPE; + let expected_markers = vec![label_marker("activity-group", "activity-group-label")]; + + let expected_for_assert = expected_markers.clone(); + mock_cfg.completion_asserts_from_expectations(|mut asserts| { + asserts + .then(move |wft| { + assert_eq!(wft.commands.len(), 1); + assert_eq!( + wft.commands[0].command_type(), + CommandType::ScheduleActivityTask + ); + assert_eq!(wft.commands[0].event_group_markers, expected_for_assert); + }) + .then(|wft| { + assert_eq!(wft.commands.len(), 1); + assert_eq!( + wft.commands[0].command_type(), + CommandType::CompleteWorkflowExecution + ); + assert!(wft.commands[0].event_group_markers.is_empty()); + }); + }); + + #[workflow] + struct ActivityWithGroupWorkflow { + event_group_markers: Vec, + } + + #[workflow_methods(factory_only)] + impl ActivityWithGroupWorkflow { + #[run(name = DEFAULT_WORKFLOW_TYPE)] + async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { + let event_group_markers = ctx.state(|wf| wf.event_group_markers.clone()); + ctx.execute_activity( + StdActivities::default, + (), + ActivityOptions::with_start_to_close_timeout(Duration::from_secs(5)) + .event_group_markers(event_group_markers) + .build(), + ) + .await?; + Ok(()) + } + } + + let mut worker = mock_sdk_cfg_with_options( + mock_cfg, + |_| {}, + |options| { + options + .register_workflow_with_factory(move || ActivityWithGroupWorkflow { + event_group_markers: expected_markers.clone(), + }) + .unwrap(); + }, + ); + let task_queue = worker.inner_mut().task_queue().to_owned(); + worker + .submit_wf( + wf_type.to_owned(), + vec![], + WorkflowStartOptions::new(task_queue, wf_id.to_owned()).build(), + ) + .await + .unwrap(); + worker.run_until_done().await.unwrap(); +} + +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] +#[tokio::test] +async fn pass_event_group_markers_on_start_child_workflow() { + let wf_id = "1"; + let wf_type = DEFAULT_WORKFLOW_TYPE; + let t = canned_histories::single_child_workflow(wf_id); + let mut mock_cfg = MockPollCfg::from_hist_builder(t); + let expected_markers = vec![label_marker("child-group", "child-group-label")]; + + let expected_for_assert = expected_markers.clone(); + mock_cfg.completion_asserts_from_expectations(|mut asserts| { + asserts + .then(move |wft| { + assert_eq!(wft.commands.len(), 1); + assert_eq!( + wft.commands[0].command_type(), + CommandType::StartChildWorkflowExecution + ); + assert_eq!(wft.commands[0].event_group_markers, expected_for_assert); + }) + .then(|wft| { + assert_eq!(wft.commands.len(), 1); + assert_eq!( + wft.commands[0].command_type(), + CommandType::CompleteWorkflowExecution + ); + assert!(wft.commands[0].event_group_markers.is_empty()); + }); + }); + + #[workflow] + struct ChildWithGroupWorkflow { + child_wf_id: String, + event_group_markers: Vec, + } + + #[workflow_methods(factory_only)] + impl ChildWithGroupWorkflow { + #[run(name = DEFAULT_WORKFLOW_TYPE)] + async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { + let (child_wf_id, event_group_markers) = + ctx.state(|wf| (wf.child_wf_id.clone(), wf.event_group_markers.clone())); + ctx.start_child_workflow( + UntypedWorkflow::new("child"), + RawValue::new(vec![]), + ChildWorkflowOptions::builder() + .workflow_id(child_wf_id) + .event_group_markers(event_group_markers) + .build(), + ) + .await?; + Ok(()) + } + } + + let child_wf_id = wf_id.to_string(); + let event_group_markers_for_wf = expected_markers.clone(); + let mut worker = mock_sdk_cfg_with_options( + mock_cfg, + |_| {}, + |options| { + options + .register_workflow_with_factory(move || ChildWithGroupWorkflow { + child_wf_id: child_wf_id.clone(), + event_group_markers: event_group_markers_for_wf.clone(), + }) + .unwrap(); + }, + ); + let task_queue = worker.inner_mut().task_queue().to_owned(); + worker + .submit_wf( + wf_type.to_owned(), + vec![], + WorkflowStartOptions::new(task_queue, wf_id.to_owned()).build(), + ) + .await + .unwrap(); + worker.run_until_done().await.unwrap(); +} + +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] +#[tokio::test] +async fn pass_event_group_markers_on_start_timer() { + let t = canned_histories::single_timer("1"); + let mut mock_cfg = MockPollCfg::from_hist_builder(t); + let wf_id = mock_cfg.hists[0].wf_id.clone(); + let wf_type = DEFAULT_WORKFLOW_TYPE; + let expected_markers = vec![label_marker("timer-group", "timer-group-label")]; + + let expected_for_assert = expected_markers.clone(); + mock_cfg.completion_asserts_from_expectations(|mut asserts| { + asserts + .then(move |wft| { + assert_eq!(wft.commands.len(), 1); + assert_eq!(wft.commands[0].command_type(), CommandType::StartTimer); + assert_eq!(wft.commands[0].event_group_markers, expected_for_assert); + }) + .then(|wft| { + assert_eq!(wft.commands.len(), 1); + assert_eq!( + wft.commands[0].command_type(), + CommandType::CompleteWorkflowExecution + ); + assert!(wft.commands[0].event_group_markers.is_empty()); + }); + }); + + #[workflow] + struct TimerWithGroupWorkflow { + event_group_markers: Vec, + } + + #[workflow_methods(factory_only)] + impl TimerWithGroupWorkflow { + #[run(name = DEFAULT_WORKFLOW_TYPE)] + async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { + let event_group_markers = ctx.state(|wf| wf.event_group_markers.clone()); + ctx.timer( + TimerOptions::builder(Duration::from_secs(1)) + .event_group_markers(event_group_markers) + .build(), + ) + .await; + Ok(()) + } + } + + let event_group_markers_for_wf = expected_markers.clone(); + let mut worker = mock_sdk_cfg_with_options( + mock_cfg, + |_| {}, + |options| { + options + .register_workflow_with_factory(move || TimerWithGroupWorkflow { + event_group_markers: event_group_markers_for_wf.clone(), + }) + .unwrap(); + }, + ); + let task_queue = worker.inner_mut().task_queue().to_owned(); + worker + .submit_wf( + wf_type.to_owned(), + vec![], + WorkflowStartOptions::new(task_queue, wf_id.to_owned()).build(), + ) + .await + .unwrap(); + worker.run_until_done().await.unwrap(); +} + +/// Local activities pose some particular challenges: the corresponding `RecordMarker` command +/// only gets created at a later point, after the local activity completes execution. +/// server, so instead of a command of their own they produce a `RecordMarker` command that +/// Core synthesizes when the activity resolves. Markers attached to the `ScheduleLocalActivity` +/// command have to survive that indirection and end up on the marker command. +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] +#[tokio::test] +async fn pass_event_group_markers_on_schedule_local_activity() { + let t = canned_histories::single_local_activity("1"); + let mut mock_cfg = MockPollCfg::from_hist_builder(t); + let expected_markers = vec![label_marker("local-activity-group", "local-activity-label")]; + + let expected_for_assert = expected_markers.clone(); + mock_cfg.completion_asserts_from_expectations(|mut asserts| { + // The activity resolves within the same workflow task that scheduled it, so the marker + // command is flushed together with the workflow completion rather than on its own. + asserts.then(move |wft| { + assert_eq!(wft.commands.len(), 2); + assert_eq!(wft.commands[0].command_type(), CommandType::RecordMarker); + assert_eq!(wft.commands[0].event_group_markers, expected_for_assert); + assert_eq!( + wft.commands[1].command_type(), + CommandType::CompleteWorkflowExecution + ); + assert!(wft.commands[1].event_group_markers.is_empty()); + }); + }); + + #[workflow] + struct LocalActivityWithGroupWorkflow { + event_group_markers: Vec, + } + + #[workflow_methods(factory_only)] + impl LocalActivityWithGroupWorkflow { + #[run(name = DEFAULT_WORKFLOW_TYPE)] + async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { + let event_group_markers = ctx.state(|wf| wf.event_group_markers.clone()); + ctx.execute_local_activity( + StdActivities::default, + (), + LocalActivityOptions::builder() + .event_group_markers(event_group_markers) + .build(), + ) + .await?; + Ok(()) + } + } + + // Unlike the tests above, this one drives a plain SDK worker off the canned history rather + // than submitting a workflow: the local activity must actually run for a marker to be + // recorded, so the worker needs the activity implementation registered too. + let mut worker = build_fake_sdk_with_options(mock_cfg, |options| { + options + .register_workflow_with_factory(move || LocalActivityWithGroupWorkflow { + event_group_markers: expected_markers.clone(), + }) + .unwrap() + .register_activities(StdActivities); + }); + worker.run().await.unwrap(); +} + +// Constants used by the real-server test below; defining them at module scope so the +// workflow body and the assertion can construct the same marker independently. +const PERSIST_TEST_MARKER_ID: &str = "persist-test"; +const PERSIST_TEST_MARKER_LABEL: &str = "persist-test-label"; +const PERSIST_TEST_LA_MARKER_ID: &str = "persist-test-la"; +const PERSIST_TEST_LA_MARKER_LABEL: &str = "persist-test-la-label"; + +#[workflow] +#[derive(Default)] +pub(crate) struct ActivityEventGroupPersistsWf; + +#[workflow_methods] +impl ActivityEventGroupPersistsWf { + #[run(name = "event_group_markers_persist_to_history_events")] + pub(crate) async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { + ctx.execute_activity( + StdActivities::default, + (), + ActivityOptions::with_start_to_close_timeout(Duration::from_secs(5)) + .event_group_markers(vec![label_marker( + PERSIST_TEST_MARKER_ID, + PERSIST_TEST_MARKER_LABEL, + )]) + .build(), + ) + .await?; + ctx.execute_local_activity( + StdActivities::default, + (), + LocalActivityOptions::builder() + .start_to_close_timeout(Duration::from_secs(5)) + .event_group_markers(vec![label_marker( + PERSIST_TEST_LA_MARKER_ID, + PERSIST_TEST_LA_MARKER_LABEL, + )]) + .build(), + ) + .await?; + Ok(()) + } +} + +/// End-to-end: a marker attached to a command must also land on the resulting history event +/// after the server persists it. Covers both an ordinary activity (`ActivityTaskScheduled`) and +/// a local activity, which surfaces as the `MarkerRecorded` event Core writes on resolution. +#[tokio::test] +async fn event_group_markers_persist_to_history_events() { + let wf_name = "event_group_markers_persist_to_history_events"; + let mut starter = CoreWfStarter::new(wf_name); + starter + .sdk_config + .register_activities(StdActivities) + .register_workflow::() + .unwrap(); + let mut worker = starter.worker().await; + + starter.start_with_worker(wf_name, &mut worker).await; + worker.run_until_done().await.unwrap(); + + let history = starter.get_history().await; + let scheduled_events: Vec<_> = history + .events + .iter() + .filter(|e| e.event_type() == EventType::ActivityTaskScheduled) + .collect(); + assert_eq!(scheduled_events.len(), 1); + assert_eq!( + scheduled_events[0].event_group_markers, + vec![label_marker( + PERSIST_TEST_MARKER_ID, + PERSIST_TEST_MARKER_LABEL + )] + ); + + let marker_events: Vec<_> = history + .events + .iter() + .filter(|e| e.event_type() == EventType::MarkerRecorded) + .collect(); + assert_eq!(marker_events.len(), 1); + assert_eq!( + marker_events[0].event_group_markers, + vec![label_marker( + PERSIST_TEST_LA_MARKER_ID, + PERSIST_TEST_LA_MARKER_LABEL + )] + ); +} + +fn label_marker(id: &str, label: &str) -> EventGroupMarker { + EventGroupMarker { + variant: Some(Variant::Label(Label { + id: id.to_string(), + label: Some(label.as_json_payload().unwrap()), + })), + } as EventGroupMarker +} diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/interceptors.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/interceptors.rs index 88dc5d07f..1550300f4 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/interceptors.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/interceptors.rs @@ -1,4 +1,4 @@ -use crate::common::{CoreWfStarter, activity_functions::StdActivities}; +use crate::common::{CoreWfStarter, WorkflowHandleExt, activity_functions::StdActivities}; use std::{ future::Future, pin::Pin, @@ -20,7 +20,8 @@ use temporalio_common::protos::temporal::api::{ use temporalio_macros::{workflow, workflow_methods}; use temporalio_sdk::{ ActivityOptions, ChildWorkflowOptions, LocalActivityOptions, NexusOperationOptions, - SyncWorkflowContext, TimerResult, WorkflowContext, WorkflowContextView, WorkflowResult, + SyncWorkflowContext, TimerResult, WorkflowContext, WorkflowContextKey, WorkflowContextView, + WorkflowResult, workflow_interceptors::{ CancellableWorkflowOutboundFuture, ExecuteWorkflowInput, ExecuteWorkflowResult, HandleQueryInput, HandleQueryResult, HandleSignalInput, HandleSignalResult, @@ -69,9 +70,13 @@ impl InboundInterceptorWorkflow { #[update_validator(set_update)] fn validate_set_update( &self, - _ctx: &WorkflowContextView, + ctx: &WorkflowContextView, input: &str, ) -> Result<(), Box> { + assert_eq!( + ctx.context_value::().as_deref(), + Some(&"validator") + ); assert!(input.ends_with("-validated")); if input.starts_with("reject") { Err("update rejected by validator".into()) @@ -88,7 +93,11 @@ impl InboundInterceptorWorkflow { } #[query] - fn get_status(&self, _ctx: &WorkflowContextView, input: String) -> String { + fn get_status(&self, ctx: &WorkflowContextView, input: String) -> String { + assert_eq!( + ctx.context_value::().as_deref(), + Some(&"query") + ); assert_eq!(input, "query-mutated"); "query-original-output".to_string() } @@ -190,7 +199,10 @@ impl WorkflowInterceptor for MutatingWorkflowInterceptor { *input = "query-mutated".to_string(); } - let result = next.run(input)?; + assert!(ctx.context_value::().is_none()); + let result = + ctx.with_context_value::("query", || next.run(input))?; + assert!(ctx.context_value::().is_none()); assert_eq!( result.downcast_ref::().map(String::as_str), Some("query-original-output") @@ -200,7 +212,7 @@ impl WorkflowInterceptor for MutatingWorkflowInterceptor { fn validate_update( &self, - _ctx: SyncWorkflowInterceptorContext, + ctx: SyncWorkflowInterceptorContext, mut input: ValidateUpdateInput, next: WorkflowNext<'_, ValidateUpdateInput, ValidateUpdateResult>, ) -> ValidateUpdateResult { @@ -213,7 +225,11 @@ impl WorkflowInterceptor for MutatingWorkflowInterceptor { if let Some(input) = input.input_mut::() { input.push_str("-validated"); } - next.run(input) + assert!(ctx.context_value::().is_none()); + let result = + ctx.with_context_value::("validator", || next.run(input)); + assert!(ctx.context_value::().is_none()); + result } } @@ -310,6 +326,465 @@ async fn workflow_interceptors_mutate_inputs_and_replace_outputs() { join!(driver, run); } +#[workflow] +#[derive(Default)] +struct AllHandlersFinishedWorkflow { + handler_started: bool, +} + +#[workflow_methods] +impl AllHandlersFinishedWorkflow { + #[run] + async fn run(ctx: &mut WorkflowContext) -> WorkflowResult { + ctx.wait_condition(|state| state.handler_started).await?; + let handlers_finished = ctx.all_handlers_finished(); + let ctx_clone = ctx.clone(); + ctx.wait_condition(move |_| ctx_clone.all_handlers_finished()) + .await?; + Ok(handlers_finished) + } + + #[signal] + fn sync_signal(&mut self, ctx: &mut SyncWorkflowContext) { + assert!(!ctx.all_handlers_finished()); + self.handler_started = true; + } + + #[signal] + async fn async_signal(ctx: &mut WorkflowContext) { + ctx.state_mut(|state| state.handler_started = true); + } + + #[signal] + fn wake(&mut self, _ctx: &mut SyncWorkflowContext) {} + + #[update_validator(async_update)] + fn validate_async_update( + &self, + _ctx: &WorkflowContextView, + reject: &bool, + ) -> Result<(), Box> { + if *reject { + Err("rejected by validator".into()) + } else { + Ok(()) + } + } + + #[update] + async fn async_update(ctx: &mut WorkflowContext, _reject: bool) { + ctx.state_mut(|state| state.handler_started = true); + } +} + +struct PostHandlerTimerInterceptor; + +impl WorkflowInterceptor for PostHandlerTimerInterceptor { + fn handle_signal<'a>( + &'a self, + ctx: WorkflowInterceptorContext, + input: HandleSignalInput, + next: WorkflowNext< + 'a, + HandleSignalInput, + WorkflowInterceptorFuture<'a, HandleSignalResult>, + >, + ) -> WorkflowInterceptorFuture<'a, HandleSignalResult> { + if input.name() == "wake" { + return next.run(input); + } + WorkflowInterceptorFuture::new(async move { + let result = next.run(input).await; + ctx.timer(Duration::from_millis(1)).await; + result + }) + } + + fn handle_update<'a>( + &'a self, + ctx: WorkflowInterceptorContext, + input: HandleUpdateInput, + next: WorkflowNext< + 'a, + HandleUpdateInput, + WorkflowInterceptorFuture<'a, HandleUpdateResult>, + >, + ) -> WorkflowInterceptorFuture<'a, HandleUpdateResult> { + WorkflowInterceptorFuture::new(async move { + let result = next.run(input).await; + ctx.timer(Duration::from_millis(1)).await; + result + }) + } +} + +#[derive(Clone, Copy)] +enum HandlerKind { + SyncSignal, + AsyncSignal, + Update, +} + +#[rstest::rstest] +#[tokio::test] +async fn all_handlers_finished_waits_for_handler_chain( + #[values( + HandlerKind::SyncSignal, + HandlerKind::AsyncSignal, + HandlerKind::Update + )] + handler_kind: HandlerKind, + #[values(false, true)] with_interceptor: bool, +) { + let mut starter = CoreWfStarter::new("all_handlers_finished_waits_for_handler_chain"); + starter + .sdk_config + .register_workflow::() + .unwrap(); + if with_interceptor { + starter.sdk_config.register_workflow_interceptors(vec![ + WorkflowInterceptorConstructor::new(|_| PostHandlerTimerInterceptor), + ]); + } + let mut worker = starter.worker().await; + + let handle = worker + .submit_workflow( + AllHandlersFinishedWorkflow::run, + (), + WorkflowStartOptions::new( + starter.get_task_queue().to_owned(), + starter.get_wf_id().to_owned(), + ) + .build(), + ) + .await + .unwrap(); + + let driver = async { + match handler_kind { + HandlerKind::SyncSignal => { + handle + .signal( + AllHandlersFinishedWorkflow::sync_signal, + (), + WorkflowSignalOptions::default(), + ) + .await + .unwrap(); + } + HandlerKind::AsyncSignal => { + handle + .signal( + AllHandlersFinishedWorkflow::async_signal, + (), + WorkflowSignalOptions::default(), + ) + .await + .unwrap(); + } + HandlerKind::Update => { + handle + .execute_update( + AllHandlersFinishedWorkflow::async_update, + false, + WorkflowExecuteUpdateOptions::default(), + ) + .await + .unwrap(); + } + } + assert_eq!( + // If interceptor wasn't registered, no timer was scheduled after the handlers so they should + // be finished after the first `wait_condition`. + !with_interceptor, + handle.get_result(Default::default()).await.unwrap() + ); + }; + let (_, worker_result) = join!(driver, worker.run_until_done()); + worker_result.unwrap(); + handle.fetch_history_and_replay(&mut worker).await.unwrap(); +} + +struct CurrentContextLabel; + +impl WorkflowContextKey for CurrentContextLabel { + type Value = &'static str; +} + +#[workflow] +#[derive(Default)] +struct WorkflowContextPropagationWorkflow { + finish: bool, +} + +#[workflow_methods] +impl WorkflowContextPropagationWorkflow { + #[run] + async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { + let run_ctx = ctx.clone(); + ctx.with_context_value::("run", async move { + let left_ctx = run_ctx.clone(); + let left = run_ctx.with_context_value::("left", async move { + left_ctx.timer(Duration::from_millis(1)).await; + assert_eq!( + left_ctx.context_value::().as_deref(), + Some(&"left") + ); + }); + let right_ctx = run_ctx.clone(); + let right = run_ctx.with_context_value::("right", async move { + right_ctx.timer(Duration::from_millis(1)).await; + assert_eq!( + right_ctx.context_value::().as_deref(), + Some(&"right") + ); + }); + temporalio_sdk::workflows::join!(left, right); + + assert_eq!( + run_ctx.context_value::().as_deref(), + Some(&"run") + ); + run_ctx.timer(Duration::from_millis(1)).await; + run_ctx.wait_condition(|state| state.finish).await + }) + .await?; + assert!(ctx.context_value::().is_none()); + Ok(()) + } + + #[signal] + async fn finish(ctx: &mut WorkflowContext) { + let signal_ctx = ctx.clone(); + ctx.with_context_value::("signal", async move { + signal_ctx.timer(Duration::from_millis(1)).await; + assert_eq!( + signal_ctx.context_value::().as_deref(), + Some(&"signal") + ); + signal_ctx.state_mut(|state| state.finish = true); + }) + .await; + assert!(ctx.context_value::().is_none()); + } +} + +struct ObserveWorkflowContextInterceptor { + observed: Arc>>, +} + +impl WorkflowInterceptor for ObserveWorkflowContextInterceptor { + fn start_timer( + &self, + ctx: WorkflowInterceptorContext, + input: StartTimerInput, + next: WorkflowNext< + 'static, + StartTimerInput, + CancellableWorkflowOutboundFuture, + >, + ) -> CancellableWorkflowOutboundFuture { + self.observed.lock().unwrap().push( + *ctx.context_value::() + .expect("timer must have workflow context"), + ); + next.run(input) + } +} + +#[tokio::test] +async fn workflow_context_is_branch_and_handler_local_during_replay() { + let mut starter = + CoreWfStarter::new("workflow_context_is_branch_and_handler_local_during_replay"); + let observed = Arc::new(Mutex::new(Vec::new())); + let interceptor_observed = observed.clone(); + starter + .sdk_config + .register_workflow::() + .unwrap() + .register_workflow_interceptors(vec![WorkflowInterceptorConstructor::new(move |_| { + ObserveWorkflowContextInterceptor { + observed: interceptor_observed.clone(), + } + })]); + let mut worker = starter.worker().await; + + let handle = worker + .submit_workflow( + WorkflowContextPropagationWorkflow::run, + (), + WorkflowStartOptions::new( + starter.get_task_queue().to_owned(), + starter.get_wf_id().to_owned(), + ) + .build(), + ) + .await + .unwrap(); + let driver = async { + handle + .signal( + WorkflowContextPropagationWorkflow::finish, + (), + WorkflowSignalOptions::default(), + ) + .await + .unwrap(); + handle.get_result(Default::default()).await.unwrap(); + }; + let (_, worker_result) = join!(driver, worker.run_until_done()); + worker_result.unwrap(); + let mut live_observed = observed.lock().unwrap().clone(); + live_observed.sort_unstable(); + assert_eq!(live_observed, ["left", "right", "run", "signal"]); + + observed.lock().unwrap().clear(); + handle.fetch_history_and_replay(&mut worker).await.unwrap(); + let mut replay_observed = observed.lock().unwrap().clone(); + replay_observed.sort_unstable(); + assert_eq!(replay_observed, ["left", "right", "run", "signal"]); +} + +#[tokio::test] +async fn rejected_update_does_not_leave_a_handler_in_progress() { + let mut starter = CoreWfStarter::new("rejected_update_does_not_leave_a_handler_in_progress"); + starter + .sdk_config + .register_workflow::() + .unwrap() + .register_workflow_interceptors(vec![WorkflowInterceptorConstructor::new(|_| { + PostHandlerTimerInterceptor + })]); + let mut worker = starter.worker().await; + + let handle = worker + .submit_workflow( + AllHandlersFinishedWorkflow::run, + (), + WorkflowStartOptions::new( + starter.get_task_queue().to_owned(), + starter.get_wf_id().to_owned(), + ) + .build(), + ) + .await + .unwrap(); + + let driver = async { + assert!( + handle + .execute_update( + AllHandlersFinishedWorkflow::async_update, + true, + WorkflowExecuteUpdateOptions::default(), + ) + .await + .is_err() + ); + handle + .signal( + AllHandlersFinishedWorkflow::sync_signal, + (), + WorkflowSignalOptions::default(), + ) + .await + .unwrap(); + assert!(!handle.get_result(Default::default()).await.unwrap()); + }; + let (_, worker_result) = join!(driver, worker.run_until_done()); + worker_result.unwrap(); +} + +struct NonTemporalPostHandlerInterceptor { + waiting: Arc, + release: Arc, +} + +impl WorkflowInterceptor for NonTemporalPostHandlerInterceptor { + fn handle_signal<'a>( + &'a self, + _ctx: WorkflowInterceptorContext, + input: HandleSignalInput, + next: WorkflowNext< + 'a, + HandleSignalInput, + WorkflowInterceptorFuture<'a, HandleSignalResult>, + >, + ) -> WorkflowInterceptorFuture<'a, HandleSignalResult> { + if input.name() != "async_signal" { + return next.run(input); + } + let waiting = self.waiting.clone(); + let release = self.release.clone(); + WorkflowInterceptorFuture::new(async move { + let result = next.run(input).await; + waiting.notify_one(); + release.notified().await; + result + }) + } +} + +#[tokio::test] +async fn all_handlers_finished_tracks_nondeterministic_futures() { + let mut starter = + CoreWfStarter::new("all_handlers_finished_tracks_non_temporal_interceptor_futures"); + starter.sdk_config.detect_nondeterministic_futures = false; + let waiting = Arc::new(Notify::new()); + let release = Arc::new(Notify::new()); + let waiting_ref = waiting.clone(); + let release_ref = release.clone(); + starter + .sdk_config + .register_workflow::() + .unwrap() + .register_workflow_interceptors(vec![WorkflowInterceptorConstructor::new(move |_| { + NonTemporalPostHandlerInterceptor { + waiting: waiting_ref.clone(), + release: release_ref.clone(), + } + })]); + let mut worker = starter.worker().await; + + let handle = worker + .submit_workflow( + AllHandlersFinishedWorkflow::run, + (), + WorkflowStartOptions::new( + starter.get_task_queue().to_owned(), + starter.get_wf_id().to_owned(), + ) + .build(), + ) + .await + .unwrap(); + + let driver = async { + handle + .signal( + AllHandlersFinishedWorkflow::async_signal, + (), + WorkflowSignalOptions::default(), + ) + .await + .unwrap(); + waiting.notified().await; + release.notify_one(); + handle + .signal( + AllHandlersFinishedWorkflow::wake, + (), + WorkflowSignalOptions::default(), + ) + .await + .unwrap(); + assert!(!handle.get_result(Default::default()).await.unwrap()); + }; + let (_, worker_result) = join!(driver, worker.run_until_done()); + worker_result.unwrap(); +} + #[workflow] #[derive(Default)] struct InboundInterceptorOrderWorkflow; diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/local_activities.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/local_activities.rs index 5e4de4c42..737ed1d42 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/local_activities.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/local_activities.rs @@ -1,7 +1,7 @@ use crate::common::{ - ActivationAssertionsInterceptor, CoreWfStarter, WorkflowHandleExt, - activity_functions::StdActivities, history_from_proto_binary, init_core_replay_preloaded, - workflows::LaProblemWorkflow, + ActivationAssertionsInterceptor, CoreWfStarter, FailOnNondeterminismInterceptor, + WorkflowHandleExt, activity_functions::StdActivities, history_from_proto_binary, + init_core_replay_preloaded, workflows::LaProblemWorkflow, }; use anyhow::anyhow; use crossbeam_queue::SegQueue; @@ -40,7 +40,7 @@ use temporalio_common::{ }, temporal::api::{ command::v1::{RecordMarkerCommandAttributes, command}, - common::v1::RetryPolicy, + common::v1::{Payload, RetryPolicy}, enums::v1::{ CommandType, EventType, TimeoutType as ProtoTimeoutType, WorkflowTaskFailedCause, }, @@ -56,7 +56,7 @@ use temporalio_sdk::{ CancellableFuture, LocalActivityOptions, TimeoutType, Worker, WorkflowContext, WorkflowContextView, WorkflowResult, activities::{ActivityContext, ActivityError}, - interceptors::{FailOnNondeterminismInterceptor, WorkerInterceptor}, + interceptors::WorkerInterceptor, }; use temporalio_sdk_core::{ PollError, TunerHolder, prost_dur, @@ -276,7 +276,7 @@ impl LocalActFanoutWf { async fn local_act_fanout() { let wf_name = "local_act_fanout"; let mut starter = CoreWfStarter::new(wf_name); - starter.sdk_config.tuner = Arc::new(TunerHolder::fixed_size(5, 1, 1, 1)); + starter.set_core_tuner(Arc::new(TunerHolder::fixed_size(5, 1, 1, 1))); starter.sdk_config.register_activities(StdActivities); starter .sdk_config @@ -944,16 +944,16 @@ async fn repro_nondeterminism_with_timer_bug() { .unwrap(); worker.run_until_done().await.unwrap(); let client = starter.get_core_client().await; - let handle = WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: wf_name.into(), - run_id: Some(handle.run_id().unwrap().to_string()), - first_execution_run_id: None, - } - .bind_untyped(client.clone()); + let handle = WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(wf_name) + .maybe_run_id(Some(handle.run_id().unwrap().to_string())) + .build() + .bind_untyped(client.clone()); handle.fetch_history_and_replay(&mut worker).await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest::rstest] #[tokio::test] async fn weird_la_nondeterminism_repro(#[values(true, false)] fix_hist: bool) { @@ -980,6 +980,7 @@ async fn weird_la_nondeterminism_repro(#[values(true, false)] fix_hist: bool) { worker.run().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn second_weird_la_nondeterminism_repro() { let mut hist = @@ -1002,6 +1003,7 @@ async fn second_weird_la_nondeterminism_repro() { worker.run().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn third_weird_la_nondeterminism_repro() { let mut hist = @@ -1111,13 +1113,12 @@ async fn la_resolve_same_time_as_other_cancel() { .unwrap(); worker.run_until_done().await.unwrap(); let client = starter.get_core_client().await; - let handle = WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: wf_name.into(), - run_id: Some(handle.run_id().unwrap().to_string()), - first_execution_run_id: None, - } - .bind_untyped(client.clone()); + let handle = WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(wf_name) + .maybe_run_id(Some(handle.run_id().unwrap().to_string())) + .build() + .bind_untyped(client.clone()); handle.fetch_history_and_replay(&mut worker).await.unwrap(); } @@ -1352,6 +1353,7 @@ async fn local_activity_with_summary() { /// This test verifies that when replaying we are able to resolve local activities whose data we /// don't see until after the workflow issues the command +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest::rstest] #[case::replay(true, true)] #[case::not_replay(false, true)] @@ -1409,6 +1411,7 @@ async fn local_act_two_wfts_before_marker(#[case] replay: bool, #[case] cached: worker.run_until_done().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn local_act_many_concurrent() { let mut t = TestHistoryBuilder::default(); @@ -1443,6 +1446,7 @@ async fn local_act_many_concurrent() { /// /// The test with shutdown verifies if we call shutdown while the local activity is running that /// shutdown does not complete until it's finished. +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest::rstest] #[case::with_shutdown(true)] #[case::normal_complete(false)] @@ -1535,6 +1539,7 @@ async fn local_act_heartbeat(#[case] shutdown_middle: bool) { runres.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest::rstest] #[case::retry_then_pass(true)] #[case::retry_until_fail(false)] @@ -1620,6 +1625,7 @@ async fn local_act_fail_and_retry(#[case] eventually_pass: bool) { assert_eq!(expected_attempts, attempts.load(Ordering::Relaxed)); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn local_act_retry_long_backoff_uses_timer() { let mut t = TestHistoryBuilder::default(); @@ -1700,6 +1706,7 @@ async fn local_act_retry_long_backoff_uses_timer() { worker.run_until_done().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn local_act_null_result() { let mut t = TestHistoryBuilder::default(); @@ -1736,6 +1743,7 @@ async fn local_act_null_result() { worker.run_until_done().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn local_act_command_immediately_follows_la_marker() { // This repro only works both when cache is off, and there is at least one heartbeat wft @@ -2035,6 +2043,7 @@ async fn la_resolve_during_legacy_query_does_not_combine(#[case] impossible_quer core.drain_pollers_and_shutdown().await; } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn test_schedule_to_start_timeout() { let mut t = TestHistoryBuilder::default(); @@ -2087,6 +2096,7 @@ async fn test_schedule_to_start_timeout() { worker.run_until_done().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest::rstest] #[case::sched_to_start(true)] #[case::sched_to_close(false)] @@ -2192,6 +2202,7 @@ async fn test_schedule_to_start_timeout_not_based_on_original_time( worker.run_until_done().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest::rstest] #[tokio::test] async fn start_to_close_timeout_allows_retries(#[values(true, false)] la_completes: bool) { @@ -2307,6 +2318,7 @@ async fn start_to_close_timeout_allows_retries(#[values(true, false)] la_complet assert_eq!(cancels.load(Ordering::Acquire), num_cancels); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn wft_failure_cancels_running_las() { let mut t = TestHistoryBuilder::default(); @@ -2371,6 +2383,7 @@ async fn wft_failure_cancels_running_las() { worker.run_until_done().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn resolved_las_not_recorded_if_wft_fails_many_times() { // We shouldn't record any LA results if the workflow activation is repeatedly failing. There @@ -2428,6 +2441,7 @@ async fn resolved_las_not_recorded_if_wft_fails_many_times() { worker.run_until_done().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn local_act_records_nonfirst_attempts_ok() { let mut t = TestHistoryBuilder::default(); @@ -2671,6 +2685,7 @@ async fn queries_can_be_received_while_heartbeating() { core.drain_pollers_and_shutdown().await; } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest] #[case::current_history(false)] #[case::old_heartbeat_history_replay(true)] @@ -2770,6 +2785,7 @@ async fn local_activity_after_wf_complete_is_discarded(#[case] old_heartbeat_his core.drain_pollers_and_shutdown().await; } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn local_act_retry_explicit_delay() { let mut t = TestHistoryBuilder::default(); @@ -2881,6 +2897,7 @@ impl LaWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest] #[case::incremental(false, true)] #[case::replay(true, true)] @@ -2981,6 +2998,77 @@ async fn one_la_success(#[case] replay: bool, #[case] completes_ok: bool) { worker.run().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] +#[rstest] +#[case::excluded(false)] +#[case::included(true)] +#[tokio::test] +async fn local_activity_marker_optionally_includes_arguments(#[case] include_arguments: bool) { + let mut history = TestHistoryBuilder::default(); + history.add_by_type(EventType::WorkflowExecutionStarted); + history.add_workflow_task_scheduled_and_started(); + + let arguments: Vec = vec![b"first".into(), b"second".into()]; + let expected_arguments = arguments.clone(); + let mut mock_cfg = MockPollCfg::from_hist_builder(history); + mock_cfg.make_poll_stream_interminable = true; + mock_cfg.completion_asserts_from_expectations(|mut asserts| { + asserts.then(move |wft| { + assert_eq!(wft.commands.len(), 2); + let marker = assert_matches!( + wft.commands[0].attributes.as_ref(), + Some(command::Attributes::RecordMarkerCommandAttributes(marker)) => marker + ); + let marker_input = marker.details.get("input"); + if include_arguments { + assert_eq!(marker_input.unwrap().payloads, expected_arguments); + } else { + assert!(marker_input.is_none()); + } + assert_eq!( + wft.commands[1].command_type(), + CommandType::CompleteWorkflowExecution + ); + }); + }); + let core = mock_worker(build_mock_pollers(mock_cfg)); + + let activation = core.poll_workflow_activation().await.unwrap(); + core.complete_workflow_activation(WorkflowActivationCompletion::from_cmd( + activation.run_id, + ScheduleLocalActivity { + seq: 1, + activity_id: "1".to_string(), + activity_type: "test_act".to_string(), + arguments, + start_to_close_timeout: Some(prost_dur!(from_secs(30))), + include_arguments_in_marker: include_arguments, + ..Default::default() + } + .into(), + )) + .await + .unwrap(); + + let activity_task = core.poll_activity_task().await.unwrap(); + core.complete_activity_task(ActivityTaskCompletion { + task_token: activity_task.task_token, + result: Some(ActivityExecutionResult::ok(b"result".into())), + }) + .await + .unwrap(); + + let resolution = core.poll_workflow_activation().await.unwrap(); + assert_matches!( + resolution.jobs.as_slice(), + [WorkflowActivationJob { + variant: Some(workflow_activation_job::Variant::ResolveActivity(_)), + }] + ); + core.complete_execution(&resolution.run_id).await; + core.drain_pollers_and_shutdown().await; +} + #[workflow] #[derive(Default)] struct TwoLaWf; @@ -3119,6 +3207,7 @@ async fn local_activity_resolutions_are_delivered_incrementally() { /// been delivered and the task answered. Resolutions queued while an earlier one is outstanding with /// lang may schedule further activities, and the markers for all of them belong on the completion /// that finally answers the task. +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn zero_cache_doesnt_evict_before_wft_is_answered() { let wfid = "fake_wf_id"; @@ -3279,6 +3368,7 @@ impl OldBatchedLocalActivityHistoryWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn old_batched_local_activity_history_replays() { let mut history = TestHistoryBuilder::default(); @@ -3399,6 +3489,7 @@ impl HeartbeatGatedActivity { /// Verifies lookahead preserves completion order when one LA resolves before a heartbeat and the /// other resolves after it. +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest] #[tokio::test] async fn mixed_la_completion_times(#[values(true, false)] replay: bool) { @@ -3485,6 +3576,7 @@ async fn mixed_la_completion_times(#[values(true, false)] replay: bool) { /// Both LAs remain outstanding through a heartbeat, then resolve in the following WFT. Replay must /// deliver their resolutions in marker order, even when that differs from schedule order. +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest] #[tokio::test] async fn two_las_with_heartbeat( @@ -3588,6 +3680,7 @@ impl ResolvedActivity { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest] #[tokio::test] async fn two_sequential_las( @@ -3712,6 +3805,7 @@ impl LaTimerLaWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest] #[case::incremental(false)] #[case::replay(true)] @@ -3788,6 +3882,7 @@ async fn las_separated_by_timer(#[case] replay: bool) { worker.run().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn one_la_heartbeating_wft_failure_still_executes() { let mut t = TestHistoryBuilder::default(); @@ -3821,6 +3916,7 @@ async fn one_la_heartbeating_wft_failure_still_executes() { worker.run().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest] #[tokio::test] async fn immediate_cancel( @@ -3882,6 +3978,7 @@ async fn immediate_cancel( worker.run().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest] #[case::incremental(false)] #[case::replay(true)] @@ -4316,9 +4413,9 @@ async fn replay_out_of_order_local_activity_markers_is_deterministic() { let workflow_id = handle.info().workflow_id.clone(); let events = handle .fetch_history(Default::default()) + .into_events() .await - .unwrap() - .into_events(); + .unwrap(); let marker_order = events .iter() .filter_map(|event| match event.attributes.as_ref() { diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/modify_wf_properties.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/modify_wf_properties.rs index ff7a1db16..d7aa163e4 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/modify_wf_properties.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/modify_wf_properties.rs @@ -64,16 +64,15 @@ async fn sends_modify_wf_props() { worker.run_until_done().await.unwrap(); let client = starter.get_core_client().await; - let description = WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: wf_id.to_string(), - run_id: Some(run_id), - first_execution_run_id: None, - } - .bind_untyped(client.clone()) - .describe(WorkflowDescribeOptions::default()) - .await - .unwrap(); + let description = WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(wf_id.to_string()) + .maybe_run_id(Some(run_id)) + .build() + .bind_untyped(client.clone()) + .describe(WorkflowDescribeOptions::default()) + .await + .unwrap(); assert_eq!( description.memo().get::(FIELD_A).unwrap(), Some("enchi".to_string()) @@ -101,6 +100,7 @@ impl ModifyPropsWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn workflow_modify_props() { let mut t = TestHistoryBuilder::default(); diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/nexus.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/nexus.rs index 34dd25f85..41f7bfc60 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/nexus.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/nexus.rs @@ -17,7 +17,7 @@ use temporalio_client::{ use temporalio_common::{ data_converters::{ GenericPayloadConverter, PayloadConverter, RawValue, SerializationContext, - SerializationContextData, + SerializationContextData, WorkflowSerializationContext, }, protos::{ coresdk::{ @@ -99,6 +99,10 @@ impl NexusBasicWf { } } +#[temporalio_macros::cloud_test_exclusion( + crate::CloudTestExclusionReason::RequiresCloudProvisioning, + "Creates a Nexus endpoint through the Operator Service, unavailable to isolated Cloud namespace credentials." +)] #[rstest::rstest] #[tokio::test] async fn nexus_basic( @@ -301,13 +305,17 @@ impl AsyncCompleter { Outcome::Succeed => Ok("completed async".to_string()), Outcome::Cancel | Outcome::CancelAfterRecordedBeforeStarted => { ctx.cancelled().await; - Err(WorkflowTermination::Cancelled) + Err(WorkflowTermination::cancelled()) } _ => Err(ApplicationFailure::new("broken").into()), } } } +#[temporalio_macros::cloud_test_exclusion( + crate::CloudTestExclusionReason::RequiresCloudProvisioning, + "Creates a Nexus endpoint through the Operator Service, unavailable to isolated Cloud namespace credentials." +)] #[rstest::rstest] #[tokio::test] async fn nexus_async( @@ -340,18 +348,16 @@ async fn nexus_async( let core_worker = starter.get_core_worker().await; let endpoint = mk_nexus_endpoint(&mut starter).await; - let schedule_to_close_timeout = if outcome == Outcome::CancelAfterRecordedBeforeStarted { - None - } else { - Some(Duration::from_secs(5)) + let schedule_to_close_timeout = match outcome { + Outcome::CancelAfterRecordedBeforeStarted => None, + Outcome::Timeout => Some(Duration::from_secs(5)), + _ => Some(Duration::from_secs(60)), }; let submitter = worker.get_submitter_handle(); let converter = PayloadConverter::default(); - let ser_ctx = SerializationContext { - data: &SerializationContextData::Workflow, - converter: &converter, - }; + let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let ser_ctx = SerializationContext::new(&context_data, &converter); let wf_handle = worker .submit_workflow( NexusAsyncWf::run, @@ -569,6 +575,10 @@ impl NexusCancelBeforeStartWf { } } +#[temporalio_macros::cloud_test_exclusion( + crate::CloudTestExclusionReason::RequiresCloudProvisioning, + "Creates a Nexus endpoint through the Operator Service, unavailable to isolated Cloud namespace credentials." +)] #[tokio::test] async fn nexus_cancel_before_start() { let wf_name = "nexus_cancel_before_start"; @@ -629,10 +639,14 @@ impl NexusRootCancellationWf { result.status, Some(nexus_operation_result::Status::Cancelled(_)) ); - Err(WorkflowTermination::Cancelled) + Err(WorkflowTermination::cancelled()) } } +#[temporalio_macros::cloud_test_exclusion( + crate::CloudTestExclusionReason::RequiresCloudProvisioning, + "Creates a Nexus endpoint through the Operator Service, unavailable to isolated Cloud namespace credentials." +)] #[tokio::test] async fn workflow_cancellation_propagates_to_started_nexus_operation() { let wf_name = "workflow_cancellation_propagates_to_started_nexus_operation"; @@ -767,6 +781,10 @@ impl NexusMustCompleteTaskWf { } } +#[temporalio_macros::cloud_test_exclusion( + crate::CloudTestExclusionReason::RequiresCloudProvisioning, + "Creates a Nexus endpoint through the Operator Service, unavailable to isolated Cloud namespace credentials." +)] #[rstest::rstest] #[tokio::test] async fn nexus_must_complete_task_to_shutdown(#[values(true, false)] use_grace_period: bool) { @@ -936,7 +954,7 @@ impl AsyncCompleterWf { } ctx.state(|wf| wf.handler_exited_tx.send(true).unwrap()); - Err(WorkflowTermination::Cancelled) + Err(WorkflowTermination::cancelled()) } #[signal(name = "proceed-to-exit")] @@ -945,6 +963,10 @@ impl AsyncCompleterWf { } } +#[temporalio_macros::cloud_test_exclusion( + crate::CloudTestExclusionReason::RequiresCloudProvisioning, + "Creates a Nexus endpoint through the Operator Service, unavailable to isolated Cloud namespace credentials." +)] #[rstest::rstest] #[tokio::test] async fn nexus_cancellation_types( diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/patches.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/patches.rs index 5b443f3a1..347001a4a 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/patches.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/patches.rs @@ -212,8 +212,12 @@ async fn patch_activation_callback_is_memoized_across_replay() { (false, false) ); assert_eq!(callback_calls.load(Ordering::Relaxed), 1); - let history = handle.fetch_history(Default::default()).await.unwrap(); - assert!(!history.events().iter().any(|event| matches!( + let history = handle + .fetch_history(Default::default()) + .into_events() + .await + .unwrap(); + assert!(!history.iter().any(|event| matches!( &event.attributes, Some(EventAttributes::MarkerRecordedEventAttributes(attrs)) if attrs.marker_name == PATCH_MARKER_NAME @@ -305,8 +309,12 @@ async fn declined_patch_can_roll_out_to_old_worker() { }); run_result.unwrap(); assert_eq!(callback_calls.load(Ordering::Relaxed), 1); - let history = handle.fetch_history(Default::default()).await.unwrap(); - assert!(!history.events().iter().any(|event| matches!( + let history = handle + .fetch_history(Default::default()) + .into_events() + .await + .unwrap(); + assert!(!history.iter().any(|event| matches!( &event.attributes, Some(EventAttributes::MarkerRecordedEventAttributes(attrs)) if attrs.marker_name == PATCH_MARKER_NAME @@ -369,8 +377,12 @@ async fn activated_patch_replays_without_consulting_declining_callback() { }); run_result.unwrap(); assert_eq!(activated_calls.load(Ordering::Relaxed), 1); - let history = handle.fetch_history(Default::default()).await.unwrap(); - assert!(history.events().iter().any(|event| matches!( + let history = handle + .fetch_history(Default::default()) + .into_events() + .await + .unwrap(); + assert!(history.iter().any(|event| matches!( &event.attributes, Some(EventAttributes::MarkerRecordedEventAttributes(attrs)) if attrs.marker_name == PATCH_MARKER_NAME @@ -756,6 +768,7 @@ impl PatchWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest] #[case::v1_breaks_on_normal_marker(false, MarkerType::NotDeprecated, 1)] #[case::v1_accepts_dep_marker(false, MarkerType::Deprecated, 1)] @@ -815,6 +828,7 @@ async fn v1_and_v4_changes( } // Note that the not-replaying and no-marker cases don't make sense and hence are absent +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest] #[case::v2_marker_new_path(false, MarkerType::NotDeprecated, 2)] #[case::v2_dep_marker_new_path(false, MarkerType::Deprecated, 2)] @@ -959,6 +973,7 @@ impl SameChangeMultipleSpotsWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest] #[case::has_change_replay(true, true)] #[case::no_change_replay(false, true)] @@ -1072,6 +1087,7 @@ impl ManyPatchesWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest] #[case::happy_path(50)] // We start exceeding the 2k size limit at 180 patches with this format @@ -1179,9 +1195,12 @@ async fn patch_marker_size_overflow_replay_is_deterministic() { // Confirm that the original execution did in fact hit the size limit: the last upsert SA // event in history must contain fewer than the total number of patches issued by the workflow. - let history = handle.fetch_history(Default::default()).await.unwrap(); + let history = handle + .fetch_history(Default::default()) + .into_events() + .await + .unwrap(); let last_upsert_patches = history - .events() .iter() .rev() .find_map(|e| match &e.attributes { diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/priority.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/priority.rs index 92629ae4b..ca0654251 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/priority.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/priority.rs @@ -1,5 +1,9 @@ use crate::shared_tests; +#[temporalio_macros::cloud_test_exclusion( + crate::CloudTestExclusionReason::NeedsCloudAdaptation, + "Uses new_cloud_or_local, which treats envconfig as local and calls a cluster-info RPC unavailable to Cloud namespace credentials." +)] #[tokio::test] async fn priority_values_sent_to_server() { shared_tests::priority::priority_values_sent_to_server().await diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/queries.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/queries.rs index 52921e57a..b42149f92 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/queries.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/queries.rs @@ -59,6 +59,7 @@ impl CompleteOnSecondPollWf { /// The error message from core when this happens is: /// "Workflow completion had a legacy query response along with other commands. /// This is not allowed and constitutes an error in the lang SDK." +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn query_only_activation_should_not_advance_workflow() { let mut t = TestHistoryBuilder::default(); @@ -119,6 +120,7 @@ async fn query_only_activation_should_not_advance_workflow() { } /// Test that a query for a non-existent handler doesn't advance the workflow either. +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn nonexistent_query_should_not_advance_workflow() { let mut t = TestHistoryBuilder::default(); @@ -211,6 +213,7 @@ impl CounterWf { /// Non-legacy queries (in the `queries` field) come bundled with new history. /// Core sends these queries in their own activation after the workflow has processed /// the history, so queries should observe state AFTER the workflow has advanced. +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn non_legacy_query_should_see_state_after_workflow_advances() { let wfid = "non_legacy_query_state_test"; @@ -352,6 +355,7 @@ impl ContextViewWf { } /// Test that WorkflowContextView contains the correct workflow information. +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn query_returns_workflow_context_view_info() { const WFID: &str = "context_view_test_wf"; @@ -440,6 +444,7 @@ impl CurrentDetailsWf { /// Verify that the query returns a proto-JSON-encoded `WorkflowMetadata` /// whose `current_details` field reflects the value set by `set_current_details`. +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn workflow_metadata_query_returns_current_details() { let wfid = "workflow_metadata_query_test"; @@ -533,6 +538,7 @@ impl NoCurrentDetailsWf { /// Verify that the query returns `{}` when `set_current_details` was never /// called, matching proto3 JSON behavior where default (empty) fields are omitted. +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn workflow_metadata_query_empty_details() { let wfid = "workflow_metadata_query_empty_test"; diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/replay.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/replay.rs index bbf9b7420..4fc04d2dd 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/replay.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/replay.rs @@ -74,6 +74,7 @@ fn fire_happy_hist(num_timers: u32) -> Worker { ) } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest] #[case::one_timer(fire_happy_hist(1), 1)] #[case::five_timers(fire_happy_hist(5), 5)] @@ -85,6 +86,7 @@ async fn replay_flag_is_correct(#[case] mut worker: Worker, #[case] _num_timers: worker.run().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test(flavor = "multi_thread")] async fn replay_flag_is_correct_partial_history() { let mut t = canned_histories::long_sequential_timers(2); @@ -101,6 +103,7 @@ async fn replay_flag_is_correct_partial_history() { worker.run().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn timer_workflow_replay() { let core = init_core_replay_preloaded( @@ -158,6 +161,7 @@ async fn timer_workflow_replay() { core.shutdown().await; } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn workflow_nondeterministic_replay() { let core = init_core_replay_preloaded( @@ -198,6 +202,7 @@ async fn workflow_nondeterministic_replay() { ); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn replay_using_wf_function() { let num_timers = 10u32; @@ -210,6 +215,7 @@ async fn replay_using_wf_function() { worker.run().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn replay_ending_wft_complete_with_commands_but_no_scheduled_started() { let mut t = TestHistoryBuilder::default(); @@ -237,12 +243,14 @@ async fn replay_abrupt_ending(mut t: TestHistoryBuilder) { }); worker.run().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn replay_ok_ending_with_terminated() { let mut t1 = canned_histories::single_timer("1"); t1.add_workflow_execution_terminated(); replay_abrupt_ending(t1).await; } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn replay_ok_ending_with_timed_out() { let mut t2 = canned_histories::single_timer("1"); @@ -250,6 +258,7 @@ async fn replay_ok_ending_with_timed_out() { replay_abrupt_ending(t2).await; } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn replay_shutdown_worker() { let mut t = canned_histories::single_timer("1"); @@ -311,6 +320,7 @@ impl SeqTimerWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest::rstest] #[tokio::test] async fn multiple_histories_replay(#[values(false, true)] use_feeder: bool) { @@ -366,6 +376,7 @@ async fn multiple_histories_replay(#[values(false, true)] use_feeder: bool) { assert_eq!(runs_ctr.lock().len(), 2); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn multiple_histories_can_handle_dupe_run_ids() { let mut hist1 = canned_histories::single_timer("1"); @@ -385,6 +396,7 @@ async fn multiple_histories_can_handle_dupe_run_ids() { } // Verifies SDK can decode patch markers before changing them to use json encoding. +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn replay_old_patch_format() { let mut worker = crate::common::replay_sdk_worker_with_options( @@ -401,6 +413,7 @@ async fn replay_old_patch_format() { worker.run().await.unwrap(); } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn replay_ends_with_empty_wft() { let core = init_core_replay_preloaded( diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/resets.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/resets.rs index d568ae84a..e66554427 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/resets.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/resets.rs @@ -1,4 +1,4 @@ -use crate::common::{CoreWfStarter, NAMESPACE, activity_functions::StdActivities}; +use crate::common::{CoreWfStarter, activity_functions::StdActivities}; use std::{ sync::{ Arc, OnceLock, @@ -7,7 +7,7 @@ use std::{ time::Duration, }; use temporalio_client::{ - WorkflowSignalOptions, WorkflowStartOptions, errors::WorkflowGetResultError, + NamespacedClient, WorkflowSignalOptions, WorkflowStartOptions, errors::WorkflowGetResultError, grpc::WorkflowService, }; use temporalio_common::protos::temporal::api::{ @@ -80,7 +80,7 @@ async fn reset_workflow() { client .reset_workflow_execution( ResetWorkflowExecutionRequest { - namespace: NAMESPACE.to_owned(), + namespace: client.namespace(), workflow_execution: Some(WorkflowExecution { workflow_id: wf_name.to_owned(), run_id, @@ -116,6 +116,7 @@ async fn reset_workflow() { struct ResetRandomseedWf { did_fail: Arc, initial_random: Arc>, + initial_named_random: Arc>, reset_started: Arc, saw_updated_random: Arc, notify: Arc, @@ -127,10 +128,13 @@ struct ResetRandomseedWf { impl ResetRandomseedWf { #[run(name = "reset_randomseed")] async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { + let named_random = ctx.random_stream("reset-test"); if ctx.state(|wf| !wf.reset_started.load(Ordering::Relaxed)) { let initial_random = ctx.random::(); + let initial_named_random = named_random.random::(); ctx.state(|wf| { let _ = wf.initial_random.set(initial_random); + let _ = wf.initial_named_random.set(initial_named_random); }); } ctx.timer(Duration::from_millis(100)).await; @@ -157,6 +161,16 @@ impl ResetRandomseedWf { initial_random, "random stream should be reseeded after reset" ); + let initial_named_random = ctx.state(|wf| { + *wf.initial_named_random + .get() + .expect("initial named random value should be recorded") + }); + assert_ne!( + named_random.random::(), + initial_named_random, + "named random stream should be reseeded after reset" + ); ctx.state(|wf| { wf.saw_updated_random.store(true, Ordering::Relaxed); }); @@ -194,10 +208,12 @@ async fn reset_randomseed() { let did_fail = Arc::new(AtomicBool::new(false)); let initial_random = Arc::new(OnceLock::new()); + let initial_named_random = Arc::new(OnceLock::new()); let reset_started = Arc::new(AtomicBool::new(false)); let saw_updated_random = Arc::new(AtomicBool::new(false)); let notify = Arc::new(Notify::new()); let notify_clone = notify.clone(); + let initial_named_random_for_wf = initial_named_random.clone(); let reset_started_for_wf = reset_started.clone(); let saw_updated_random_for_wf = saw_updated_random.clone(); starter @@ -205,6 +221,7 @@ async fn reset_randomseed() { .register_workflow_with_factory(move || ResetRandomseedWf { did_fail: did_fail.clone(), initial_random: initial_random.clone(), + initial_named_random: initial_named_random_for_wf.clone(), reset_started: reset_started_for_wf.clone(), saw_updated_random: saw_updated_random_for_wf.clone(), notify: notify_clone.clone(), @@ -242,7 +259,7 @@ async fn reset_randomseed() { client .reset_workflow_execution( ResetWorkflowExecutionRequest { - namespace: NAMESPACE.to_owned(), + namespace: client.namespace(), workflow_execution: Some(WorkflowExecution { workflow_id: wf_name.to_owned(), run_id: run_id.clone(), diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/signals.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/signals.rs index ea9d8792a..8bb1fd9c2 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/signals.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/signals.rs @@ -1,9 +1,13 @@ use crate::common::{ActivationAssertionsInterceptor, CoreWfStarter}; -use std::collections::HashMap; -use temporalio_client::{WorkflowStartOptions, WorkflowStartSignal}; +use futures::future::BoxFuture; +use std::{collections::HashMap, sync::Arc}; +use temporalio_client::{ + ClientInterceptor, Next, SignalWithStartWorkflowInput, StartWorkflowOutput, + WorkflowStartOptions, errors::WorkflowStartError, +}; use temporalio_common::protos::{ coresdk::{ - AsJsonPayloadExt, IntoPayloadsExt, + AsJsonPayloadExt, workflow_activation::{ ResolveSignalExternalWorkflow, WorkflowActivationJob, workflow_activation_job, }, @@ -20,7 +24,11 @@ use temporalio_sdk_core::replay::{DEFAULT_WORKFLOW_TYPE, TestHistoryBuilder}; use temporalio_macros::{workflow, workflow_methods}; use temporalio_sdk::{ ApplicationFailure, CancellableFuture, ChildWorkflowOptions, SignalWorkflowOptions, - SyncWorkflowContext, WorkflowContext, WorkflowResult, + SyncWorkflowContext, WorkflowContext, WorkflowResult, WorkflowSignalError, + workflow_interceptors::{ + HandleSignalInput, HandleSignalResult, WorkflowInterceptor, WorkflowInterceptorConstructor, + WorkflowInterceptorContext, WorkflowInterceptorFuture, WorkflowNext, + }, }; use temporalio_sdk_core::test_help::MockPollCfg; use uuid::Uuid; @@ -48,7 +56,7 @@ impl SignalSender { ) .await; if expect_failure { - assert!(sigres.is_err()); + assert_matches!(sigres, Err(WorkflowSignalError::NotFound(_))); } else { sigres.unwrap(); } @@ -114,14 +122,49 @@ impl SignalWithCreateWfReceiver { } #[signal(name = "signame")] - fn handle_signal(&mut self, ctx: &mut SyncWorkflowContext, input: String) { + fn handle_signal(&mut self, _ctx: &mut SyncWorkflowContext, input: String) { assert_eq!(input, "tada"); - let headers = ctx.headers(); + self.received = true; + } +} + +struct SignalWithStartHeaderClientInterceptor; + +impl ClientInterceptor for SignalWithStartHeaderClientInterceptor { + fn signal_with_start_workflow<'a>( + &'a self, + mut input: SignalWithStartWorkflowInput, + next: Next< + 'a, + SignalWithStartWorkflowInput, + BoxFuture<'a, Result>, + >, + ) -> BoxFuture<'a, Result> { + input.options.header = + Some(HashMap::from([("tupac".to_string(), Payload::from("shakur"))]).into()); + next.run(input) + } +} + +struct SignalHeaderWorkflowInterceptor; + +impl WorkflowInterceptor for SignalHeaderWorkflowInterceptor { + fn handle_signal<'a>( + &'a self, + _ctx: WorkflowInterceptorContext, + input: HandleSignalInput, + next: WorkflowNext< + 'a, + HandleSignalInput, + WorkflowInterceptorFuture<'a, HandleSignalResult>, + >, + ) -> WorkflowInterceptorFuture<'a, HandleSignalResult> { + assert_eq!(input.name(), SIGNAME); assert_eq!( - *headers.get("tupac").expect("tupac header exists"), + *input.headers().get("tupac").expect("tupac header exists"), b"shakur".into() ); - self.received = true; + next.run(input) } } @@ -164,22 +207,27 @@ async fn sends_signal_with_create_wf() { starter .sdk_config .register_workflow::() - .unwrap(); + .unwrap() + .register_workflow_interceptors(vec![WorkflowInterceptorConstructor::new(|_| { + SignalHeaderWorkflowInterceptor + })]); let mut worker = starter.worker().await; - let client = starter.get_core_client().await; - let mut header: HashMap = HashMap::new(); - header.insert("tupac".into(), "shakur".into()); + let mut client = starter.get_core_client().await; + client + .options_mut() + .client_interceptors + .push(Arc::new(SignalWithStartHeaderClientInterceptor)); let task_queue = worker.inner_mut().task_queue().to_string(); - let start_signal = WorkflowStartSignal::new(SIGNAME) - .maybe_input(vec!["tada".to_string().as_json_payload().unwrap()].into_payloads()) - .maybe_header(Some(header.into())) - .build(); - let options = WorkflowStartOptions::new(task_queue, "sends_signal_with_create_wf") - .start_signal(start_signal) - .build(); + let options = WorkflowStartOptions::new(task_queue, "sends_signal_with_create_wf").build(); let handle = client - .start_workflow(SignalWithCreateWfReceiver::run, (), options) + .signal_with_start_workflow( + SignalWithCreateWfReceiver::run, + (), + SignalWithCreateWfReceiver::handle_signal, + "tada".to_string(), + options, + ) .await .expect("request succeeds.qed"); @@ -289,6 +337,7 @@ impl SignalSenderCanned { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[rstest::rstest] #[case::succeeds(false)] #[case::fails(true)] @@ -360,6 +409,7 @@ impl CancelsBeforeSending { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn cancels_before_sending() { let mut t = TestHistoryBuilder::default(); diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/stickyness.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/stickyness.rs index 7e70b7941..6cb3b5a66 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/stickyness.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/stickyness.rs @@ -9,8 +9,8 @@ use std::{ use temporalio_client::WorkflowStartOptions; use temporalio_common::protos::temporal::api::enums::v1::EventType; use temporalio_macros::{workflow, workflow_methods}; -use temporalio_sdk::{WorkflowContext, WorkflowResult}; -use temporalio_sdk_core::{PollerBehavior, TunerHolder}; +use temporalio_sdk::{WorkflowContext, WorkflowResult, runtime::PollerBehavior}; +use temporalio_sdk_core::TunerHolder; use tokio::sync::Barrier; #[tokio::test] @@ -121,7 +121,7 @@ impl CacheMissWf { async fn cache_miss_ok() { let wf_name = "cache_miss_ok"; let mut starter = CoreWfStarter::new(wf_name); - starter.sdk_config.tuner = Arc::new(TunerHolder::fixed_size(2, 1, 1, 1)); + starter.set_core_tuner(Arc::new(TunerHolder::fixed_size(2, 1, 1, 1))); starter.sdk_config.max_cached_workflows = 0_usize; starter.sdk_config.workflow_task_poller_behavior = Some(PollerBehavior::SimpleMaximum(1_usize)); diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/timers.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/timers.rs index 5ccdb1e39..91a66392f 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/timers.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/timers.rs @@ -175,6 +175,7 @@ impl CancelAlreadyFiredTimerWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn cancel_unpolled_timer_after_both_timers_fire_same_activation() { let mut t = canned_histories::parallel_timer("1", "2"); @@ -203,6 +204,7 @@ impl HappyTimerWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn test_fire_happy_path_inc() { let t = canned_histories::single_timer("1"); @@ -241,6 +243,7 @@ impl MismatchedTimerWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn mismatched_timer_ids_errors() { let t = canned_histories::single_timer("badid"); @@ -273,6 +276,7 @@ impl CancelTimerWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn incremental_cancellation() { let t = canned_histories::cancel_timer("2", "1"); @@ -315,6 +319,7 @@ impl CancelBeforeSentWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn cancel_before_sent_to_server() { let mut t = TestHistoryBuilder::default(); @@ -370,15 +375,16 @@ impl WaitConditionWakerWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn wait_condition_waker_in_futures_unordered() { let t = canned_histories::single_timer_wf_completes("1"); let mock_cfg = MockPollCfg::from_hist_builder(t); let mut worker = crate::common::build_fake_sdk_with_options(mock_cfg, |options| { + // FuturesUnordered uses internal wakers that forward wake calls outside the + // SdkWakeGuard scope. + options.detect_nondeterministic_futures = false; options.register_workflow::().unwrap(); }); - // FuturesUnordered uses internal wakers that forward wake calls outside the - // SdkWakeGuard scope. - worker.set_detect_nondeterministic_futures(false); worker.run().await.unwrap(); } diff --git a/crates/sdk-core/tests/integ_tests/workflow_tests/upsert_search_attrs.rs b/crates/sdk-core/tests/integ_tests/workflow_tests/upsert_search_attrs.rs index 5af2768ff..2076ad7d7 100644 --- a/crates/sdk-core/tests/integ_tests/workflow_tests/upsert_search_attrs.rs +++ b/crates/sdk-core/tests/integ_tests/workflow_tests/upsert_search_attrs.rs @@ -48,6 +48,10 @@ impl SearchAttrUpdater { } } +#[temporalio_macros::cloud_test_exclusion( + crate::CloudTestExclusionReason::RequiresCloudProvisioning, + "Uses custom search attributes that isolated Cloud CI does not provision." +)] #[tokio::test] async fn sends_upsert() { let wf_name = "sends_upsert_search_attrs"; @@ -122,6 +126,7 @@ impl UpsertTestWf { } } +#[temporalio_macros::cloud_test_exclusion(crate::CloudTestExclusionReason::DoesNotUseServer)] #[tokio::test] async fn upsert_search_attrs_from_workflow() { let mut t = TestHistoryBuilder::default(); diff --git a/crates/sdk-core/tests/main.rs b/crates/sdk-core/tests/main.rs index f88abc713..007ed83d2 100644 --- a/crates/sdk-core/tests/main.rs +++ b/crates/sdk-core/tests/main.rs @@ -7,6 +7,15 @@ extern crate assert_matches; mod common; +#[cfg(test)] +pub(crate) enum CloudTestExclusionReason { + DoesNotUseServer, + RequiresLocalServer, + RequiresOssOnlyApis, + RequiresCloudProvisioning, + NeedsCloudAdaptation, +} + #[cfg(test)] mod shared_tests; @@ -24,6 +33,7 @@ mod integ_tests { mod polling_tests; mod queries_tests; mod schedule_tests; + mod standalone_activity_tests; mod update_tests; mod visibility_tests; mod worker_heartbeat_tests; diff --git a/crates/sdk-core/tests/manual_tests.rs b/crates/sdk-core/tests/manual_tests.rs index 02c5b45ab..a60845ebe 100644 --- a/crates/sdk-core/tests/manual_tests.rs +++ b/crates/sdk-core/tests/manual_tests.rs @@ -29,8 +29,9 @@ use temporalio_macros::{activities, workflow, workflow_methods}; use temporalio_sdk::{ ActivityOptions, SyncWorkflowContext, WorkflowContext, WorkflowResult, activities::{ActivityContext, ActivityError}, + runtime::{AutoscalingOptions, PollerBehavior}, }; -use temporalio_sdk_core::{CoreRuntime, PollerBehavior, TunerHolder}; +use temporalio_sdk_core::{CoreRuntime, TunerHolder}; use tracing::info; struct JitteryEchoActivities; @@ -133,17 +134,21 @@ async fn poller_load_spiky() { let rt = CoreRuntime::new_assume_tokio(get_integ_runtime_options(telemopts)).unwrap(); let mut starter = CoreWfStarter::new_with_runtime("poller_load", rt); starter.sdk_config.max_cached_workflows = 5000; - starter.sdk_config.tuner = Arc::new(TunerHolder::fixed_size(1000, 1000, 100, 100)); - starter.sdk_config.workflow_task_poller_behavior = Some(PollerBehavior::Autoscaling { - minimum: 1, - maximum: 200, - initial: 5, - }); - starter.sdk_config.activity_task_poller_behavior = Some(PollerBehavior::Autoscaling { - minimum: 1, - maximum: 200, - initial: 5, - }); + starter.set_core_tuner(Arc::new(TunerHolder::fixed_size(1000, 1000, 100, 100))); + starter.sdk_config.workflow_task_poller_behavior = Some(PollerBehavior::Autoscaling( + AutoscalingOptions::builder() + .minimum(1) + .maximum(200) + .initial(5) + .build(), + )); + starter.sdk_config.activity_task_poller_behavior = Some(PollerBehavior::Autoscaling( + AutoscalingOptions::builder() + .minimum(1) + .maximum(200) + .initial(5) + .build(), + )); starter .sdk_config .register_activities(JitteryEchoActivities) @@ -171,13 +176,12 @@ async fn poller_load_spiky() { .await .unwrap(); workflow_handles.push( - WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: wfid, - run_id: Some(rid), - first_execution_run_id: None, - } - .bind_untyped(client.clone()), + WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(wfid) + .maybe_run_id(Some(rid)) + .build() + .bind_untyped(client.clone()), ); } info!("Done starting workflows"); @@ -209,13 +213,12 @@ async fn poller_load_spiky() { .await .unwrap(); workflow_handles.push( - WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: wfid, - run_id: Some(rid), - first_execution_run_id: None, - } - .bind_untyped(client.clone()), + WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(wfid) + .maybe_run_id(Some(rid)) + .build() + .bind_untyped(client.clone()), ); } stream::iter(workflow_handles) @@ -277,12 +280,14 @@ async fn poller_load_sustained() { let rt = CoreRuntime::new_assume_tokio(get_integ_runtime_options(telemopts)).unwrap(); let mut starter = CoreWfStarter::new_with_runtime("poller_load", rt); starter.sdk_config.max_cached_workflows = 5000; - starter.sdk_config.tuner = Arc::new(TunerHolder::fixed_size(1000, 100, 100, 100)); - starter.sdk_config.workflow_task_poller_behavior = Some(PollerBehavior::Autoscaling { - minimum: 1, - maximum: 200, - initial: 5, - }); + starter.set_core_tuner(Arc::new(TunerHolder::fixed_size(1000, 100, 100, 100))); + starter.sdk_config.workflow_task_poller_behavior = Some(PollerBehavior::Autoscaling( + AutoscalingOptions::builder() + .minimum(1) + .maximum(200) + .initial(5) + .build(), + )); starter .sdk_config .register_workflow::() @@ -308,13 +313,12 @@ async fn poller_load_sustained() { .await .unwrap(); workflow_handles.push( - WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: wfid, - run_id: Some(rid), - first_execution_run_id: None, - } - .bind_untyped(client.clone()), + WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(wfid) + .maybe_run_id(Some(rid)) + .build() + .bind_untyped(client.clone()), ); } info!("Done starting workflows"); @@ -354,17 +358,21 @@ async fn poller_load_spike_then_sustained() { let rt = CoreRuntime::new_assume_tokio(get_integ_runtime_options(telemopts)).unwrap(); let mut starter = CoreWfStarter::new_with_runtime("poller_load", rt); starter.sdk_config.max_cached_workflows = 5000; - starter.sdk_config.tuner = Arc::new(TunerHolder::fixed_size(1000, 100, 100, 100)); - starter.sdk_config.workflow_task_poller_behavior = Some(PollerBehavior::Autoscaling { - minimum: 1, - maximum: 200, - initial: 5, - }); - starter.sdk_config.activity_task_poller_behavior = Some(PollerBehavior::Autoscaling { - minimum: 1, - maximum: 200, - initial: 5, - }); + starter.set_core_tuner(Arc::new(TunerHolder::fixed_size(1000, 100, 100, 100))); + starter.sdk_config.workflow_task_poller_behavior = Some(PollerBehavior::Autoscaling( + AutoscalingOptions::builder() + .minimum(1) + .maximum(200) + .initial(5) + .build(), + )); + starter.sdk_config.activity_task_poller_behavior = Some(PollerBehavior::Autoscaling( + AutoscalingOptions::builder() + .minimum(1) + .maximum(200) + .initial(5) + .build(), + )); starter .sdk_config .register_activities(JitteryEchoActivities) @@ -392,13 +400,12 @@ async fn poller_load_spike_then_sustained() { .await .unwrap(); workflow_handles.push( - WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: wfid, - run_id: Some(rid), - first_execution_run_id: None, - } - .bind_untyped(client.clone()), + WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(wfid) + .maybe_run_id(Some(rid)) + .build() + .bind_untyped(client.clone()), ); } info!("Done starting workflows"); @@ -429,13 +436,12 @@ async fn poller_load_spike_then_sustained() { .await .unwrap(); workflow_handles.push( - WorkflowExecutionInfo { - namespace: client.namespace(), - workflow_id: wfid, - run_id: Some(rid), - first_execution_run_id: None, - } - .bind_untyped(client.clone()), + WorkflowExecutionInfo::builder() + .namespace(client.namespace()) + .workflow_id(wfid) + .maybe_run_id(Some(rid)) + .build() + .bind_untyped(client.clone()), ); tokio::time::sleep(Duration::from_secs(1)).await; } diff --git a/crates/sdk-core/tests/runner.rs b/crates/sdk-core/tests/runner.rs index 427a3d90f..7b1aa6da4 100644 --- a/crates/sdk-core/tests/runner.rs +++ b/crates/sdk-core/tests/runner.rs @@ -1,11 +1,13 @@ +mod cloud_namespace; + // All non-main.rs tests ignore dead common code so that the linter doesn't complain about about it. #[allow(dead_code)] mod common; use crate::common::integ_dev_server_config; use anyhow::{anyhow, bail}; -use clap::Parser; -use common::INTEG_SERVER_TARGET_ENV_VAR; +use clap::{Parser, Subcommand}; +use common::{INTEG_SERVER_TARGET_ENV_VAR, TEST_ENV_CONFIG_SERVER_ENV_VAR}; use std::{ env, path::{Path, PathBuf}, @@ -22,6 +24,9 @@ const INTEG_TEST_SERVER_USED_ENV_VAR: &str = "INTEG_TEST_SERVER_ON"; #[derive(clap::Parser)] #[command(author, version, about, long_about = None)] struct Cli { + #[command(subcommand)] + command: Option, + /// Test harness to run. Anything defined as a `[[test]]` in core's `Cargo.toml` is valid. #[arg(short, long, default_value = "integ_tests")] test_name: String, @@ -42,10 +47,31 @@ struct Cli { /// If set, only run the build, not any tests just_build: bool, + /// Run only tests that are eligible for Temporal Cloud + #[arg(long)] + cloud: bool, + /// The rest of the arguments will be passed through to the test harness harness_args: Vec, } +#[derive(Subcommand)] +enum RunnerCommand { + /// Manage an isolated Temporal Cloud namespace for integration tests + CloudNamespace { + #[command(subcommand)] + command: CloudNamespaceCommand, + }, +} + +#[derive(Subcommand)] +enum CloudNamespaceCommand { + /// Create a namespace and write its full name to GITHUB_OUTPUT + Create, + /// Delete a namespace and wait for deletion to finish + Delete { namespace: String }, +} + #[derive(Copy, Clone, PartialEq, Eq, clap::ValueEnum)] enum ServerKind { /// Use Temporal-cli @@ -54,22 +80,41 @@ enum ServerKind { TestServer, /// Do not automatically start any server External, + /// Load the server connection configuration from envconfig without starting a server + #[value(name = "envconfig")] + EnvConfig, } #[tokio::main] async fn main() -> Result<(), anyhow::Error> { let Cli { + command, test_name, server_kind, cargo_test_args, test_executable, just_build, + cloud, harness_args, } = Cli::parse(); + if let Some(RunnerCommand::CloudNamespace { command }) = command { + return match command { + CloudNamespaceCommand::Create => cloud_namespace::create_namespace().await, + CloudNamespaceCommand::Delete { namespace } => { + cloud_namespace::delete_namespace(namespace).await + } + }; + } + if cloud && test_name != "integ_tests" { + bail!("Cloud filtering is only defined for the integ_tests target"); + } + if cloud && test_executable.is_some() { + bail!("Cloud filtering requires Cargo to build the test target with cloud-test-mode"); + } let cargo = env::var("CARGO").unwrap_or_else(|_| "cargo".to_string()); // Try building first, so that we error early on build failures & don't start server // Unclear why --all-features doesn't work here - let test_args_preamble = [ + let mut test_args_preamble = [ "test", "--features", "temporalio-common/serde_serialize", @@ -79,13 +124,15 @@ async fn main() -> Result<(), anyhow::Error> { "ephemeral-server", "--features", "temporalio-sdk-core/otel", - "--test", - &test_name, ] .into_iter() .map(ToString::to_string) - .chain(cargo_test_args) .collect::>(); + if cloud { + test_args_preamble.extend(["--features".to_owned(), "cloud-test-mode".to_owned()]); + } + test_args_preamble.extend(["--test".to_owned(), test_name.clone()]); + test_args_preamble.extend(cargo_test_args); if test_executable.is_none() { let mut build_cmd = Command::new(&cargo); strip_cargo_env_vars(&mut build_cmd); @@ -100,7 +147,6 @@ async fn main() -> Result<(), anyhow::Error> { if just_build { return Ok(()); } - let (server, envs) = match server_kind { ServerKind::TemporalCLI => { let config = @@ -135,6 +181,12 @@ async fn main() -> Result<(), anyhow::Error> { println!("========================================================"); (None, vec![]) } + ServerKind::EnvConfig => { + println!("========================================================"); + println!("Not starting up a server. Loading its configuration from envconfig."); + println!("========================================================"); + (None, vec![(TEST_ENV_CONFIG_SERVER_ENV_VAR, "true")]) + } }; let mut cmd = if let Some(test_executable) = test_executable { @@ -159,7 +211,12 @@ async fn main() -> Result<(), anyhow::Error> { format!("http://{}", &srv.target), ); } - let status = cmd.envs(envs).current_dir(project_root()).status().await?; + let status = cmd + .env_remove(TEST_ENV_CONFIG_SERVER_ENV_VAR) + .envs(envs) + .current_dir(project_root()) + .status() + .await?; if let Some(mut srv) = server { srv.shutdown().await?; diff --git a/crates/sdk-core/tests/shared_tests/mod.rs b/crates/sdk-core/tests/shared_tests/mod.rs index 61a6ae402..833710170 100644 --- a/crates/sdk-core/tests/shared_tests/mod.rs +++ b/crates/sdk-core/tests/shared_tests/mod.rs @@ -383,10 +383,10 @@ pub(crate) async fn shutdown_during_active_timer_activity_workflows() { let history = client .get_workflow_handle::(wf_id) .fetch_history(WorkflowFetchHistoryOptions::default()) + .into_events() .await .unwrap(); let bad_events: Vec<_> = history - .events() .iter() .filter(|e| match &e.attributes { Some(history_event::Attributes::WorkflowTaskFailedEventAttributes(f)) diff --git a/crates/sdk-core/tests/shared_tests/priority.rs b/crates/sdk-core/tests/shared_tests/priority.rs index dc02337dd..92ce162c1 100644 --- a/crates/sdk-core/tests/shared_tests/priority.rs +++ b/crates/sdk-core/tests/shared_tests/priority.rs @@ -19,11 +19,11 @@ pub(crate) async fn priority_values_sent_to_server() { } else { return; }; - starter.workflow_options.priority = Priority { - priority_key: Some(1), - fairness_key: Some("fair-wf".to_string()), - fairness_weight: Some(4.2), - }; + starter.workflow_options.priority = Priority::builder() + .priority_key(1) + .fairness_key("fair-wf") + .fairness_weight(4.2) + .build(); let child_type = "child-wf"; struct PriorityActivities; @@ -33,11 +33,11 @@ pub(crate) async fn priority_values_sent_to_server() { async fn echo(ctx: ActivityContext, echo_me: String) -> Result { assert_eq!( ctx.info().priority, - Priority { - priority_key: Some(5), - fairness_key: Some("fair-act".to_string()), - fairness_weight: Some(1.1) - } + Priority::builder() + .priority_key(5) + .fairness_key("fair-act") + .fairness_weight(1.1) + .build() ); Ok(echo_me) } @@ -60,11 +60,13 @@ pub(crate) async fn priority_values_sent_to_server() { RawValue::new(vec![]), ChildWorkflowOptions::builder() .workflow_id(format!("{}-child", ctx.task_queue())) - .priority(Priority { - priority_key: Some(4), - fairness_key: Some("fair-child".to_string()), - fairness_weight: Some(1.23), - }) + .priority( + Priority::builder() + .priority_key(4) + .fairness_key("fair-child") + .fairness_weight(1.23) + .build(), + ) .build(), ) .await?; @@ -72,11 +74,13 @@ pub(crate) async fn priority_values_sent_to_server() { PriorityActivities::echo, "hello".to_string(), ActivityOptions::with_start_to_close_timeout(Duration::from_secs(5)) - .priority(Priority { - priority_key: Some(5), - fairness_key: Some("fair-act".to_string()), - fairness_weight: Some(1.1), - }) + .priority( + Priority::builder() + .priority_key(5) + .fairness_key("fair-act") + .fairness_weight(1.1) + .build(), + ) .do_not_eagerly_execute(true) .build(), ); @@ -96,11 +100,11 @@ pub(crate) async fn priority_values_sent_to_server() { async fn run(ctx: &mut WorkflowContext) -> WorkflowResult<()> { assert_eq!( ctx.info().priority(), - Priority { - priority_key: Some(4), - fairness_key: Some("fair-child".to_string()), - fairness_weight: Some(1.23) - } + Priority::builder() + .priority_key(4) + .fairness_key("fair-child") + .fairness_weight(1.23) + .build() ); Ok(()) } @@ -131,9 +135,9 @@ pub(crate) async fn priority_values_sent_to_server() { .unwrap(); let events = handle .fetch_history(Default::default()) + .into_events() .await - .unwrap() - .into_events(); + .unwrap(); let workflow_init_event = events .iter() .find_map(|e| { diff --git a/crates/sdk-core/tests/wasm_workflow_tests.rs b/crates/sdk-core/tests/wasm_workflow_tests.rs index e597bf233..e6babde46 100644 --- a/crates/sdk-core/tests/wasm_workflow_tests.rs +++ b/crates/sdk-core/tests/wasm_workflow_tests.rs @@ -95,7 +95,10 @@ async fn wasm_patch_activation_callback_can_decline() { let callback_input = input.clone(); let callback: PatchActivationCallback = Arc::new(move |value| { callback_calls.fetch_add(1, Ordering::Relaxed); - *callback_input.lock().unwrap() = Some(value); + *callback_input.lock().unwrap() = Some(( + value.workflow_info.workflow_type().to_string(), + value.patch_id, + )); false }); @@ -108,11 +111,8 @@ async fn wasm_patch_activation_callback_can_decline() { assert_eq!(marker_count, 0); let input = input.lock().unwrap(); let input = input.as_ref().unwrap(); - assert_eq!( - input.workflow_info.workflow_type(), - WASM_PATCH_ACTIVATION_WORKFLOW_TYPE - ); - assert_eq!(input.patch_id, WASM_PATCH_ID); + assert_eq!(input.0, WASM_PATCH_ACTIVATION_WORKFLOW_TYPE); + assert_eq!(input.1, WASM_PATCH_ID); } #[tokio::test] @@ -205,7 +205,7 @@ async fn wasm_patch_activation_callback_panic_fails_workflow_task() { #[tokio::test] async fn wasm_task_failure_preserves_wit_failure_details() { - let component_path = build_wasm_hello_component().await; + let component_path = build_wasm_task_failure_component().await; let component = WasmWorkflowComponent::from_file(WASM_COMPONENT_ID, component_path) .expect("sample WASM component should be loadable"); @@ -389,6 +389,11 @@ async fn build_wasm_patch_activation_component() -> PathBuf { build_wasm_component(fixture_dir, "temporal_wasm_patch_activation_workflow.wasm").await } +async fn build_wasm_task_failure_component() -> PathBuf { + let fixture_dir = repository_root().join("crates/sdk-core/tests/fixtures/wasm_task_failure"); + build_wasm_component(fixture_dir, "temporal_wasm_task_failure_workflow.wasm").await +} + fn repository_root() -> PathBuf { PathBuf::from(env!("CARGO_MANIFEST_DIR")) .ancestors() diff --git a/crates/sdk-core/tests/workflows_procmacro.rs b/crates/sdk-core/tests/workflows_procmacro.rs index 2481a2304..90ba02dbc 100644 --- a/crates/sdk-core/tests/workflows_procmacro.rs +++ b/crates/sdk-core/tests/workflows_procmacro.rs @@ -1,6 +1,5 @@ #[test] fn workflows_procmacro_build_tests() { let t = trybuild::TestCases::new(); - t.pass("tests/workflows_trybuild/*_pass.rs"); t.compile_fail("tests/workflows_trybuild/*_fail.rs"); } diff --git a/crates/sdk-core/tests/workflows_trybuild/basic_pass.rs b/crates/sdk-core/tests/workflows_trybuild/basic_pass.rs deleted file mode 100644 index 785e8f389..000000000 --- a/crates/sdk-core/tests/workflows_trybuild/basic_pass.rs +++ /dev/null @@ -1,49 +0,0 @@ -use temporalio_macros::{workflow, workflow_methods}; -use temporalio_sdk::{SyncWorkflowContext, WorkflowContext, WorkflowContextView, WorkflowResult}; - -#[workflow] -pub struct MyWorkflow { - counter: u32, -} - -#[workflow_methods] -impl MyWorkflow { - #[init] - pub fn new(_ctx: &WorkflowContextView, _input: String) -> Self { - Self { counter: 0 } - } - - // Async run uses &self - #[run] - pub async fn run(_ctx: &mut WorkflowContext) -> WorkflowResult { - Ok("hi".to_owned()) - } - - // Sync signal uses &mut self - #[signal(name = "increment")] - pub fn increment_counter(&mut self, _ctx: &mut SyncWorkflowContext, amount: u32) { - self.counter += amount; - } - - #[signal] - pub async fn async_signal(_ctx: &mut WorkflowContext) {} - - // Query uses &self with read-only context - #[query] - pub fn get_counter(&self, _ctx: &WorkflowContextView) -> u32 { - self.counter - } - - #[update(name = "double")] - pub fn double_counter(&mut self, _ctx: &mut SyncWorkflowContext) -> u32 { - self.counter *= 2; - self.counter - } - - #[update] - pub async fn async_update(_ctx: &mut WorkflowContext, val: i32) -> i32 { - val * 2 - } -} - -fn main() {} diff --git a/crates/sdk-core/tests/workflows_trybuild/minimal_pass.rs b/crates/sdk-core/tests/workflows_trybuild/minimal_pass.rs deleted file mode 100644 index cde1467dd..000000000 --- a/crates/sdk-core/tests/workflows_trybuild/minimal_pass.rs +++ /dev/null @@ -1,21 +0,0 @@ -use temporalio_macros::{workflow, workflow_methods}; -use temporalio_sdk::{WorkflowContext, WorkflowResult}; - -#[workflow] -pub struct MinimalWorkflow; - -#[workflow_methods] -impl MinimalWorkflow { - #[run] - pub async fn run(_ctx: &mut WorkflowContext) -> WorkflowResult<()> { - Ok(()) - } -} - -impl Default for MinimalWorkflow { - fn default() -> Self { - Self - } -} - -fn main() {} diff --git a/crates/sdk/Cargo.toml b/crates/sdk/Cargo.toml index c925bfa96..2c4d61bcd 100644 --- a/crates/sdk/Cargo.toml +++ b/crates/sdk/Cargo.toml @@ -1,7 +1,8 @@ [package] name = "temporalio-sdk" -version = "0.6.0" +version = "1.0.0" edition = "2024" +rust-version = "1.92.0" authors = ["Spencer Judge "] license-file = { workspace = true } description = "Temporal Rust SDK" @@ -12,6 +13,9 @@ categories = ["development-tools"] readme = "README.md" autoexamples = false +[package.metadata.docs.rs] +features = ["experimental"] + [dependencies] async-trait = "0.1" anyhow = "1.0" @@ -39,29 +43,30 @@ tokio-stream = { version = "0.1", default-features = false } tracing = "0.1" uuid = { version = "1.18", default-features = false, features = ["v4"] } wasmtime = { version = "44", optional = true, features = ["component-model"] } +url = { version = "2.5", optional = true } [dependencies.temporalio-sdk-core] path = "../sdk-core" -version = "0.6" +version = "=0.9.0" default-features = false [dependencies.temporalio-workflow] path = "../workflow" -version = "0.6" +version = "~1.0.0" [dependencies.temporalio-common] path = "../common" -version = "0.6" +version = "~1.0.0" default-features = false [dependencies.temporalio-client] path = "../client" -version = "0.6" +version = "~1.0.0" default-features = false [dependencies.temporalio-macros] path = "../macros" -version = "0.6" +version = "~1.0.0" [dev-dependencies] futures = "0.3" @@ -70,12 +75,14 @@ rstest = "0.26" [features] default = ["envconfig", "prometheus"] envconfig = ["temporalio-sdk-core/envconfig"] +experimental = ["temporalio-client/experimental", "temporalio-workflow/experimental"] prometheus = ["temporalio-sdk-core/prometheus"] otel = ["temporalio-sdk-core/otel"] examples = ["serde/derive", "dep:serde_json", "envconfig"] wasm-examples = [] wasm-workflows = ["dep:wasmtime"] dynamic-tls = ["temporalio-client/dynamic-tls"] +testing = ["temporalio-sdk-core/ephemeral-server", "dep:url"] [dependencies.serde_json] version = "1" @@ -111,6 +118,11 @@ name = "activity-interceptor-starter" path = "examples/activity_interceptor/starter.rs" required-features = ["examples"] +[[example]] +name = "workflow-context" +path = "examples/workflow_context.rs" +required-features = ["examples"] + [[example]] name = "timer-examples-worker" path = "examples/timer_examples/worker.rs" @@ -221,6 +233,36 @@ name = "schedules-starter" path = "examples/schedules/starter.rs" required-features = ["examples"] +[[example]] +name = "standalone-activities-worker" +path = "examples/standalone_activities/worker.rs" +required-features = ["examples"] + +[[example]] +name = "standalone-activities-execute" +path = "examples/standalone_activities/execute_activity.rs" +required-features = ["examples"] + +[[example]] +name = "standalone-activities-start" +path = "examples/standalone_activities/start_activity.rs" +required-features = ["examples"] + +[[example]] +name = "standalone-activities-get-handle" +path = "examples/standalone_activities/get_activity_handle.rs" +required-features = ["examples"] + +[[example]] +name = "standalone-activities-list" +path = "examples/standalone_activities/list_activities.rs" +required-features = ["examples"] + +[[example]] +name = "standalone-activities-count" +path = "examples/standalone_activities/count_activities.rs" +required-features = ["examples"] + [[example]] name = "wasm-workflows" path = "examples/wasm_workflows/src/lib.rs" diff --git a/crates/sdk/README.md b/crates/sdk/README.md index d5ff9a191..acdea1514 100644 --- a/crates/sdk/README.md +++ b/crates/sdk/README.md @@ -3,11 +3,8 @@ [![crates.io](https://img.shields.io/crates/v/temporalio-sdk.svg)](https://crates.io/crates/temporalio-sdk) [![docs.rs](https://docs.rs/temporalio-sdk/badge.svg)](https://docs.rs/temporalio-sdk) -This crate contains a Public Preview Rust SDK. The SDK is built on top of -Core and provides a native Rust experience for writing Temporal workflows and activities. - -⚠️ **The SDK is in Public Preview and under active development.** The API can and -will continue to evolve. +This crate contains the Temporal Rust SDK. The SDK is built on top of Core and provides a native +Rust experience for writing Temporal workflows and activities. ## Quick Start @@ -88,7 +85,7 @@ use temporalio_sdk::{Runtime, Worker, WorkerOptions}; #[tokio::main] async fn main() -> Result<(), Box> { - let runtime = Runtime::new_assume_tokio(Default::default())?; + let runtime = Runtime::from_current_tokio(Default::default())?; let (conn_options, client_options) = ClientOptions::load_from_config( LoadClientConfigProfileOptions::default() )?; @@ -105,21 +102,44 @@ async fn main() -> Result<(), Box> { } ``` -## Crate Features +### Testing -The SDK enables a few convenience integrations by default. Users who want a smaller dependency -graph can disable defaults and opt back into the integrations they use: +Enable the `testing` feature to run activities directly or start an isolated Temporal CLI dev +server for workflow tests. + +Activity test inputs and outputs are ordinary Rust values. Register an activity implementer when +testing an instance activity: + +```rust +let env = ActivityEnvironment::builder() + .register_activities(MyActivities { counter: Default::default() }) + .build(); -```toml -temporalio-sdk = { version = "0.3", default-features = false, features = ["envconfig"] } +assert_eq!(env.run(MyActivities::greet, "Rust".to_owned()).await?, "Hello, Rust!"); ``` -- `envconfig` - enabled by default. Adds `ClientOptions::load_from_config` and related helpers for - loading connection settings from environment variables and `temporal.toml` files. -- `prometheus` - enabled by default. Adds the Prometheus metrics exporter in - `temporalio_common::telemetry` for serving SDK metrics from a HTTP endpoint. -- `otel` - optional. Adds the OpenTelemetry metrics exporter in `temporalio_common::telemetry` for - sending SDK metrics to an OpenTelemetry collector. +Workflow tests can use a local server with the normal client and worker APIs. Local environments +own their server and expose a consuming shutdown method: + +```rust +let env = WorkflowEnvironment::start_local(LocalWorkflowEnvironmentOptions::default()).await?; +let client = env.client().clone(); +// Construct workflow starters and workers with `client`. +env.shutdown().await?; +``` + +## Crate Features + +The SDK enables a few convenience integrations by default. Users who want a smaller dependency +graph can disable defaults and opt back into the integrations they use. + +- `envconfig`: Support for loading connection settings from environment variables and `temporal.toml` files. | +- `prometheus`: The Prometheus metrics exporter for `temporalio_common::telemetry`. | +- `otel`: The OpenTelemetry metrics exporter for `temporalio_common::telemetry`. | +- `experimental`: Rust SDK, client, and Workflow APIs that are still under development and may change or be removed. | +- `testing`: The `testing` module, direct activity test support, and local Temporal CLI dev-server lifecycle management. | +- `dynamic-tls`: Dynamic mTLS client-certificate resolution for transparent certificate rotation. | +- `wasm-workflows`: Support WebAssembly workflow components through Wasmtime for workers and workflow replay. | ## Workflows in detail diff --git a/crates/sdk/examples/README.md b/crates/sdk/examples/README.md index 48a943330..eca8a3783 100644 --- a/crates/sdk/examples/README.md +++ b/crates/sdk/examples/README.md @@ -32,6 +32,7 @@ See each example's README for details and expected output. - [Hello World](hello_world/) — Basic workflow that calls a single activity - [Activity Heartbeating](activity_heartbeating/) — Long-running activity with heartbeating and resume-on-retry - [Activity Inbound Interceptor](activity_interceptor/) — Wrapping inbound activity execution and inspecting typed inputs/outputs +- [Standalone Activities](standalone_activities/) — Activities started directly from a client, without a workflow - [Timer Examples](timer_examples/) — Workflow timers, racing timers against activities, and timer cancellation - [Message Passing](message_passing/) — Signals, queries, and updates on a workflow - [Child Workflows](child_workflows/) — Starting and collecting results from child workflows diff --git a/crates/sdk/examples/activity_heartbeating/worker.rs b/crates/sdk/examples/activity_heartbeating/worker.rs index a5279818f..c73cd90a0 100644 --- a/crates/sdk/examples/activity_heartbeating/worker.rs +++ b/crates/sdk/examples/activity_heartbeating/worker.rs @@ -8,7 +8,7 @@ use workflows::{HeartbeatingActivities, HeartbeatingWorkflow}; #[tokio::main] async fn main() -> Result<(), Box> { - let runtime = Runtime::new_assume_tokio(Default::default())?; + let runtime = Runtime::from_current_tokio(Default::default())?; let (conn_opts, client_opts) = ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; let connection = Connection::connect(conn_opts).await?; diff --git a/crates/sdk/examples/activity_interceptor/worker.rs b/crates/sdk/examples/activity_interceptor/worker.rs index 9503d8e72..30c40e467 100644 --- a/crates/sdk/examples/activity_interceptor/worker.rs +++ b/crates/sdk/examples/activity_interceptor/worker.rs @@ -52,7 +52,7 @@ impl ActivityInboundInterceptor for LoggingActivityInterceptor { #[tokio::main] async fn main() -> Result<(), Box> { - let runtime = Runtime::new_assume_tokio(Default::default())?; + let runtime = Runtime::from_current_tokio(Default::default())?; let (conn_opts, client_opts) = ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; let connection = Connection::connect(conn_opts).await?; diff --git a/crates/sdk/examples/cancellation/worker.rs b/crates/sdk/examples/cancellation/worker.rs index 0410af6dc..f7afde533 100644 --- a/crates/sdk/examples/cancellation/worker.rs +++ b/crates/sdk/examples/cancellation/worker.rs @@ -8,7 +8,7 @@ use workflows::{CancellationActivities, CancellationWorkflow}; #[tokio::main] async fn main() -> Result<(), Box> { - let runtime = Runtime::new_assume_tokio(Default::default())?; + let runtime = Runtime::from_current_tokio(Default::default())?; let (conn_opts, client_opts) = ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; let connection = Connection::connect(conn_opts).await?; diff --git a/crates/sdk/examples/child_workflows/worker.rs b/crates/sdk/examples/child_workflows/worker.rs index 3a8a94651..3bdbff138 100644 --- a/crates/sdk/examples/child_workflows/worker.rs +++ b/crates/sdk/examples/child_workflows/worker.rs @@ -8,7 +8,7 @@ use workflows::{GreetingChildWorkflow, ParentWorkflow}; #[tokio::main] async fn main() -> Result<(), Box> { - let runtime = Runtime::new_assume_tokio(Default::default())?; + let runtime = Runtime::from_current_tokio(Default::default())?; let (conn_opts, client_opts) = ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; let connection = Connection::connect(conn_opts).await?; diff --git a/crates/sdk/examples/continue_as_new/worker.rs b/crates/sdk/examples/continue_as_new/worker.rs index 5e68e63d0..817138035 100644 --- a/crates/sdk/examples/continue_as_new/worker.rs +++ b/crates/sdk/examples/continue_as_new/worker.rs @@ -8,7 +8,7 @@ use workflows::ContinueAsNewWorkflow; #[tokio::main] async fn main() -> Result<(), Box> { - let runtime = Runtime::new_assume_tokio(Default::default())?; + let runtime = Runtime::from_current_tokio(Default::default())?; let (conn_opts, client_opts) = ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; let connection = Connection::connect(conn_opts).await?; diff --git a/crates/sdk/examples/hello_world/worker.rs b/crates/sdk/examples/hello_world/worker.rs index 8480cc063..e067ddd4e 100644 --- a/crates/sdk/examples/hello_world/worker.rs +++ b/crates/sdk/examples/hello_world/worker.rs @@ -8,7 +8,7 @@ use workflows::{GreetingActivities, HelloWorldWorkflow}; #[tokio::main] async fn main() -> Result<(), Box> { - let runtime = Runtime::new_assume_tokio(Default::default())?; + let runtime = Runtime::from_current_tokio(Default::default())?; let (conn_opts, client_opts) = ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; let connection = Connection::connect(conn_opts).await?; diff --git a/crates/sdk/examples/local_activities/worker.rs b/crates/sdk/examples/local_activities/worker.rs index 684ac4a0d..8ace835b4 100644 --- a/crates/sdk/examples/local_activities/worker.rs +++ b/crates/sdk/examples/local_activities/worker.rs @@ -8,7 +8,7 @@ use workflows::{GreetingActivities, LocalActivitiesWorkflow}; #[tokio::main] async fn main() -> Result<(), Box> { - let runtime = Runtime::new_assume_tokio(Default::default())?; + let runtime = Runtime::from_current_tokio(Default::default())?; let (conn_opts, client_opts) = ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; let connection = Connection::connect(conn_opts).await?; diff --git a/crates/sdk/examples/message_passing/worker.rs b/crates/sdk/examples/message_passing/worker.rs index 66248ee87..8a931618c 100644 --- a/crates/sdk/examples/message_passing/worker.rs +++ b/crates/sdk/examples/message_passing/worker.rs @@ -8,7 +8,7 @@ use workflows::MessagePassingWorkflow; #[tokio::main] async fn main() -> Result<(), Box> { - let runtime = Runtime::new_assume_tokio(Default::default())?; + let runtime = Runtime::from_current_tokio(Default::default())?; let (conn_opts, client_opts) = ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; let connection = Connection::connect(conn_opts).await?; diff --git a/crates/sdk/examples/polling/worker.rs b/crates/sdk/examples/polling/worker.rs index 92ca74228..bc07e277a 100644 --- a/crates/sdk/examples/polling/worker.rs +++ b/crates/sdk/examples/polling/worker.rs @@ -8,7 +8,7 @@ use workflows::{PollingActivities, PollingWorkflow}; #[tokio::main] async fn main() -> Result<(), Box> { - let runtime = Runtime::new_assume_tokio(Default::default())?; + let runtime = Runtime::from_current_tokio(Default::default())?; let (conn_opts, client_opts) = ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; let connection = Connection::connect(conn_opts).await?; diff --git a/crates/sdk/examples/saga/worker.rs b/crates/sdk/examples/saga/worker.rs index e1fe5bb6d..bc35c6efd 100644 --- a/crates/sdk/examples/saga/worker.rs +++ b/crates/sdk/examples/saga/worker.rs @@ -8,7 +8,7 @@ use workflows::{BookingActivities, SagaWorkflow}; #[tokio::main] async fn main() -> Result<(), Box> { - let runtime = Runtime::new_assume_tokio(Default::default())?; + let runtime = Runtime::from_current_tokio(Default::default())?; let (conn_opts, client_opts) = ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; let connection = Connection::connect(conn_opts).await?; diff --git a/crates/sdk/examples/schedules/worker.rs b/crates/sdk/examples/schedules/worker.rs index 7ce66dd5d..9495052da 100644 --- a/crates/sdk/examples/schedules/worker.rs +++ b/crates/sdk/examples/schedules/worker.rs @@ -8,7 +8,7 @@ use workflows::{ScheduledActivities, ScheduledWorkflow}; #[tokio::main] async fn main() -> Result<(), Box> { - let runtime = Runtime::new_assume_tokio(Default::default())?; + let runtime = Runtime::from_current_tokio(Default::default())?; let (conn_opts, client_opts) = ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; let connection = Connection::connect(conn_opts).await?; diff --git a/crates/sdk/examples/search_attributes/worker.rs b/crates/sdk/examples/search_attributes/worker.rs index 1c1cf3371..42cc9c0dd 100644 --- a/crates/sdk/examples/search_attributes/worker.rs +++ b/crates/sdk/examples/search_attributes/worker.rs @@ -8,7 +8,7 @@ use workflows::SearchAttributesWorkflow; #[tokio::main] async fn main() -> Result<(), Box> { - let runtime = Runtime::new_assume_tokio(Default::default())?; + let runtime = Runtime::from_current_tokio(Default::default())?; let (conn_opts, client_opts) = ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; let connection = Connection::connect(conn_opts).await?; diff --git a/crates/sdk/examples/standalone_activities/README.md b/crates/sdk/examples/standalone_activities/README.md new file mode 100644 index 000000000..adbe20aa9 --- /dev/null +++ b/crates/sdk/examples/standalone_activities/README.md @@ -0,0 +1,50 @@ +# Standalone Activities + +This sample shows activities that run on their own, started directly from a client instead of being +orchestrated by a workflow. The activity and the worker are written exactly as they would be for a +workflow activity; only the way the activity is started differs. + +The client crate has no single "execute" call. To run an activity and wait for it, start it and then +await the handle's result. + +### Running this sample + +1. `temporal server start-dev` to start the Temporal server. Standalone activities require + Temporal CLI v1.9.1 or later. +2. In another terminal, start the worker: + +```bash + cargo run --features examples --example standalone-activities-worker +``` + +3. In another terminal, start an activity and wait for its result: + +```bash + cargo run --features examples --example standalone-activities-execute +``` + +It should print: + + Activity result: Hello, Temporal! + +### The other starters + +Start an activity without waiting for it, printing the activity and run IDs: + +```bash + cargo run --features examples --example standalone-activities-start +``` + +Get a handle to an activity started earlier, describe it, and read its result. Run one of the two +starters above first, since this looks up the activity ID they use: + +```bash + cargo run --features examples --example standalone-activities-get-handle +``` + +List and count the activity executions on this sample's task queue: + +```bash + cargo run --features examples --example standalone-activities-list + cargo run --features examples --example standalone-activities-count +``` diff --git a/crates/sdk/examples/standalone_activities/activities.rs b/crates/sdk/examples/standalone_activities/activities.rs new file mode 100644 index 000000000..688fa0d11 --- /dev/null +++ b/crates/sdk/examples/standalone_activities/activities.rs @@ -0,0 +1,17 @@ +#![allow(unreachable_pub)] +use temporalio_macros::activities; +use temporalio_sdk::activities::{ActivityContext, ActivityError}; + +pub struct GreetingActivities; + +#[activities] +impl GreetingActivities { + #[activity] + pub async fn compose_greeting( + _ctx: ActivityContext, + input: (String, String), + ) -> Result { + let (greeting, name) = input; + Ok(format!("{greeting}, {name}!")) + } +} diff --git a/crates/sdk/examples/standalone_activities/count_activities.rs b/crates/sdk/examples/standalone_activities/count_activities.rs new file mode 100644 index 000000000..f9bd76d63 --- /dev/null +++ b/crates/sdk/examples/standalone_activities/count_activities.rs @@ -0,0 +1,23 @@ +use temporalio_client::{ + Client, ClientOptions, Connection, envconfig::LoadClientConfigProfileOptions, +}; + +#[tokio::main] +async fn main() -> Result<(), Box> { + let (conn_opts, client_opts) = + ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; + let connection = Connection::connect(conn_opts).await?; + let client = Client::new(connection, client_opts)?; + + let count = client + .count_activities("TaskQueue = 'standalone-activities'", Default::default()) + .await?; + + println!("Total: {}", count.count()); + // Non-empty only when the query has a GROUP BY clause. + for group in count.groups() { + println!(" {:?} => {}", group.get::(0), group.count()); + } + + Ok(()) +} diff --git a/crates/sdk/examples/standalone_activities/execute_activity.rs b/crates/sdk/examples/standalone_activities/execute_activity.rs new file mode 100644 index 000000000..e5fafe738 --- /dev/null +++ b/crates/sdk/examples/standalone_activities/execute_activity.rs @@ -0,0 +1,38 @@ +mod activities; + +use std::time::Duration; + +use activities::GreetingActivities; +use temporalio_client::{ + ActivityStartOptions, Client, ClientOptions, Connection, + envconfig::LoadClientConfigProfileOptions, +}; + +#[tokio::main] +async fn main() -> Result<(), Box> { + let (conn_opts, client_opts) = + ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; + let connection = Connection::connect(conn_opts).await?; + let client = Client::new(connection, client_opts)?; + + let options = ActivityStartOptions::with_start_to_close_timeout( + "standalone-activities", + "standalone-activity-id", + Duration::from_secs(10), + ) + .build(); + + // There is no single "execute" call: start the activity, then await its result. + let handle = client + .start_activity( + GreetingActivities::compose_greeting, + ("Hello".to_string(), "Temporal".to_string()), + options, + ) + .await?; + + let result = handle.result().await?; + println!("Activity result: {result}"); + + Ok(()) +} diff --git a/crates/sdk/examples/standalone_activities/get_activity_handle.rs b/crates/sdk/examples/standalone_activities/get_activity_handle.rs new file mode 100644 index 000000000..28a0869ca --- /dev/null +++ b/crates/sdk/examples/standalone_activities/get_activity_handle.rs @@ -0,0 +1,37 @@ +mod activities; + +use activities::GreetingActivities; +use temporalio_client::{ + ActivityDescribeOptions, ActivityExecutionInfoLike, Client, ClientOptions, Connection, + envconfig::LoadClientConfigProfileOptions, +}; + +#[tokio::main] +async fn main() -> Result<(), Box> { + let (conn_opts, client_opts) = + ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; + let connection = Connection::connect(conn_opts).await?; + let client = Client::new(connection, client_opts)?; + + // Passing `None` for the run ID targets the latest run with this activity ID. + let handle = client.get_activity_handle( + GreetingActivities::compose_greeting, + "standalone-activity-id", + None, + ); + + let description = handle + .describe( + ActivityDescribeOptions::builder() + .include_outcome(true) + .build(), + ) + .await?; + + println!("Status: {:?}", description.status()); + println!("Type: {}", description.activity_type()); + println!("Attempt: {}", description.attempt()); + println!("Activity result: {}", handle.result().await?); + + Ok(()) +} diff --git a/crates/sdk/examples/standalone_activities/list_activities.rs b/crates/sdk/examples/standalone_activities/list_activities.rs new file mode 100644 index 000000000..10606ff54 --- /dev/null +++ b/crates/sdk/examples/standalone_activities/list_activities.rs @@ -0,0 +1,30 @@ +use futures::StreamExt; +use temporalio_client::{ + ActivityExecutionInfoLike, Client, ClientOptions, Connection, + envconfig::LoadClientConfigProfileOptions, +}; + +#[tokio::main] +async fn main() -> Result<(), Box> { + let (conn_opts, client_opts) = + ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; + let connection = Connection::connect(conn_opts).await?; + let client = Client::new(connection, client_opts)?; + + // List Standalone Activity Executions on this Task Queue. Only Standalone Activity Executions + // are returned -- Activities scheduled inside Workflows are not. The stream pages lazily. + let mut executions = + client.list_activities("TaskQueue = 'standalone-activities'", Default::default()); + + while let Some(execution) = executions.next().await { + let execution = execution?; + println!( + "{} {} {:?}", + execution.activity_id(), + execution.activity_type(), + execution.status() + ); + } + + Ok(()) +} diff --git a/crates/sdk/examples/standalone_activities/start_activity.rs b/crates/sdk/examples/standalone_activities/start_activity.rs new file mode 100644 index 000000000..9202120f9 --- /dev/null +++ b/crates/sdk/examples/standalone_activities/start_activity.rs @@ -0,0 +1,41 @@ +mod activities; + +use std::time::Duration; + +use activities::GreetingActivities; +use temporalio_client::{ + ActivityStartOptions, Client, ClientOptions, Connection, + envconfig::LoadClientConfigProfileOptions, +}; + +#[tokio::main] +async fn main() -> Result<(), Box> { + let (conn_opts, client_opts) = + ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; + let connection = Connection::connect(conn_opts).await?; + let client = Client::new(connection, client_opts)?; + + let options = ActivityStartOptions::with_start_to_close_timeout( + "standalone-activities", + "standalone-activity-id", + Duration::from_secs(10), + ) + .build(); + + // Returns as soon as the server has durably enqueued the activity. + let handle = client + .start_activity( + GreetingActivities::compose_greeting, + ("Hello".to_string(), "Temporal".to_string()), + options, + ) + .await?; + + println!( + "Started activity, id: {} run_id: {:?}", + handle.activity_id(), + handle.run_id() + ); + + Ok(()) +} diff --git a/crates/sdk/examples/standalone_activities/worker.rs b/crates/sdk/examples/standalone_activities/worker.rs new file mode 100644 index 000000000..ee52f0ead --- /dev/null +++ b/crates/sdk/examples/standalone_activities/worker.rs @@ -0,0 +1,30 @@ +// The worker registers the activity but never calls it by name; the client programs do. +#![allow(dead_code)] + +mod activities; + +use activities::GreetingActivities; +use temporalio_client::{ + Client, ClientOptions, Connection, envconfig::LoadClientConfigProfileOptions, +}; +use temporalio_sdk::{Runtime, Worker, WorkerOptions}; + +#[tokio::main] +async fn main() -> Result<(), Box> { + let runtime = Runtime::from_current_tokio(Default::default())?; + let (conn_opts, client_opts) = + ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; + let connection = Connection::connect(conn_opts).await?; + let client = Client::new(connection, client_opts)?; + + // A Worker that only runs Standalone Activities needs no registered workflows. + let worker_options = WorkerOptions::new("standalone-activities") + .register_activities(GreetingActivities) + .build(); + + let mut worker = Worker::new(&runtime, client, worker_options)?; + println!("Worker started on task queue: standalone-activities"); + worker.run().await?; + + Ok(()) +} diff --git a/crates/sdk/examples/timer_examples/worker.rs b/crates/sdk/examples/timer_examples/worker.rs index 03760b56b..f16bab808 100644 --- a/crates/sdk/examples/timer_examples/worker.rs +++ b/crates/sdk/examples/timer_examples/worker.rs @@ -8,7 +8,7 @@ use workflows::{TimerActivities, TimerWorkflow}; #[tokio::main] async fn main() -> Result<(), Box> { - let runtime = Runtime::new_assume_tokio(Default::default())?; + let runtime = Runtime::from_current_tokio(Default::default())?; let (conn_opts, client_opts) = ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; let connection = Connection::connect(conn_opts).await?; diff --git a/crates/sdk/examples/updatable_timer/worker.rs b/crates/sdk/examples/updatable_timer/worker.rs index 8a038f22b..2937370a0 100644 --- a/crates/sdk/examples/updatable_timer/worker.rs +++ b/crates/sdk/examples/updatable_timer/worker.rs @@ -8,7 +8,7 @@ use workflows::UpdatableTimerWorkflow; #[tokio::main] async fn main() -> Result<(), Box> { - let runtime = Runtime::new_assume_tokio(Default::default())?; + let runtime = Runtime::from_current_tokio(Default::default())?; let (conn_opts, client_opts) = ClientOptions::load_from_config(LoadClientConfigProfileOptions::default())?; let connection = Connection::connect(conn_opts).await?; diff --git a/crates/sdk/examples/wasm_workflows/src/lib.rs b/crates/sdk/examples/wasm_workflows/src/lib.rs index 966519a07..c7dd0e607 100644 --- a/crates/sdk/examples/wasm_workflows/src/lib.rs +++ b/crates/sdk/examples/wasm_workflows/src/lib.rs @@ -1,21 +1,4 @@ -use temporalio_workflow::{ - WorkflowContext, WorkflowResult, - common::protos::temporal::api::{ - enums::v1::WorkflowTaskFailedCause, - failure::v1::{ApplicationFailureInfo, Failure, failure::FailureInfo}, - }, - component::{StaticWorkflowComponent, instantiate_component_workflow}, - runtime::{ - guest::WorkflowInstance, - host::WorkflowHost, - types::{ - ActivationJobResult, ActivationResult, MAIN_ROUTINE_ID, MainRoutineCompletion, - RoutineCompletion, RoutinePollResult, TaskFailure, WorkflowDefinitionDescriptor, - WorkflowFailure, WorkflowInit, - }, - }, - workflow, workflow_methods, -}; +use temporalio_workflow::{WorkflowContext, WorkflowResult, workflow, workflow_methods}; #[workflow] #[derive(Default)] @@ -29,96 +12,4 @@ impl HelloWorkflow { } } -struct WasmTaskFailureWorkflow; - -impl WorkflowInstance for WasmTaskFailureWorkflow { - fn activate( - &mut self, - activation: temporalio_workflow::runtime::types::WorkflowActivation, - _waker: &std::task::Waker, - ) -> Result { - Ok(ActivationResult { - job_results: activation - .jobs - .iter() - .map(|_| ActivationJobResult::None) - .collect(), - }) - } - - fn poll_routine( - &mut self, - routine_id: u64, - _waker: &std::task::Waker, - ) -> Result { - if routine_id != MAIN_ROUTINE_ID { - return Err(Box::new(Failure { - message: format!("unexpected routine id {routine_id}"), - ..Default::default() - })); - } - - Ok(RoutinePollResult { - completion: Some(RoutineCompletion::Main(MainRoutineCompletion::TaskFailed( - TaskFailure { - failure: Box::new(Failure { - message: "structured wasm workflow task failure".to_string(), - failure_info: Some(FailureInfo::ApplicationFailureInfo( - ApplicationFailureInfo { - r#type: "WasmTaskFailure".to_string(), - non_retryable: true, - ..Default::default() - }, - )), - ..Default::default() - }), - force_cause: Some(WorkflowTaskFailedCause::NonDeterministicError as u32), - }, - ))), - made_progress: true, - pending_state: None, - }) - } -} - -struct WasmTestWorkflowModule; - -impl StaticWorkflowComponent for WasmTestWorkflowModule { - fn list_workflows() -> Vec { - vec![ - ::definition(), - WorkflowDefinitionDescriptor { - workflow_type: "WasmTaskFailureWorkflow".to_string(), - has_init: false, - init_takes_input: false, - signals: vec![], - queries: vec![], - updates: vec![], - }, - ] - } - - fn instantiate_workflow( - workflow_type: &str, - init: WorkflowInit, - host: std::rc::Rc, - ) -> Result, WorkflowFailure> { - match workflow_type { - name if name - == ::name() => - { - instantiate_component_workflow::(init, host) - } - "WasmTaskFailureWorkflow" => Ok(Box::new(WasmTaskFailureWorkflow)), - _ => Err(Box::new(Failure { - message: format!("No workflow named '{workflow_type}' exported by this component"), - ..Default::default() - })), - } - } -} - -type WasmTestWorkflowComponentExport = - temporalio_workflow::component::ExportedComponent; - -temporalio_workflow::__temporalio_export_workflow_component!(WasmTestWorkflowComponentExport); +temporalio_workflow::export_workflow_module!([HelloWorkflow]); diff --git a/crates/sdk/examples/workflow_context.rs b/crates/sdk/examples/workflow_context.rs new file mode 100644 index 000000000..75c2aae51 --- /dev/null +++ b/crates/sdk/examples/workflow_context.rs @@ -0,0 +1,83 @@ +//! Establish application context in workflow code and observe it in an outbound interceptor. + +use std::{sync::Arc, time::Duration}; +use temporalio_common::protos::temporal::api::common::v1::Payload; +use temporalio_macros::{activities, workflow, workflow_methods}; +use temporalio_sdk::{ + ActivityOptions, WorkflowContext, WorkflowContextKey, WorkflowResult, + activities::{ActivityContext, ActivityError}, + workflow_interceptors::{ + CancellableWorkflowOutboundFuture, ScheduleActivityInput, ScheduleActivityResult, + WorkflowInterceptor, WorkflowInterceptorContext, WorkflowNext, + }, +}; + +struct CurrentSpan; + +impl WorkflowContextKey for CurrentSpan { + type Value = String; +} + +#[workflow] +#[derive(Default)] +struct ContextWorkflow; + +#[workflow_methods] +impl ContextWorkflow { + #[run] + async fn run(ctx: &mut WorkflowContext, name: String) -> WorkflowResult { + let scoped_ctx = ctx.clone(); + ctx.with_context_value::(format!("greet-{name}"), async move { + scoped_ctx + .execute_activity( + GreetingActivities::greet, + name, + ActivityOptions::start_to_close_timeout(Duration::from_secs(10)), + ) + .await + .map_err(Into::into) + }) + .await + } +} + +struct GreetingActivities; + +#[activities] +impl GreetingActivities { + #[activity] + async fn greet(_ctx: ActivityContext, name: String) -> Result { + Ok(format!("Hello, {name}!")) + } +} + +struct ContextHeaderInterceptor; + +impl WorkflowInterceptor for ContextHeaderInterceptor { + fn schedule_activity( + &self, + ctx: WorkflowInterceptorContext, + mut input: ScheduleActivityInput, + next: WorkflowNext< + 'static, + ScheduleActivityInput, + CancellableWorkflowOutboundFuture, + >, + ) -> CancellableWorkflowOutboundFuture { + if let Some(span) = ctx.context_value::() { + input.headers_mut().insert( + "example-span".to_owned(), + Payload { + metadata: [("encoding".to_owned(), b"binary/plain".to_vec())].into(), + data: span.as_bytes().to_vec(), + ..Default::default() + }, + ); + } + next.run(input) + } +} + +fn main() { + let _interceptor: Arc = Arc::new(ContextHeaderInterceptor); +} diff --git a/crates/sdk/src/activities.rs b/crates/sdk/src/activities.rs index 381cea3e8..f2164b270 100644 --- a/crates/sdk/src/activities.rs +++ b/crates/sdk/src/activities.rs @@ -62,6 +62,8 @@ use futures_util::{ future::{BoxFuture, ready}, }; use prost_types::{Duration, Timestamp}; +#[cfg(feature = "testing")] +use std::any::Any; use std::{ collections::HashMap, fmt::Debug, @@ -72,11 +74,11 @@ use std::{ use temporalio_client::{Client, ClientOptions, Priority, WorkflowExecutionInfo, WorkflowHandle}; pub use temporalio_common::ActivityError; use temporalio_common::{ - ActivityDefinition, HasWorkflowDefinition, RetryPolicy, WorkflowExecution, + ActivityDefinition, HasWorkflowDefinition, RetryPolicy, data_converters::{ - DataConverter, DecodablePayloads, GenericPayloadConverter, PayloadConversionError, - PayloadConverter, RawValue, SerializationContext, SerializationContextData, - TemporalDeserializable, TemporalSerializable, + ActivitySerializationContext, DataConverter, DecodablePayloads, GenericPayloadConverter, + PayloadConversionError, PayloadConverter, RawValue, SerializationContext, + SerializationContextData, TemporalDeserializable, TemporalSerializable, }, error::ApplicationFailure, protos::{ @@ -88,18 +90,115 @@ use temporalio_common::{ use temporalio_sdk_core::Worker as CoreWorker; use tokio_util::sync::CancellationToken; +#[cfg(feature = "testing")] +pub(crate) type ActivityHeartbeatCallback = Arc) + Send + Sync>; + /// Used within activities to get info, heartbeat management etc. #[derive(Clone)] pub struct ActivityContext { - worker: Arc, - client_options: ClientOptions, + backend: ActivityContextBackend, cancellation_token: CancellationToken, heartbeat_details: ActivityHeartbeatDetails, header_fields: HashMap, info: ActivityInfo, } +#[derive(Clone)] +enum ActivityContextBackend { + Worker { + worker: Arc, + client_options: ClientOptions, + }, + #[cfg(feature = "testing")] + Test { + client: Option, + heartbeat_callback: Option, + }, +} + +impl ActivityContextBackend { + async fn record_heartbeat( + &self, + task_token: &[u8], + details: T, + ) -> Result<(), PayloadConversionError> + where + T: TemporalSerializable + 'static, + { + match self { + Self::Worker { + worker, + client_options, + } => { + let details = client_options + .data_converter + .to_payloads( + &SerializationContextData::Activity(ActivitySerializationContext::new()), + &details, + ) + .await?; + worker.record_activity_heartbeat(ActivityHeartbeat { + task_token: task_token.to_vec(), + details, + }); + } + #[cfg(feature = "testing")] + Self::Test { + heartbeat_callback, .. + } => { + if let Some(callback) = heartbeat_callback { + callback(Box::new(details)); + } + } + } + Ok(()) + } + + fn client(&self) -> Client { + match self { + Self::Worker { + worker, + client_options, + } => { + let connection = worker.get_client_connection().expect( + "activity context client is unavailable because the worker was not created from a Temporal client", + ); + Client::new(connection, client_options.clone()) + .expect("client construction from a worker connection should be infallible") + } + #[cfg(feature = "testing")] + Self::Test { client, .. } => client + .as_ref() + .expect("ActivityEnvironment was created without a Client. Pass one during construction to have one availalbe at runtime") + .clone(), + } + } +} + impl ActivityContext { + #[cfg(feature = "testing")] + pub(crate) fn new_for_test( + info: ActivityInfo, + header_fields: HashMap, + payload_converter: PayloadConverter, + cancellation_token: CancellationToken, + heartbeat_details: Vec, + client: Option, + heartbeat_callback: Option, + ) -> Self { + let heartbeat_details = ActivityHeartbeatDetails::new(heartbeat_details, payload_converter); + Self { + backend: ActivityContextBackend::Test { + client, + heartbeat_callback, + }, + cancellation_token, + heartbeat_details, + header_fields, + info, + } + } + pub(crate) fn new( worker: Arc, client_options: ClientOptions, @@ -139,20 +238,27 @@ impl ActivityContext { heartbeat_details, client_options.data_converter.payload_converter().clone(), ); + let (workflow_id, workflow_run_id) = workflow_execution + .map(|we| (we.workflow_id, we.run_id)) + .unzip(); + let activity_run_id = (workflow_id.is_none() && !run_id.is_empty()).then_some(run_id); ( ActivityContext { - worker, - client_options, + backend: ActivityContextBackend::Worker { + worker, + client_options, + }, cancellation_token, heartbeat_details, header_fields, info: ActivityInfo { task_token, task_queue, - workflow_type, - workflow_namespace, - workflow_execution: workflow_execution.map(Into::into), + workflow_type: (!workflow_type.is_empty()).then_some(workflow_type), + namespace: workflow_namespace, + workflow_id, + workflow_run_id, activity_id, activity_type, heartbeat_timeout: heartbeat_timeout.try_into_or_none(), @@ -165,7 +271,7 @@ impl ActivityContext { retry_policy: retry_policy.map(Into::into), is_local, priority: priority.map(Into::into).unwrap_or_default(), - run_id: (!run_id.is_empty()).then_some(run_id), + activity_run_id, }, }, input, @@ -195,15 +301,9 @@ impl ActivityContext { T: TemporalSerializable + 'static, { if !self.info.is_local { - let details = self - .client_options - .data_converter - .to_payloads(&SerializationContextData::Activity, &details) + self.backend + .record_heartbeat(&self.info.task_token, details) .await?; - self.worker.record_activity_heartbeat(ActivityHeartbeat { - task_token: self.info.task_token.clone(), - details, - }) } Ok(()) } @@ -215,27 +315,24 @@ impl ActivityContext { /// Return a client targeting the same Temporal service and namespace as this activity's worker. pub fn client(&self) -> Client { - let connection = self.worker.get_client_connection().expect( - "activity context client is unavailable because the worker was not created from a \ - Temporal client", - ); - Client::new(connection, self.client_options.clone()) - .expect("client construction from a worker connection should be infallible") + self.backend.client() } /// Return a workflow handle for the workflow execution that started this activity, if any. pub fn workflow_handle(&self) -> Option> { - let workflow_execution = self.info.workflow_execution.as_ref()?; - let run_id = (!workflow_execution.run_id().is_empty()) - .then(|| workflow_execution.run_id().to_owned()); + let workflow_id = self.info.workflow_id.clone()?; + let run_id = self.info.workflow_run_id.clone(); + let first_execution_run_id = run_id.clone(); + let client = self.client(); + Some(WorkflowHandle::new( - self.client(), - WorkflowExecutionInfo { - namespace: self.client_options.namespace.clone(), - workflow_id: workflow_execution.workflow_id().to_owned(), - run_id: run_id.clone(), - first_execution_run_id: run_id, - }, + client.clone(), + WorkflowExecutionInfo::builder() + .namespace(client.options().namespace.clone()) + .workflow_id(workflow_id) + .maybe_run_id(run_id) + .maybe_first_execution_run_id(first_execution_run_id) + .build(), )) } @@ -262,7 +359,7 @@ impl ActivityHeartbeatDetails { payloads: DecodablePayloads::new( payloads, payload_converter, - SerializationContextData::Activity, + SerializationContextData::Activity(ActivitySerializationContext::new()), ), } } @@ -295,12 +392,14 @@ impl ActivityHeartbeatDetails { pub struct ActivityInfo { /// An opaque token representing a specific Activity task. pub task_token: Vec, - /// The type of the workflow that invoked this activity. - pub workflow_type: String, - /// The namespace of the workflow that invoked this activity. - pub workflow_namespace: String, - /// The execution of the workflow that invoked this activity. - pub workflow_execution: Option, + /// The type of the workflow that invoked this activity. None for standalone activities. + pub workflow_type: Option, + /// The namespace of this activity. + pub namespace: String, + /// ID of the workflow that invoked this activity. None for standalone activities. + pub workflow_id: Option, + /// Run ID of the workflow that invoked this activity. None for standalone activities. + pub workflow_run_id: Option, /// The ID of this activity. pub activity_id: String, /// The type of this activity. @@ -326,7 +425,7 @@ pub struct ActivityInfo { /// Priority of this activity. If unset uses [Priority::default]. pub priority: Priority, /// Run ID of this activity execution. Only set for standalone activities. - pub run_id: Option, + pub activity_run_id: Option, } /// Deadline calculation. This is a port of @@ -417,15 +516,27 @@ fn call_execute_activity<'a>( } } -#[doc(hidden)] +/// Implemented by `#[activities]` for types that provide activity methods. +/// +/// This trait supports registration and direct execution infrastructure. Applications normally +/// use the generated implementation rather than implementing it manually. pub trait ActivityImplementer { + /// Register every activity method implemented by this type. fn register_all(self: Arc, defs: &mut ActivityDefinitions); } -#[doc(hidden)] +/// Direct execution support generated for each activity marker by `#[activities]`. +/// +/// Applications normally use the generated implementation rather than implementing this trait +/// manually. pub trait ExecutableActivity: ActivityDefinition + Sized { + /// Type containing the activity implementation. type Implementer: ActivityImplementer + Send + Sync + 'static; + /// Whether this activity requires an implementation instance. + const REQUIRES_INSTANCE: bool; + /// Return this activity's definition marker. fn definition() -> Self; + /// Execute the activity with already-typed input. fn execute( receiver: Option>, ctx: ActivityContext, @@ -433,9 +544,6 @@ pub trait ExecutableActivity: ActivityDefinition + Sized { ) -> BoxFuture<'static, Result>; } -#[doc(hidden)] -pub trait HasOnlyStaticMethods {} - /// Contains activity registrations in a form ready for execution by workers. #[derive(Default, Clone)] pub struct ActivityDefinitions { @@ -443,6 +551,7 @@ pub struct ActivityDefinitions { } impl ActivityDefinitions { + #[cfg(feature = "experimental")] pub(crate) fn extend(&mut self, other: &Self) { self.activities.extend(other.activities.clone()); } @@ -468,10 +577,9 @@ impl ActivityDefinitions { // Codec application happens at the SDK/Core boundary, so activity // implementations work with the payload converter directly. let pc = dc.payload_converter(); - let ctx = SerializationContext { - data: &SerializationContextData::Activity, - converter: pc, - }; + let context_data = + SerializationContextData::Activity(ActivitySerializationContext::new()); + let ctx = SerializationContext::new(&context_data, pc); let input: AD::Input = pc.from_payloads(&ctx, payloads)?; let input = ExecuteActivityInput::new(c, Box::new(input)); let leaf = activity_inbound_base::(instance); @@ -556,14 +664,20 @@ pub(crate) fn activity_error_to_core_result( ) -> ActivityExecutionResult { match err { ActivityError::Application(app) => ActivityExecutionResult::fail(dc.to_failure( - &SerializationContextData::Activity, + &SerializationContextData::Activity(ActivitySerializationContext::new()), OutgoingError::Activity(OutgoingActivityError::Application(app)), )), ActivityError::Cancelled { details } => ActivityExecutionResult::cancel(dc.to_failure( - &SerializationContextData::Activity, + &SerializationContextData::Activity(ActivitySerializationContext::new()), OutgoingError::Activity(OutgoingActivityError::Cancelled { details }), )), ActivityError::WillCompleteAsync => ActivityExecutionResult::will_complete_async(), + other => ActivityExecutionResult::fail(dc.to_failure( + &SerializationContextData::Activity(ActivitySerializationContext::new()), + OutgoingError::Activity(OutgoingActivityError::Application(Box::new( + ApplicationFailure::new(anyhow::anyhow!("Unsupported activity error: {other:?}")), + ))), + )), } } @@ -586,10 +700,10 @@ mod test { let payload_converter = PayloadConverter::default(); let payload = payload_converter .to_payload( - &SerializationContext { - data: &SerializationContextData::Activity, - converter: &payload_converter, - }, + &SerializationContext::new( + &SerializationContextData::Activity(ActivitySerializationContext::new()), + &payload_converter, + ), &"progress".to_owned(), ) .unwrap(); diff --git a/crates/sdk/src/error.rs b/crates/sdk/src/error.rs index f50ee9f97..38be11f34 100644 --- a/crates/sdk/src/error.rs +++ b/crates/sdk/src/error.rs @@ -1,8 +1,52 @@ //! Shared SDK error re-exports. pub use crate::workflow_registry::WorkflowRegistrationError; +#[cfg(feature = "experimental")] use temporalio_client::PluginApplyError; -pub use temporalio_sdk_core::WorkerValidationError; +use temporalio_sdk_core::WorkerValidationError as CoreWorkerValidationError; + +/// Errors that can occur while creating an SDK runtime. +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum RuntimeError { + /// Runtime initialization failed. + #[error("runtime initialization failed: {0}")] + Initialization(#[source] Box), + /// No Tokio runtime is active on the current thread. + #[error("no Tokio runtime is active on the current thread")] + NoCurrentTokioRuntime, +} + +impl RuntimeError { + pub(crate) fn from_core(error: anyhow::Error) -> Self { + Self::Initialization(error.into_boxed_dyn_error()) + } +} + +/// Errors encountered while validating a worker before polling begins. +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum WorkerValidationError { + /// The configured namespace could not be described. + #[error("namespace {namespace} was not found or otherwise could not be described: {source}")] + NamespaceDescribeError { + /// The underlying server error. + #[source] + source: temporalio_client::tonic::Status, + /// The namespace that could not be described. + namespace: String, + }, +} + +impl WorkerValidationError { + pub(crate) fn from_core(error: CoreWorkerValidationError) -> Self { + match error { + CoreWorkerValidationError::NamespaceDescribeError { source, namespace } => { + Self::NamespaceDescribeError { source, namespace } + } + } + } +} /// Errors that can occur while creating a worker. /// @@ -11,6 +55,7 @@ pub use temporalio_sdk_core::WorkerValidationError; #[non_exhaustive] pub enum WorkerCreateError { /// A plugin failed while configuring worker options. + #[cfg(feature = "experimental")] #[error(transparent)] Plugin(#[from] PluginApplyError), /// Worker initialization failed after plugin configuration completed. @@ -38,6 +83,7 @@ pub enum WorkerRunError { pub use temporalio_common::error::{ ActivityExecutionError, ApplicationErrorCategory, ApplicationFailure, - ChildWorkflowExecutionError, ChildWorkflowStartError, OutgoingActivityError, OutgoingError, - OutgoingWorkflowError, RetryState, TimeoutType, WorkflowSignalError, + CancelExternalWorkflowError, ChildWorkflowExecutionError, ChildWorkflowStartError, + OutgoingActivityError, OutgoingError, OutgoingWorkflowError, RetryState, TimeoutType, + WorkflowSignalError, }; diff --git a/crates/sdk/src/interceptors.rs b/crates/sdk/src/interceptors.rs index 166cd0240..809ae12c9 100644 --- a/crates/sdk/src/interceptors.rs +++ b/crates/sdk/src/interceptors.rs @@ -4,20 +4,17 @@ use crate::{ Worker, WorkerRunError, activities::{ActivityContext, ActivityError, ActivityInfo}, }; -use anyhow::bail; use futures_util::future::{BoxFuture, LocalBoxFuture}; -use std::{ - any::Any, - collections::HashMap, - sync::{Arc, OnceLock}, -}; +#[cfg(feature = "experimental")] +use std::sync::OnceLock; +use std::{any::Any, collections::HashMap, sync::Arc}; use temporalio_common::{ data_converters::{ GenericPayloadConverter, PayloadConversionError, SerializationContext, TemporalSerializable, }, protos::{ coresdk::{ - workflow_activation::{WorkflowActivation, remove_from_cache::EvictionReason}, + workflow_activation::WorkflowActivation, workflow_completion::WorkflowActivationCompletion, }, temporal::api::common::v1::Payload, @@ -47,49 +44,6 @@ mod activity_execution_value { } } -/// Implementors can intercept certain actions that happen within the Worker. -/// -/// Advanced usage only. -/// **Experimental:** This API may change or be removed. -#[async_trait::async_trait(?Send)] -pub trait WorkerInterceptor: Send + Sync { - /// Intercept the running of a worker. - fn run_worker<'a>( - &'a self, - input: RunWorkerInput<'a>, - next: Next<'a, RunWorkerInput<'a>, LocalBoxFuture<'a, Result<(), WorkerRunError>>>, - ) -> LocalBoxFuture<'a, Result<(), WorkerRunError>> { - next.run(input) - } - - /// Intercept the running of a worker created for workflow replay. - fn with_workflow_replay_worker<'a>( - &'a self, - input: WithWorkflowReplayWorkerInput<'a>, - next: Next< - 'a, - WithWorkflowReplayWorkerInput<'a>, - LocalBoxFuture<'a, Result<(), WorkerRunError>>, - >, - ) -> LocalBoxFuture<'a, Result<(), WorkerRunError>> { - next.run(input) - } - - /// Called every time a workflow activation completes (just before sending the completion to - /// core). - async fn on_workflow_activation_completion(&self, _completion: &WorkflowActivationCompletion) {} - /// Called after the worker has initiated shutdown and the workflow/activity polling loops - /// have exited, but just before waiting for the inner core worker shutdown - fn on_shutdown(&self, _sdk_worker: &Worker) {} - /// Called every time a workflow is about to be activated - async fn on_workflow_activation( - &self, - _activation: &WorkflowActivation, - ) -> Result<(), anyhow::Error> { - Ok(()) - } -} - /// Continuation for an interceptor operation. /// /// Interceptor implementations call [`Next::run`] to invoke the next step of the chain. @@ -97,76 +51,137 @@ pub struct Next<'a, I, O> { inner: Box O + Send + 'a>, } -/// Input to [`WorkerInterceptor::run_worker`]. -#[derive(Debug)] -#[non_exhaustive] -pub struct RunWorkerInput<'a> { - /// The worker being run. - pub worker: &'a mut Worker, -} +impl<'a, I, O> Next<'a, I, O> { + pub(crate) fn new(f: impl FnOnce(I) -> O + Send + 'a) -> Self { + Self { inner: Box::new(f) } + } -impl<'a> RunWorkerInput<'a> { - pub(crate) fn new(worker: &'a mut Worker) -> Self { - Self { worker } + /// Continue the call chain with the provided input. + pub fn run(self, input: I) -> O { + (self.inner)(input) } } -/// Input to [`WorkerInterceptor::with_workflow_replay_worker`]. -#[derive(Debug)] -#[non_exhaustive] -pub struct WithWorkflowReplayWorkerInput<'a> { - /// The worker created for this replay operation. - pub worker: &'a mut Worker, -} +#[cfg_attr(not(feature = "experimental"), allow(unreachable_pub))] +mod worker_lifecycle { + use super::*; + + /// Implementors can intercept certain actions that happen within the Worker. + /// + /// Advanced usage only. + /// **Experimental:** This API may change or be removed. + #[async_trait::async_trait(?Send)] + pub trait WorkerInterceptor: Send + Sync { + /// Intercept the running of a worker. + fn run_worker<'a>( + &'a self, + input: RunWorkerInput<'a>, + next: Next<'a, RunWorkerInput<'a>, LocalBoxFuture<'a, Result<(), WorkerRunError>>>, + ) -> LocalBoxFuture<'a, Result<(), WorkerRunError>> { + next.run(input) + } + + /// Intercept the running of a worker created for workflow replay. + fn with_workflow_replay_worker<'a>( + &'a self, + input: WithWorkflowReplayWorkerInput<'a>, + next: Next< + 'a, + WithWorkflowReplayWorkerInput<'a>, + LocalBoxFuture<'a, Result<(), WorkerRunError>>, + >, + ) -> LocalBoxFuture<'a, Result<(), WorkerRunError>> { + next.run(input) + } -impl<'a> WithWorkflowReplayWorkerInput<'a> { - pub(crate) fn new(worker: &'a mut Worker) -> Self { - Self { worker } + /// Called every time a workflow activation completes (just before sending the completion to + /// core). + async fn on_workflow_activation_completion( + &self, + _completion: &WorkflowActivationCompletion, + ) { + } + /// Called after the worker has initiated shutdown and the workflow/activity polling loops + /// have exited, but just before waiting for the inner core worker shutdown + fn on_shutdown(&self, _sdk_worker: &Worker) {} + /// Called every time a workflow is about to be activated + async fn on_workflow_activation( + &self, + _activation: &WorkflowActivation, + ) -> Result<(), anyhow::Error> { + Ok(()) + } } -} -pub(crate) fn call_run_worker<'a>( - interceptors: &'a [Arc], - input: RunWorkerInput<'a>, - terminal: Next<'a, RunWorkerInput<'a>, LocalBoxFuture<'a, Result<(), WorkerRunError>>>, -) -> LocalBoxFuture<'a, Result<(), WorkerRunError>> { - if let Some((interceptor, remaining)) = interceptors.split_first() { - let next = Next::new(move |input| call_run_worker(remaining, input, terminal)); - interceptor.run_worker(input, next) - } else { - terminal.run(input) + /// Input to [`WorkerInterceptor::run_worker`]. + #[derive(Debug)] + #[non_exhaustive] + pub struct RunWorkerInput<'a> { + /// The worker being run. + pub worker: &'a mut Worker, } -} -pub(crate) fn call_with_workflow_replay_worker<'a>( - interceptors: &'a [Arc], - input: WithWorkflowReplayWorkerInput<'a>, - terminal: Next< - 'a, - WithWorkflowReplayWorkerInput<'a>, - LocalBoxFuture<'a, Result<(), WorkerRunError>>, - >, -) -> LocalBoxFuture<'a, Result<(), WorkerRunError>> { - if let Some((interceptor, remaining)) = interceptors.split_first() { - let next = - Next::new(move |input| call_with_workflow_replay_worker(remaining, input, terminal)); - interceptor.with_workflow_replay_worker(input, next) - } else { - terminal.run(input) + impl<'a> RunWorkerInput<'a> { + pub(crate) fn new(worker: &'a mut Worker) -> Self { + Self { worker } + } } -} -impl<'a, I, O> Next<'a, I, O> { - pub(crate) fn new(f: impl FnOnce(I) -> O + Send + 'a) -> Self { - Self { inner: Box::new(f) } + /// Input to [`WorkerInterceptor::with_workflow_replay_worker`]. + #[derive(Debug)] + #[non_exhaustive] + pub struct WithWorkflowReplayWorkerInput<'a> { + /// The worker created for this replay operation. + pub worker: &'a mut Worker, } - /// Continue the call chain with the provided input. - pub fn run(self, input: I) -> O { - (self.inner)(input) + impl<'a> WithWorkflowReplayWorkerInput<'a> { + pub(crate) fn new(worker: &'a mut Worker) -> Self { + Self { worker } + } + } + + pub(crate) fn call_run_worker<'a>( + interceptors: &'a [Arc], + input: RunWorkerInput<'a>, + terminal: Next<'a, RunWorkerInput<'a>, LocalBoxFuture<'a, Result<(), WorkerRunError>>>, + ) -> LocalBoxFuture<'a, Result<(), WorkerRunError>> { + if let Some((interceptor, remaining)) = interceptors.split_first() { + let next = Next::new(move |input| call_run_worker(remaining, input, terminal)); + interceptor.run_worker(input, next) + } else { + terminal.run(input) + } + } + + pub(crate) fn call_with_workflow_replay_worker<'a>( + interceptors: &'a [Arc], + input: WithWorkflowReplayWorkerInput<'a>, + terminal: Next< + 'a, + WithWorkflowReplayWorkerInput<'a>, + LocalBoxFuture<'a, Result<(), WorkerRunError>>, + >, + ) -> LocalBoxFuture<'a, Result<(), WorkerRunError>> { + if let Some((interceptor, remaining)) = interceptors.split_first() { + let next = Next::new(move |input| { + call_with_workflow_replay_worker(remaining, input, terminal) + }); + interceptor.with_workflow_replay_worker(input, next) + } else { + terminal.run(input) + } } } +#[cfg(not(feature = "experimental"))] +pub(crate) use worker_lifecycle::{ + RunWorkerInput, WithWorkflowReplayWorkerInput, WorkerInterceptor, +}; +#[cfg(feature = "experimental")] +pub use worker_lifecycle::{RunWorkerInput, WithWorkflowReplayWorkerInput, WorkerInterceptor}; +pub(crate) use worker_lifecycle::{call_run_worker, call_with_workflow_replay_worker}; + /// Activity execution data passed to [`ActivityInboundInterceptor::execute_activity`]. #[non_exhaustive] pub struct ExecuteActivityInput { @@ -260,31 +275,14 @@ pub trait ActivityInboundInterceptor: Send + Sync + 'static { } } -/// An interceptor which causes the worker's run function to exit early if nondeterminism errors are -/// encountered -pub struct FailOnNondeterminismInterceptor {} -#[async_trait::async_trait(?Send)] -impl WorkerInterceptor for FailOnNondeterminismInterceptor { - async fn on_workflow_activation( - &self, - activation: &WorkflowActivation, - ) -> Result<(), anyhow::Error> { - if matches!( - activation.eviction_reason(), - Some(EvictionReason::Nondeterminism) - ) { - bail!("Workflow is being evicted because of nondeterminism! {activation}"); - } - Ok(()) - } -} - /// An interceptor that allows you to fetch the exit value of the workflow if and when it is set +#[cfg(feature = "experimental")] #[derive(Default)] pub struct ReturnWorkflowExitValueInterceptor { result_value: Arc>, } +#[cfg(feature = "experimental")] impl ReturnWorkflowExitValueInterceptor { /// Can be used to fetch the workflow result if/when it is determined pub fn result_handle(&self) -> Arc> { @@ -293,6 +291,7 @@ impl ReturnWorkflowExitValueInterceptor { } #[async_trait::async_trait(?Send)] +#[cfg(feature = "experimental")] impl WorkerInterceptor for ReturnWorkflowExitValueInterceptor { async fn on_workflow_activation_completion(&self, c: &WorkflowActivationCompletion) { if let Some(v) = c.complete_workflow_execution_value() { diff --git a/crates/sdk/src/lib.rs b/crates/sdk/src/lib.rs index 3dfebee73..803f1fa68 100644 --- a/crates/sdk/src/lib.rs +++ b/crates/sdk/src/lib.rs @@ -1,12 +1,11 @@ +#![cfg_attr(docsrs, feature(doc_cfg))] #![warn(missing_docs)] // error if there are missing docs -//! This crate defines a Public Preview Temporal Rust SDK. +//! This crate defines the Temporal Rust SDK. //! //! The SDK is built on top of Core and provides a native Rust experience for writing Temporal //! Workflows and Activities. //! -//! The SDK is in Public Preview and under active development. The API can and will continue to evolve. -//! //! An example of running an activity worker: //! ```no_run //! use std::str::FromStr; @@ -35,16 +34,18 @@ //! async fn main() -> Result<(), Box> { //! let connection_options = //! ConnectionOptions::new(Url::from_str("http://localhost:7233")?).build(); -//! let runtime = Runtime::new_assume_tokio(Default::default())?; +//! let runtime = Runtime::from_current_tokio(Default::default())?; //! let connection = Connection::connect(connection_options).await?; //! let client = Client::new(connection, ClientOptions::new("my_namespace").build())?; //! //! let worker_options = WorkerOptions::new("task_queue") //! .deployment_options( -//! WorkerDeploymentOptions::new(WorkerDeploymentVersion { -//! deployment_name: "my_deployment".to_owned(), -//! build_id: "my_build_id".to_owned(), -//! }) +//! WorkerDeploymentOptions::new( +//! WorkerDeploymentVersion::builder() +//! .deployment_name("my_deployment") +//! .build_id("my_build_id") +//! .build(), +//! ) //! .build(), //! ) //! .register_activities(MyActivities) @@ -64,9 +65,12 @@ extern crate self as temporalio_sdk; pub mod activities; pub mod error; pub mod interceptors; +#[cfg(feature = "experimental")] /// Experimental APIs for configuring clients and workers with reusable plugins. pub mod plugins; pub mod runtime; +#[cfg(feature = "testing")] +pub mod testing; mod workflow_executor; mod workflow_future; pub mod workflow_interceptors; @@ -77,31 +81,35 @@ pub mod workflow_replayer; mod workflow_wasm; pub mod workflows; +#[cfg(feature = "experimental")] +pub use crate::plugins::{ + ClientAndWorkerPlugin, SimplePlugin, SimplePluginBuilder, SimplePluginOption, WorkerPlugin, +}; pub use crate::{ error::{ - ActivityExecutionError, ApplicationFailure, ChildWorkflowExecutionError, - ChildWorkflowStartError, OutgoingActivityError, OutgoingError, OutgoingWorkflowError, - RetryState, TimeoutType, WorkerCreateError, WorkerRunError, WorkerValidationError, - WorkflowRegistrationError, WorkflowSignalError, - }, - plugins::{ - ClientAndWorkerPlugin, SimplePlugin, SimplePluginBuilder, SimplePluginOption, WorkerPlugin, - WorkflowDefinitions, + ActivityExecutionError, ApplicationFailure, CancelExternalWorkflowError, + ChildWorkflowExecutionError, ChildWorkflowStartError, OutgoingActivityError, OutgoingError, + OutgoingWorkflowError, RetryState, RuntimeError, TimeoutType, WorkerCreateError, + WorkerRunError, WorkerValidationError, WorkflowRegistrationError, WorkflowSignalError, }, + workflow_registry::WorkflowDefinitions, }; pub use runtime::Runtime; pub use temporalio_client::Namespace; pub use temporalio_workflow::{ ActivityCancellationType, ActivityCloseTimeouts, ActivityOptions, BaseWorkflowContext, CancellableFuture, CancellableFutureWithReason, ChildWorkflowCancellationType, - ChildWorkflowOptions, ContinueAsNewOptions, ContinueAsNewVersioningBehavior, - ExternalWorkflowHandle, LocalActivityOptions, MemoValue, NexusOperationCancellationType, - NexusOperationOptions, ParentClosePolicy, PatchActivationCallback, SignalWorkflowOptions, - StartChildWorkflowExecutionFailedCause, StartChildWorkflowOutput, StartedChildWorkflow, - StartedNexusOperation, SyncWorkflowContext, TimerOptions, TimerResult, VersioningIntent, - WaitConditionOptions, WorkflowCancellationError, WorkflowCancellationToken, WorkflowContext, - WorkflowContextView, WorkflowIdReusePolicy, WorkflowRandomValue, WorkflowResult, - WorkflowTermination, + ChildWorkflowOptions, ContinueAsNewOptions, ExternalWorkflowHandle, LocalActivityOptions, + MemoValue, ParentClosePolicy, SignalWorkflowOptions, StartChildWorkflowExecutionFailedCause, + StartChildWorkflowOutput, StartedChildWorkflow, SyncWorkflowContext, TimerOptions, TimerResult, + VersioningIntent, WaitConditionOptions, WorkflowCancellationError, WorkflowCancellationToken, + WorkflowContext, WorkflowContextFuture, WorkflowContextKey, WorkflowContextView, + WorkflowIdReusePolicy, WorkflowRandomValue, WorkflowResult, WorkflowTermination, +}; +#[cfg(feature = "experimental")] +pub use temporalio_workflow::{ + ContinueAsNewVersioningBehavior, NexusOperationCancellationType, NexusOperationOptions, + PatchActivationCallback, PatchActivationInput, StartedNexusOperation, }; #[cfg(feature = "wasm-workflows")] pub use workflow_wasm::WasmWorkflowComponent; @@ -128,14 +136,19 @@ use std::{ time::Duration, }; use temporalio_client::{Client, ClientOptions, NamespacedClient}; +#[cfg(feature = "experimental")] +use temporalio_common::protos::temporal::api::worker::v1::PluginInfo; use temporalio_common::{ ActivityDefinition, WorkflowDefinition, - data_converters::{DataConverter, SerializationContext, SerializationContextData}, + data_converters::{ + ActivitySerializationContext, DataConverter, SerializationContext, + SerializationContextData, WorkflowSerializationContext, + }, payload_visitor::{decode_payloads, encode_payloads}, protos::{ TaskToken, coresdk::{ - ActivityTaskCompletion, AsJsonPayloadExt, + ActivityTaskCompletion, activity_result::ActivityExecutionResult, activity_task::{ActivityTask, activity_task}, workflow_activation::{WorkflowActivation, workflow_activation_job::Variant}, @@ -143,13 +156,14 @@ use temporalio_common::{ }, temporal::api::{ common::v1::Payload, enums::v1::WorkflowTaskFailedCause, failure::v1::Failure, - worker::v1::PluginInfo, }, }, worker::{WorkerDeploymentOptions, WorkerTaskTypes, build_id_from_current_exe}, }; -use temporalio_sdk_core::{PollError, init_worker}; -use temporalio_workflow::runtime::entry::WorkflowImplementation; +use temporalio_sdk_core::{ + PollError, Worker as CoreWorker, WorkerConfig, WorkerVersioningStrategy, init_worker, +}; +use temporalio_workflow::{InternalPatchActivationCallback, workflows::WorkflowImplementation}; use tokio::sync::{ Notify, mpsc::{UnboundedSender, unbounded_channel}, @@ -159,10 +173,7 @@ use tokio_util::sync::CancellationToken; use tracing::{Instrument, Span, field}; use uuid::Uuid; -use crate::runtime::{ - CoreWorker, PollerBehavior, TunerBuilder, WorkerConfig, WorkerTuner, WorkerVersioningStrategy, - WorkflowErrorType, -}; +use crate::runtime::{PollerBehavior, WorkflowErrorType, worker_tuner::WorkerTuner}; /// Contains options for configuring a worker. /// @@ -193,9 +204,11 @@ pub struct WorkerOptions { workflow_interceptor_constructors: Vec, #[builder(field)] + #[cfg(feature = "experimental")] worker_plugins: Vec>, #[builder(field)] + #[cfg(feature = "experimental")] client_plugin_names: HashSet, #[cfg(feature = "wasm-workflows")] @@ -217,10 +230,10 @@ pub struct WorkerOptions { /// or failures. #[builder(default = 1000)] pub max_cached_workflows: usize, - /// Set a [crate::WorkerTuner] for this worker, which controls how many slots are available for - /// the different kinds of tasks. - #[builder(default = Arc::new(TunerBuilder::default().build()))] - pub tuner: Arc, + /// Set a [`runtime::worker_tuner::WorkerTuner`] for this worker, which controls how many slots + /// are available for the different kinds of tasks. + #[builder(into, default)] + pub tuner: WorkerTuner, /// Controls how polling for Workflow tasks will happen on this worker's task queue. See also /// [WorkerConfig::nonsticky_to_sticky_poll_ratio]. If using SimpleMaximum, Must be at least 2 /// when `max_cached_workflows` > 0, or is an error. @@ -296,6 +309,14 @@ pub struct WorkerOptions { /// exceed the namespace error limits; oversized payloads are sent to server, which enforces the /// limit. Defaults to false. /// NOTE: Experimental + #[cfg(feature = "experimental")] + #[cfg_attr( + docsrs, + builder(setters( + some_fn(name = disable_payload_error_limit_impl, vis = "pub(crate)"), + option_fn(name = maybe_disable_payload_error_limit_impl, vis = "pub(crate)") + )) + )] #[builder(default = false)] pub disable_payload_error_limit: bool, /// Experimental callback that decides whether the first non-replay call to @@ -305,9 +326,70 @@ pub struct WorkerOptions { /// `true` records the patch marker; returning `false` leaves the patch inactive for the /// workflow run. For registered WASM workflow components, the callback remains on the worker /// host and is invoked through the workflow component's synchronous host interface. + #[cfg(feature = "experimental")] + #[cfg_attr( + docsrs, + builder(setters( + some_fn(name = patch_activation_callback_impl, vis = "pub(crate)"), + option_fn(name = maybe_patch_activation_callback_impl, vis = "pub(crate)") + )) + )] pub patch_activation_callback: Option, } +// Bon does not propagate `doc(cfg)` to generated setters, so these docs-only methods forward to +// renamed generated implementations. +#[cfg(all(feature = "experimental", docsrs))] +impl WorkerOptionsBuilder { + /// Set whether payloads over the namespace error limit are sent to the server. + #[doc(cfg(feature = "experimental"))] + pub fn disable_payload_error_limit( + self, + value: bool, + ) -> WorkerOptionsBuilder> + where + S::DisablePayloadErrorLimit: worker_options_builder::IsUnset, + { + self.disable_payload_error_limit_impl(value) + } + + /// Set the payload error limit override from an optional value. + #[doc(cfg(feature = "experimental"))] + pub fn maybe_disable_payload_error_limit( + self, + value: Option, + ) -> WorkerOptionsBuilder> + where + S::DisablePayloadErrorLimit: worker_options_builder::IsUnset, + { + self.maybe_disable_payload_error_limit_impl(value) + } + + /// Set the callback used to decide whether a patch should activate. + #[doc(cfg(feature = "experimental"))] + pub fn patch_activation_callback( + self, + value: PatchActivationCallback, + ) -> WorkerOptionsBuilder> + where + S::PatchActivationCallback: worker_options_builder::IsUnset, + { + self.patch_activation_callback_impl(value) + } + + /// Set the patch activation callback from an optional value. + #[doc(cfg(feature = "experimental"))] + pub fn maybe_patch_activation_callback( + self, + value: Option, + ) -> WorkerOptionsBuilder> + where + S::PatchActivationCallback: worker_options_builder::IsUnset, + { + self.maybe_patch_activation_callback_impl(value) + } +} + impl WorkerOptionsBuilder { pub(crate) fn with_workflows(mut self, workflows: WorkflowDefinitions) -> Self { self.workflows = workflows; @@ -330,6 +412,7 @@ impl WorkerOptionsBuilder { self } + #[cfg(feature = "experimental")] pub(crate) fn with_worker_plugins( mut self, worker_plugins: Vec>, @@ -350,12 +433,14 @@ impl WorkerOptionsBuilder { /// Register a worker plugin. /// /// **Experimental:** This API may change or be removed. + #[cfg(feature = "experimental")] pub fn worker_plugin(mut self, plugin: P) -> Self { self.worker_plugins.push(Arc::new(plugin)); self } /// Append a worker interceptor. Interceptors run in registration order. + #[cfg(feature = "experimental")] pub fn worker_interceptor(mut self, interceptor: I) -> Self { self.worker_interceptors.push(Arc::new(interceptor)); self @@ -468,6 +553,7 @@ fn def_build_id() -> WorkerDeploymentOptions { impl WorkerOptions { /// Append a worker interceptor. Interceptors run in registration order. + #[cfg(feature = "experimental")] pub fn worker_interceptor( &mut self, interceptor: I, @@ -581,6 +667,27 @@ impl WorkerOptions { if !workflows_registered && !activities_registered { return Err("At least one workflow or activity must be registered".to_owned()); } + #[cfg(feature = "experimental")] + let disable_payload_error_limit = self.disable_payload_error_limit; + #[cfg(not(feature = "experimental"))] + let disable_payload_error_limit = false; + #[cfg(feature = "experimental")] + let plugin_info = self + .client_plugin_names + .iter() + .map(|name| PluginInfo { + name: name.clone(), + version: String::new(), + }) + .chain(self.worker_plugins.iter().map(|registration| PluginInfo { + name: registration.name().to_owned(), + version: String::new(), + })) + .collect(); + #[cfg(not(feature = "experimental"))] + let plugin_info = HashSet::new(); + + let tuner = self.tuner.to_core()?; WorkerConfig::builder() .namespace(namespace) @@ -595,10 +702,19 @@ impl WorkerOptions { }) })) .max_cached_workflows(self.max_cached_workflows) - .tuner(self.tuner.clone()) - .maybe_workflow_task_poller_behavior(self.workflow_task_poller_behavior) - .maybe_activity_task_poller_behavior(self.activity_task_poller_behavior) - .maybe_nexus_task_poller_behavior(self.nexus_task_poller_behavior) + .tuner(tuner) + .maybe_workflow_task_poller_behavior( + self.workflow_task_poller_behavior + .map(PollerBehavior::into_core), + ) + .maybe_activity_task_poller_behavior( + self.activity_task_poller_behavior + .map(PollerBehavior::into_core), + ) + .maybe_nexus_task_poller_behavior( + self.nexus_task_poller_behavior + .map(PollerBehavior::into_core), + ) .task_types(WorkerTaskTypes { enable_workflows: workflows_registered, enable_local_activities: workflows_registered && activities_registered, @@ -617,22 +733,30 @@ impl WorkerOptions { .versioning_strategy(WorkerVersioningStrategy::WorkerDeploymentBased( self.deployment_options.clone(), )) - .workflow_failure_errors(self.workflow_failure_errors.clone()) - .workflow_types_to_failure_errors(self.workflow_types_to_failure_errors.clone()) - .plugins( - self.client_plugin_names + .workflow_failure_errors( + self.workflow_failure_errors .iter() - .map(|name| PluginInfo { - name: name.clone(), - version: String::new(), + .cloned() + .map(WorkflowErrorType::into_core) + .collect(), + ) + .workflow_types_to_failure_errors( + self.workflow_types_to_failure_errors + .iter() + .map(|(workflow_type, error_types)| { + ( + workflow_type.clone(), + error_types + .iter() + .cloned() + .map(WorkflowErrorType::into_core) + .collect(), + ) }) - .chain(self.worker_plugins.iter().map(|registration| PluginInfo { - name: registration.name().to_owned(), - version: String::new(), - })) .collect(), ) - .disable_payload_error_limit(self.disable_payload_error_limit) + .plugins(plugin_info) + .disable_payload_error_limit(disable_payload_error_limit) .build() } } @@ -670,7 +794,7 @@ struct WorkflowHalf { workflow_removed_from_map: Notify, detect_nondeterministic_futures: bool, #[debug(skip)] - patch_activation_callback: Option, + patch_activation_callback: Option, } #[derive(Debug)] struct WorkflowData { @@ -736,7 +860,7 @@ async fn encode_workflow_completion( if let Err(err) = encode_payloads( completion, data_converter.codec(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .await { @@ -759,7 +883,7 @@ async fn encode_activity_completion( if let Err(err) = encode_payloads( completion, data_converter.codec(), - &SerializationContextData::Activity, + &SerializationContextData::Activity(ActivitySerializationContext::new()), ) .await { @@ -776,18 +900,22 @@ impl Worker { pub fn new( runtime: &Runtime, client: Client, - mut options: WorkerOptions, + options: WorkerOptions, ) -> Result { + #[cfg(feature = "experimental")] + let mut options = options; + #[cfg(feature = "experimental")] plugins::apply_worker_plugins(client.options(), &mut options)?; let wc = options .to_core_options(client.namespace(), client.identity()) .map_err(|error| WorkerCreateError::Initialization(anyhow!(error)))?; - let core = init_worker(runtime, wc, client.connection().clone()) + let core = init_worker(runtime.core(), wc, client.connection().clone()) .map_err(WorkerCreateError::Initialization)?; Self::new_from_core_options_prepared(Arc::new(core), client.options().clone(), options) } // TODO [rust-sdk-branch]: Eliminate this constructor in favor of passing in fake connection + #[cfg(feature = "experimental")] #[doc(hidden)] pub fn new_from_core(worker: Arc, data_converter: DataConverter) -> Self { let client_options = ClientOptions::new(worker.get_config().namespace.clone()) @@ -805,12 +933,16 @@ impl Worker { } // TODO [rust-sdk-branch]: Eliminate this constructor in favor of passing in fake connection + #[cfg(feature = "experimental")] #[doc(hidden)] pub fn new_from_core_options( worker: Arc, client_options: ClientOptions, - mut options: WorkerOptions, + options: WorkerOptions, ) -> Result { + #[cfg(feature = "experimental")] + let mut options = options; + #[cfg(feature = "experimental")] plugins::apply_worker_plugins(&client_options, &mut options)?; Self::new_from_core_options_prepared(worker, client_options, options) } @@ -838,8 +970,11 @@ impl Worker { activity_inbound_interceptors, workflow_interceptor_constructors, ); - me.set_detect_nondeterministic_futures(options.detect_nondeterministic_futures); - me.workflow_half.patch_activation_callback = options.patch_activation_callback; + me.workflow_half.detect_nondeterministic_futures = options.detect_nondeterministic_futures; + #[cfg(feature = "experimental")] + { + me.workflow_half.patch_activation_callback = options.patch_activation_callback; + } #[cfg(feature = "wasm-workflows")] me.workflow_half .workflow_definitions @@ -890,13 +1025,6 @@ impl Worker { &self.common.task_queue } - #[doc(hidden)] - /// Set whether nondeterministic future detection is enabled for workflows on this worker. Users - /// should use [WorkerOptions] to set this. TODO: Only needs to exist due to test setup. - pub fn set_detect_nondeterministic_futures(&mut self, enabled: bool) { - self.workflow_half.detect_nondeterministic_futures = enabled; - } - /// Return a handle that can be used to initiate shutdown. This is useful because [Worker::run] /// takes self mutably, so you may want to obtain a handle for shutting down before running. pub fn shutdown_handle(&self) -> impl Fn() + use<> { @@ -927,6 +1055,7 @@ impl Worker { .worker .validate() .await + .map_err(WorkerValidationError::from_core) .map_err(WorkerRunError::Validation)?; let shutdown_token = CancellationToken::new(); let (common, wf_half, act_half) = self.split_apart(); @@ -1012,7 +1141,7 @@ impl Worker { if let Err(err) = decode_payloads( &mut activation, common.data_converter.codec(), - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) .await { @@ -1095,12 +1224,15 @@ impl Worker { message: "activity polling failed".to_owned(), source: Box::new(source), })?; - if let Err(err) = decode_payloads( - &mut activity, - common.data_converter.codec(), - &SerializationContextData::Activity, - ) - .await + if let Err(err) = + decode_payloads( + &mut activity, + common.data_converter.codec(), + &SerializationContextData::Activity( + ActivitySerializationContext::new(), + ), + ) + .await { error!(error = %err, "Failed decoding activity task"); let mut completion = ActivityTaskCompletion { @@ -1138,7 +1270,9 @@ impl Worker { task_token, }) => { let failure = common.data_converter.to_failure( - &SerializationContextData::Activity, + &SerializationContextData::Activity( + ActivitySerializationContext::new(), + ), OutgoingError::Activity(OutgoingActivityError::Application( ApplicationFailure::builder(source) .type_name("NotFoundError".to_owned()) @@ -1204,11 +1338,6 @@ impl Worker { self.common.worker.worker_instance_key() } - #[doc(hidden)] - pub fn core_worker(&self) -> Arc { - self.common.worker.clone() - } - fn split_apart(&mut self) -> (&mut CommonWorker, &mut WorkflowHalf, &mut ActivityHalf) { ( &mut self.common, @@ -1395,10 +1524,12 @@ impl ActivityHalf { tokio::spawn(async move { let act_fut = async move { - if let Some(info) = &ctx.info().workflow_execution { - Span::current() - .record("temporalWorkflowID", info.workflow_id()) - .record("temporalRunID", info.run_id()); + let span = Span::current(); + if let Some(workflow_id) = &ctx.info().workflow_id { + span.record("temporalWorkflowID", workflow_id); + } + if let Some(workflow_run_id) = &ctx.info().workflow_run_id { + span.record("temporalRunID", workflow_run_id); } (act_fn)(args, data_converter, ctx, activity_inbound_interceptors).await } @@ -1409,10 +1540,10 @@ impl ActivityHalf { // Codec application happens at the SDK/Core boundary, so activity // implementations work with the payload converter directly. let pc = codec_data_converter.payload_converter(); - let ctx = SerializationContext { - data: &SerializationContextData::Activity, - converter: pc, - }; + let context_data = SerializationContextData::Activity( + ActivitySerializationContext::new(), + ); + let ctx = SerializationContext::new(&context_data, pc); match output.serialize_payload(&ctx) { Ok(payload) => ActivityExecutionResult::ok(payload), Err(err) => { @@ -1447,21 +1578,6 @@ impl ActivityHalf { } } -/// Activity functions may return these values when exiting -#[derive(Debug)] -pub enum ActExitValue { - /// Completion requires an asynchronous callback - WillCompleteAsync, - /// Finish with a result - Normal(T), -} - -impl From for ActExitValue { - fn from(t: T) -> Self { - Self::Normal(t) - } -} - /// Attempts to turn caught panics into something printable fn panic_formatter(panic: Box) -> Box { _panic_formatter::<&str>(panic) @@ -1594,7 +1710,7 @@ mod tests { let codec = Arc::new(FailingEncodeCodec::default()); let data_converter = DataConverter::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), codec.clone(), ); let mut completion = WorkflowActivationCompletion::from_cmd( @@ -1625,7 +1741,7 @@ mod tests { let codec = Arc::new(FailingEncodeCodec::default()); let data_converter = DataConverter::new( PayloadConverter::default(), - DefaultFailureConverter, + DefaultFailureConverter::default(), codec.clone(), ); let mut completion = ActivityTaskCompletion { @@ -1749,30 +1865,6 @@ mod tests { .unwrap(); } - #[test] - fn simple_plugin_workflow_function_merges_definitions() { - let plugin = SimplePlugin::builder("simple") - .workflows(|existing: Option| { - assert!(existing.is_some()); - let mut workflows = WorkflowDefinitions::new(); - workflows.register_workflow::().unwrap(); - workflows - }) - .build(); - let client_options = ClientOptions::new("namespace").build(); - let mut worker_options = WorkerOptions::new("task_q") - .register_workflow::() - .unwrap() - .worker_plugin(plugin) - .build(); - - crate::plugins::apply_worker_plugins(&client_options, &mut worker_options).unwrap(); - - let workflows = format!("{:?}", worker_options.workflows()); - assert!(workflows.contains("MyWorkflow")); - assert!(workflows.contains("OtherWorkflow")); - } - #[rstest::rstest] #[case::workflow_only(true, false, Ok(WorkerTaskTypes::workflow_only()))] #[case::activity_only(false, true, Ok(WorkerTaskTypes::activity_only()))] @@ -1908,24 +2000,6 @@ mod tests { assert_eq!(config.client_identity_override, expected); } - #[rstest::rstest] - #[case::default_enforces_error_limit(None, false)] - #[case::opt_out_disables_error_limit(Some(true), true)] - #[case::explicit_enable_error_limit(Some(false), false)] - #[test] - fn disable_payload_error_limit_propagates( - #[case] override_value: Option, - #[case] expected: bool, - ) { - let config = WorkerOptions::new("task_q") - .register_activities(MyActivities {}) - .maybe_disable_payload_error_limit(override_value) - .build() - .to_core_options("ns".into(), String::new()) - .unwrap(); - assert_eq!(config.disable_payload_error_limit, expected); - } - #[test] fn max_eager_activity_reservations_per_workflow_task_propagates() { let config = WorkerOptions::new("task_q") @@ -1936,4 +2010,51 @@ mod tests { .unwrap(); assert_eq!(config.max_eager_activity_reservations_per_workflow_task, 7); } + + #[cfg(feature = "experimental")] + mod experimental_tests { + use super::*; + + #[test] + fn simple_plugin_workflow_function_merges_definitions() { + let plugin = SimplePlugin::builder("simple") + .workflows(|existing: Option| { + assert!(existing.is_some()); + let mut workflows = WorkflowDefinitions::new(); + workflows.register_workflow::().unwrap(); + workflows + }) + .build(); + let client_options = ClientOptions::new("namespace").build(); + let mut worker_options = WorkerOptions::new("task_q") + .register_workflow::() + .unwrap() + .worker_plugin(plugin) + .build(); + + crate::plugins::apply_worker_plugins(&client_options, &mut worker_options).unwrap(); + + let workflows = format!("{:?}", worker_options.workflows()); + assert!(workflows.contains("MyWorkflow")); + assert!(workflows.contains("OtherWorkflow")); + } + + #[rstest::rstest] + #[case::default_enforces_error_limit(None, false)] + #[case::opt_out_disables_error_limit(Some(true), true)] + #[case::explicit_enable_error_limit(Some(false), false)] + #[test] + fn disable_payload_error_limit_propagates( + #[case] override_value: Option, + #[case] expected: bool, + ) { + let config = WorkerOptions::new("task_q") + .register_activities(MyActivities {}) + .maybe_disable_payload_error_limit(override_value) + .build() + .to_core_options("ns".into(), String::new()) + .unwrap(); + assert_eq!(config.disable_payload_error_limit, expected); + } + } } diff --git a/crates/sdk/src/plugins.rs b/crates/sdk/src/plugins.rs index 7aa851a91..510973643 100644 --- a/crates/sdk/src/plugins.rs +++ b/crates/sdk/src/plugins.rs @@ -454,7 +454,7 @@ pub(crate) fn apply_workflow_replayer_plugins( Ok(()) } -#[cfg(test)] +#[cfg(all(test, feature = "experimental"))] mod tests { use super::*; use std::{ diff --git a/crates/sdk/src/runtime.rs b/crates/sdk/src/runtime.rs index 840be9e28..ac16330b7 100644 --- a/crates/sdk/src/runtime.rs +++ b/crates/sdk/src/runtime.rs @@ -4,25 +4,106 @@ //! primary workflow and activity APIs. Create a [`crate::Runtime`] before connecting a client, //! then pass it to [`crate::Worker::new`]. -use std::{ - ops::{Deref, DerefMut}, - time::Duration, -}; +use std::time::Duration; -use temporalio_common::telemetry::{TelemetryInstance, TelemetryOptions}; -use temporalio_sdk_core::{CoreRuntime, RuntimeOptions as CoreRuntimeOptions}; - -pub use temporalio_sdk_core::{ - ActivitySlotKind, FixedSizeSlotSupplier, LocalActivitySlotKind, NexusSlotKind, PollerBehavior, - ResourceBasedSlotsOptions, ResourceBasedSlotsOptionsBuilder, ResourceBasedTuner, - ResourceBasedTunerConfig, ResourceController, ResourceSlotOptions, SlotInfo, SlotInfoTrait, - SlotKind, SlotKindType, SlotMarkUsedContext, SlotReleaseContext, SlotReservationContext, - SlotSupplier, SlotSupplierOptions, SlotSupplierPermit, TokioRuntimeBuilder, TunerBuilder, - TunerHolder, TunerHolderOptions, TunerHolderOptionsBuilder, Worker as CoreWorker, WorkerConfig, - WorkerConfigBuilder, WorkerTuner, WorkerVersioningStrategy, WorkflowErrorType, - WorkflowSlotKind, init_replay_worker, replay, +use temporalio_common::telemetry::TelemetryOptions; +use temporalio_sdk_core::{ + CoreRuntime, PollerBehavior as CorePollerBehavior, RuntimeOptions as CoreRuntimeOptions, + TokioRuntimeBuilder as CoreTokioRuntimeBuilder, WorkflowErrorType as CoreWorkflowErrorType, }; +use crate::error::RuntimeError; + +/// Worker concurrency tuning. +pub mod worker_tuner; + +// Keep these public only with the raw-worker APIs while they are migrated separately. +// Worker::new_from_core, Worker::new_from_core_options, Worker::with_new_core_worker +#[cfg(feature = "experimental")] +pub use temporalio_sdk_core::{Worker as CoreWorker, WorkerConfig}; + +/// Wraps a Tokio runtime builder so the SDK can install its per-thread telemetry state. +#[derive(bon::Builder)] +#[builder(state_mod(vis = "pub"))] +#[non_exhaustive] +pub struct TokioRuntimeBuilder { + /// The Tokio runtime builder used to create the runtime. + pub inner: tokio::runtime::Builder, +} + +impl Default for TokioRuntimeBuilder { + fn default() -> Self { + Self { + inner: tokio::runtime::Builder::new_multi_thread(), + } + } +} + +impl TokioRuntimeBuilder { + fn into_core(self) -> CoreTokioRuntimeBuilder> { + CoreTokioRuntimeBuilder { + inner: self.inner, + lang_on_thread_start: None, + } + } +} + +/// Options for automatically scaling the number of concurrent task polls. +#[derive(bon::Builder, Clone, Copy, Debug, PartialEq)] +#[builder(state_mod(vis = "pub"))] +#[non_exhaustive] +pub struct AutoscalingOptions { + /// Minimum number of concurrent polls. Cannot be zero. + pub minimum: usize, + /// Maximum number of concurrent polls. Must be at least `minimum`. + pub maximum: usize, + /// Initial number of concurrent polls. Must be between `minimum` and `maximum`. + pub initial: usize, +} + +/// Controls how many concurrent task polls a worker issues. +#[derive(Clone, Copy, Debug, PartialEq)] +#[non_exhaustive] +pub enum PollerBehavior { + /// Poll whenever a slot is available, up to the supplied maximum. + SimpleMaximum(usize), + /// Adjust concurrent polls using feedback from the server. + Autoscaling(AutoscalingOptions), +} + +impl PollerBehavior { + pub(crate) fn into_core(self) -> CorePollerBehavior { + match self { + PollerBehavior::SimpleMaximum(maximum) => CorePollerBehavior::SimpleMaximum(maximum), + PollerBehavior::Autoscaling(AutoscalingOptions { + minimum, + maximum, + initial, + }) => CorePollerBehavior::Autoscaling { + minimum, + maximum, + initial, + }, + } + } +} + +/// Workflow-processing errors that may be configured to fail the workflow execution. +#[derive(Clone, Debug, Eq, PartialEq, Hash)] +#[non_exhaustive] +pub enum WorkflowErrorType { + /// A workflow produced commands that do not match its recorded history. + Nondeterminism, +} + +impl WorkflowErrorType { + pub(crate) fn into_core(self) -> CoreWorkflowErrorType { + match self { + WorkflowErrorType::Nondeterminism => CoreWorkflowErrorType::Nondeterminism, + } + } +} + /// Configuration for the Rust SDK runtime. Construct with [`RuntimeOptions::builder`]. #[derive(bon::Builder)] #[builder(finish_fn(vis = "", name = build_internal))] @@ -66,12 +147,12 @@ impl RuntimeOptionsBuilder { } } -impl From for CoreRuntimeOptions { - fn from(options: RuntimeOptions) -> Self { +impl RuntimeOptions { + fn into_core(self) -> CoreRuntimeOptions { CoreRuntimeOptions::builder() - .telemetry_options(options.telemetry_options) - .heartbeat_interval(options.heartbeat_interval) - .disable_environment_info(options.disable_environment_info) + .telemetry_options(self.telemetry_options) + .heartbeat_interval(self.heartbeat_interval) + .disable_environment_info(self.disable_environment_info) .build() .expect("SDK runtime options have already been validated") } @@ -82,50 +163,62 @@ pub struct Runtime(CoreRuntime); impl Runtime { /// Creates a runtime with a newly constructed Tokio runtime. - pub fn new( + /// + /// # Errors + /// Returns an error if telemetry or the Tokio runtime cannot be initialized. + pub fn new( options: RuntimeOptions, - tokio_builder: TokioRuntimeBuilder, - ) -> Result - where - F: Fn() + Send + Sync + 'static, - { - CoreRuntime::new(options.into(), tokio_builder).map(Self) + tokio_builder: TokioRuntimeBuilder, + ) -> Result { + CoreRuntime::new(options.into_core(), tokio_builder.into_core()) + .map(Self) + .map_err(RuntimeError::from_core) } /// Creates a runtime using the currently active Tokio runtime. /// - /// # Panics - /// Panics if there is no currently active Tokio runtime. - pub fn new_assume_tokio(options: RuntimeOptions) -> Result { - CoreRuntime::new_assume_tokio(options.into()).map(Self) + /// # Errors + /// Returns [`RuntimeError::NoCurrentTokioRuntime`] if there is no currently active Tokio + /// runtime, or [`RuntimeError::Initialization`] if telemetry cannot be initialized. + pub fn from_current_tokio(options: RuntimeOptions) -> Result { + tokio::runtime::Handle::try_current().map_err(|_| RuntimeError::NoCurrentTokioRuntime)?; + CoreRuntime::new_assume_tokio(options.into_core()) + .map(Self) + .map_err(RuntimeError::from_core) } - /// Creates a runtime from an initialized telemetry instance using the currently active Tokio - /// runtime. + /// Creates a runtime using the currently active Tokio runtime. /// - /// # Panics - /// Panics if there is no currently active Tokio runtime. - pub fn new_assume_tokio_initialized_telem( - telemetry: TelemetryInstance, - heartbeat_interval: Option, - ) -> Self { - Self(CoreRuntime::new_assume_tokio_initialized_telem( - telemetry, - heartbeat_interval, - )) + /// # Errors + /// Returns [`RuntimeError::NoCurrentTokioRuntime`] if there is no currently active Tokio + /// runtime, or [`RuntimeError::Initialization`] if telemetry cannot be initialized. + #[deprecated(note = "use `Runtime::from_current_tokio` instead")] + pub fn new_assume_tokio(options: RuntimeOptions) -> Result { + Self::from_current_tokio(options) } -} - -impl Deref for Runtime { - type Target = CoreRuntime; - fn deref(&self) -> &Self::Target { + pub(crate) fn core(&self) -> &CoreRuntime { &self.0 } } -impl DerefMut for Runtime { - fn deref_mut(&mut self) -> &mut Self::Target { - &mut self.0 +#[cfg(test)] +mod tests { + use super::{Runtime, TokioRuntimeBuilder}; + use crate::error::RuntimeError; + + #[test] + fn from_current_tokio_without_runtime_returns_error() { + assert!(matches!( + Runtime::from_current_tokio(Default::default()), + Err(RuntimeError::NoCurrentTokioRuntime) + )); + } + + #[test] + fn tokio_runtime_builder_constructs_with_an_inner_builder() { + let _builder = TokioRuntimeBuilder::builder() + .inner(tokio::runtime::Builder::new_current_thread()) + .build(); } } diff --git a/crates/sdk/src/runtime/worker_tuner.rs b/crates/sdk/src/runtime/worker_tuner.rs new file mode 100644 index 000000000..8eafe35e1 --- /dev/null +++ b/crates/sdk/src/runtime/worker_tuner.rs @@ -0,0 +1,461 @@ +//! Worker concurrency tuning for SDK workers. + +use std::{fmt::Debug, sync::Arc, time::Duration}; + +use temporalio_sdk_core::{ + ResourceBasedSlotsOptions as CoreResourceBasedSlotsOptions, + ResourceBasedTunerConfig as CoreResourceBasedTunerConfig, + ResourceController as CoreResourceController, ResourceSlotOptions as CoreResourceSlotOptions, + SlotKind as CoreSlotKind, SlotSupplierOptions as CoreSlotSupplierOptions, + TunerHolderOptions as CoreTunerHolderOptions, WorkerTuner as CoreWorkerTuner, +}; + +const DEFAULT_FIXED_SIZE_SLOTS: usize = 100; + +/// A worker tuner configuration. +#[derive(Clone, Debug)] +#[non_exhaustive] +pub enum WorkerTuner { + /// A tuner that creates a resource controller scoped to this worker. + ResourceBased(ResourceBasedTuner), + /// A resource-based tuner using a controller that may be shared by multiple workers. + ResourceBasedWithController(ResourceBasedTunerWithController), + /// A tuner composed from independently selected slot suppliers. + TunerHolder(TunerHolder), +} + +impl WorkerTuner { + pub(crate) fn to_core(&self) -> Result, String> { + match self { + Self::ResourceBased(tuner) => tuner.to_tuner_holder().to_core(), + Self::ResourceBasedWithController(tuner) => tuner.to_tuner_holder().to_core(), + Self::TunerHolder(tuner) => tuner.to_core(), + } + } +} + +impl Default for WorkerTuner { + fn default() -> Self { + TunerHolder::builder() + .workflow_task_slot_supplier(FixedSizeSlotSupplier::new(DEFAULT_FIXED_SIZE_SLOTS)) + .activity_task_slot_supplier(FixedSizeSlotSupplier::new(DEFAULT_FIXED_SIZE_SLOTS)) + .local_activity_task_slot_supplier(FixedSizeSlotSupplier::new(DEFAULT_FIXED_SIZE_SLOTS)) + .nexus_task_slot_supplier(FixedSizeSlotSupplier::new(DEFAULT_FIXED_SIZE_SLOTS)) + .build() + .into() + } +} + +impl From for WorkerTuner { + fn from(value: ResourceBasedTuner) -> Self { + Self::ResourceBased(value) + } +} + +impl From for WorkerTuner { + fn from(value: ResourceBasedTunerWithController) -> Self { + Self::ResourceBasedWithController(value) + } +} + +impl From for WorkerTuner { + fn from(value: TunerHolder) -> Self { + Self::TunerHolder(value) + } +} + +/// A tuner composed from independently selected slot suppliers. +#[derive(Clone, Debug, bon::Builder)] +#[builder(state_mod(vis = "pub"))] +#[non_exhaustive] +pub struct TunerHolder { + /// Supplies workflow-task slots. + #[builder(into)] + pub workflow_task_slot_supplier: SlotSupplier, + /// Supplies activity-task slots. + #[builder(into)] + pub activity_task_slot_supplier: SlotSupplier, + /// Supplies local-activity slots. + #[builder(into)] + pub local_activity_task_slot_supplier: SlotSupplier, + /// Supplies Nexus-task slots. + #[builder(into)] + pub nexus_task_slot_supplier: SlotSupplier, +} + +impl TunerHolder { + fn to_core(&self) -> Result, String> { + let configurations = [ + self.workflow_task_slot_supplier.resource_configuration(), + self.activity_task_slot_supplier.resource_configuration(), + self.local_activity_task_slot_supplier + .resource_configuration(), + self.nexus_task_slot_supplier.resource_configuration(), + ]; + let mut configurations = configurations.into_iter().flatten(); + let resource_based_config = configurations.next(); + if let Some(first) = resource_based_config + && configurations.any(|other| !first.is_compatible_with(other)) + { + return Err( + "cannot construct worker tuner with multiple different resource-based tuner configurations" + .to_owned(), + ); + } + + CoreTunerHolderOptions::builder() + .workflow_slot_options( + self.workflow_task_slot_supplier + .to_core(ResourceKind::Workflow), + ) + .activity_slot_options( + self.activity_task_slot_supplier + .to_core(ResourceKind::Activity), + ) + .local_activity_slot_options( + self.local_activity_task_slot_supplier + .to_core(ResourceKind::Activity), + ) + .nexus_slot_options(self.nexus_task_slot_supplier.to_core(ResourceKind::Nexus)) + .maybe_resource_based_config(resource_based_config.map(ResourceBasedConfig::to_core)) + .build() + .map_err(|error| error.to_string())? + .build_tuner_holder() + .map(|tuner| Arc::new(tuner) as Arc) + .map_err(|error| error.to_string()) + } +} + +/// A resource-based tuner that creates a controller scoped to its worker. +#[derive(Clone, Debug, bon::Builder)] +#[builder(state_mod(vis = "pub"))] +#[non_exhaustive] +pub struct ResourceBasedTuner { + /// Target memory and CPU usage. + pub tuner_options: ResourceBasedTunerOptions, + /// Workflow-task slot options, or `None` to use defaults. + pub workflow_task_slot_options: Option, + /// Activity-task slot options, or `None` to use defaults. + pub activity_task_slot_options: Option, + /// Local-activity slot options, or `None` to use defaults. + pub local_activity_task_slot_options: Option, + /// Nexus-task slot options, or `None` to use defaults. + pub nexus_task_slot_options: Option, +} + +impl ResourceBasedTuner { + fn to_tuner_holder(&self) -> TunerHolder { + TunerHolder::builder() + .workflow_task_slot_supplier(ResourceBasedSlotsForType::new( + self.tuner_options, + self.workflow_task_slot_options.unwrap_or_default(), + )) + .activity_task_slot_supplier(ResourceBasedSlotsForType::new( + self.tuner_options, + self.activity_task_slot_options.unwrap_or_default(), + )) + .local_activity_task_slot_supplier(ResourceBasedSlotsForType::new( + self.tuner_options, + self.local_activity_task_slot_options.unwrap_or_default(), + )) + .nexus_task_slot_supplier(ResourceBasedSlotsForType::new( + self.tuner_options, + self.nexus_task_slot_options.unwrap_or_default(), + )) + .build() + } +} + +/// A resource-based tuner governed by a controller shared across workers. +#[derive(Clone, Debug, bon::Builder)] +#[builder(state_mod(vis = "pub"))] +#[non_exhaustive] +pub struct ResourceBasedTunerWithController { + /// The shared resource controller. + pub controller: ResourceBasedController, + /// Workflow-task slot options, or `None` to use defaults. + pub workflow_task_slot_options: Option, + /// Activity-task slot options, or `None` to use defaults. + pub activity_task_slot_options: Option, + /// Local-activity slot options, or `None` to use defaults. + pub local_activity_task_slot_options: Option, + /// Nexus-task slot options, or `None` to use defaults. + pub nexus_task_slot_options: Option, +} + +impl ResourceBasedTunerWithController { + fn to_tuner_holder(&self) -> TunerHolder { + TunerHolder::builder() + .workflow_task_slot_supplier(ResourceBasedSlotsForType::with_controller( + self.controller.clone(), + self.workflow_task_slot_options.unwrap_or_default(), + )) + .activity_task_slot_supplier(ResourceBasedSlotsForType::with_controller( + self.controller.clone(), + self.activity_task_slot_options.unwrap_or_default(), + )) + .local_activity_task_slot_supplier(ResourceBasedSlotsForType::with_controller( + self.controller.clone(), + self.local_activity_task_slot_options.unwrap_or_default(), + )) + .nexus_task_slot_supplier(ResourceBasedSlotsForType::with_controller( + self.controller.clone(), + self.nexus_task_slot_options.unwrap_or_default(), + )) + .build() + } +} + +/// Target resource usage for resource-based tuning. +#[derive(Clone, Copy, Debug, PartialEq, bon::Builder)] +#[builder(state_mod(vis = "pub"))] +#[non_exhaustive] +pub struct ResourceBasedTunerOptions { + /// Target system memory usage as a fraction from zero to one. + pub target_memory_usage: f64, + /// Target system CPU usage as a fraction from zero to one. + pub target_cpu_usage: f64, +} + +impl ResourceBasedTunerOptions { + fn to_core(self) -> CoreResourceBasedSlotsOptions { + CoreResourceBasedSlotsOptions::builder() + .target_mem_usage(self.target_memory_usage) + .target_cpu_usage(self.target_cpu_usage) + .build() + } +} + +/// Coordinates resource-based slot allocation across multiple workers. +#[derive(Clone, derive_more::Debug)] +pub struct ResourceBasedController(#[debug(skip)] Arc); + +impl ResourceBasedController { + /// Creates a controller using system-wide memory and CPU measurements. + pub fn new(options: ResourceBasedTunerOptions) -> Self { + Self(Arc::new(CoreResourceController::new(options.to_core()))) + } +} + +/// Per-task-type options for resource-based slots. +#[derive(Clone, Copy, Debug, Default, Eq, PartialEq, bon::Builder)] +#[builder(state_mod(vis = "pub"))] +#[non_exhaustive] +pub struct ResourceBasedSlotOptions { + /// Slots issued without consulting resource usage, or `None` to use the task-type default. + pub minimum_slots: Option, + /// Maximum slots that may be issued, or `None` to use the task-type default. + pub maximum_slots: Option, + /// Minimum delay between slots above the minimum, or `None` to use the task-type default. + pub ramp_throttle: Option, +} + +impl ResourceBasedSlotOptions { + fn to_core(self, kind: ResourceKind) -> CoreResourceSlotOptions { + let defaults = kind.slot_defaults(); + CoreResourceSlotOptions::new( + self.minimum_slots.unwrap_or(defaults.minimum_slots), + self.maximum_slots.unwrap_or(defaults.maximum_slots), + self.ramp_throttle.unwrap_or(defaults.ramp_throttle), + ) + } +} + +/// Resource-based slot settings for one task type. +#[derive(Clone, Debug)] +pub struct ResourceBasedSlotsForType { + configuration: ResourceBasedConfig, + /// Per-task-type slot settings. + pub slot_options: ResourceBasedSlotOptions, +} + +impl ResourceBasedSlotsForType { + /// Creates settings that construct a resource controller with the worker. + pub fn new( + tuner_options: ResourceBasedTunerOptions, + slot_options: ResourceBasedSlotOptions, + ) -> Self { + Self { + configuration: ResourceBasedConfig::Options(tuner_options), + slot_options, + } + } + + /// Creates settings governed by a shared resource controller. + pub fn with_controller( + controller: ResourceBasedController, + slot_options: ResourceBasedSlotOptions, + ) -> Self { + Self { + configuration: ResourceBasedConfig::Controller(controller), + slot_options, + } + } + + /// Returns target options when the controller is scoped to the worker. + pub fn tuner_options(&self) -> Option { + match self.configuration { + ResourceBasedConfig::Options(options) => Some(options), + ResourceBasedConfig::Controller(_) => None, + } + } + + /// Returns the shared controller, when configured. + pub fn controller(&self) -> Option<&ResourceBasedController> { + match &self.configuration { + ResourceBasedConfig::Options(_) => None, + ResourceBasedConfig::Controller(controller) => Some(controller), + } + } +} + +#[derive(Clone, Debug)] +enum ResourceBasedConfig { + Options(ResourceBasedTunerOptions), + Controller(ResourceBasedController), +} + +impl ResourceBasedConfig { + fn is_compatible_with(&self, other: &Self) -> bool { + match (self, other) { + (Self::Options(left), Self::Options(right)) => left == right, + (Self::Controller(left), Self::Controller(right)) => Arc::ptr_eq(&left.0, &right.0), + (Self::Options(_), Self::Controller(_)) | (Self::Controller(_), Self::Options(_)) => { + false + } + } + } + + fn to_core(&self) -> CoreResourceBasedTunerConfig { + match self { + Self::Options(options) => CoreResourceBasedTunerConfig::Options(options.to_core()), + Self::Controller(controller) => { + CoreResourceBasedTunerConfig::Controller(controller.0.clone()) + } + } + } +} + +/// A fixed-size slot supplier. +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +#[non_exhaustive] +pub struct FixedSizeSlotSupplier { + /// Maximum number of slots that may be issued. + pub num_slots: usize, +} + +impl FixedSizeSlotSupplier { + /// Creates a fixed-size supplier. + pub fn new(num_slots: usize) -> Self { + Self { num_slots } + } +} + +/// A fixed-size or resource-based slot supplier. +#[derive(Clone, Debug)] +#[non_exhaustive] +pub enum SlotSupplier { + /// A supplier with a fixed concurrency limit. + FixedSize(FixedSizeSlotSupplier), + /// A supplier governed by resource usage. + ResourceBased(ResourceBasedSlotsForType), +} + +impl From for SlotSupplier { + fn from(value: FixedSizeSlotSupplier) -> Self { + Self::FixedSize(value) + } +} + +impl From for SlotSupplier { + fn from(value: ResourceBasedSlotsForType) -> Self { + Self::ResourceBased(value) + } +} + +impl SlotSupplier { + fn resource_configuration(&self) -> Option<&ResourceBasedConfig> { + match self { + Self::ResourceBased(options) => Some(&options.configuration), + Self::FixedSize(_) => None, + } + } + + fn to_core(&self, kind: ResourceKind) -> CoreSlotSupplierOptions { + match self { + Self::FixedSize(supplier) => CoreSlotSupplierOptions::FixedSize { + slots: supplier.num_slots, + }, + Self::ResourceBased(options) => { + CoreSlotSupplierOptions::ResourceBased(options.slot_options.to_core(kind)) + } + } + } +} + +#[derive(Clone, Copy)] +enum ResourceKind { + Workflow, + Activity, + Nexus, +} + +impl ResourceKind { + fn slot_defaults(self) -> ResourceSlotDefaults { + match self { + Self::Workflow => ResourceSlotDefaults { + minimum_slots: 2, + maximum_slots: 1_000, + ramp_throttle: Duration::from_millis(10), + }, + Self::Activity | Self::Nexus => ResourceSlotDefaults { + minimum_slots: 1, + maximum_slots: 2_000, + ramp_throttle: Duration::from_millis(50), + }, + } + } +} + +struct ResourceSlotDefaults { + minimum_slots: usize, + maximum_slots: usize, + ramp_throttle: Duration, +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn composite_tuner_rejects_different_resource_options() { + let first_options = ResourceBasedTunerOptions::builder() + .target_memory_usage(0.5) + .target_cpu_usage(0.5) + .build(); + let second_options = ResourceBasedTunerOptions::builder() + .target_memory_usage(0.6) + .target_cpu_usage(0.5) + .build(); + let result = WorkerTuner::from( + TunerHolder::builder() + .workflow_task_slot_supplier(ResourceBasedSlotsForType::new( + first_options, + Default::default(), + )) + .activity_task_slot_supplier(ResourceBasedSlotsForType::new( + second_options, + Default::default(), + )) + .local_activity_task_slot_supplier(FixedSizeSlotSupplier::new(1)) + .nexus_task_slot_supplier(FixedSizeSlotSupplier::new(1)) + .build(), + ) + .to_core(); + assert!( + result + .err() + .is_some_and(|error| error.contains("different resource-based tuner")) + ); + } +} diff --git a/crates/sdk/src/testing.rs b/crates/sdk/src/testing.rs new file mode 100644 index 000000000..1e87f32c7 --- /dev/null +++ b/crates/sdk/src/testing.rs @@ -0,0 +1,868 @@ +//! Test environments for running activity code and workflow workers. +//! +//! Activity inputs, outputs, and outbound heartbeat details stay typed. Previous heartbeat details +//! are serialized with the configured [`PayloadConverter`]; payload codecs and failure converters +//! are not used by [`ActivityEnvironment`]. +//! +//! ``` +//! use std::sync::Arc; +//! use temporalio_macros::activities; +//! use temporalio_sdk::{ +//! activities::{ActivityContext, ActivityError}, +//! testing::ActivityEnvironment, +//! }; +//! +//! struct GreetingActivities { +//! greeting: String, +//! } +//! +//! #[activities] +//! impl GreetingActivities { +//! #[activity] +//! async fn greet( +//! self: Arc, +//! _ctx: ActivityContext, +//! name: String, +//! ) -> Result { +//! Ok(format!("{}, {name}!", self.greeting)) +//! } +//! } +//! +//! # async fn example() { +//! let env = ActivityEnvironment::builder() +//! .register_activities(GreetingActivities { +//! greeting: "Hello".to_owned(), +//! }) +//! .build(); +//! +//! assert_eq!( +//! env.run(GreetingActivities::greet, "Temporal".to_owned()) +//! .await +//! .unwrap(), +//! "Hello, Temporal!" +//! ); +//! # } +//! ``` +//! +//! [`WorkflowEnvironment::start_local`] owns a Temporal CLI dev server. Its client can be passed +//! to ordinary workflow starters and workers, while the local-server type state makes shutdown +//! available only on environments that own a server. +//! +//! ```no_run +//! use temporalio_sdk::testing::{LocalWorkflowEnvironmentOptions, WorkflowEnvironment}; +//! +//! # async fn example() -> Result<(), Box> { +//! let env = WorkflowEnvironment::start_local(LocalWorkflowEnvironmentOptions::default()).await?; +//! let client = env.client().clone(); +//! // Construct workflow starters and workers with `client`. +//! # drop(client); +//! env.shutdown().await?; +//! # Ok(()) +//! # } +//! ``` + +use crate::activities::{ + ActivityContext, ActivityDefinitions, ActivityError, ActivityHeartbeatCallback, + ActivityImplementer, ActivityInfo, ExecutableActivity, +}; +use std::{ + any::Any, + collections::HashMap, + path::PathBuf, + sync::Arc, + time::{Duration, SystemTime}, +}; +use temporalio_client::{ + Client, ClientOptions, ConnectionOptions, Priority, errors::ClientConnectError, +}; +use temporalio_common::{ + RetryPolicy, + data_converters::{ + ActivitySerializationContext, GenericPayloadConverter, PayloadConversionError, + PayloadConverter, SerializationContext, SerializationContextData, TemporalSerializable, + }, + protos::temporal::api::common::v1::Payload, +}; +use tokio_util::sync::CancellationToken; +use url::Url; + +use temporalio_sdk_core::ephemeral_server::{ + EphemeralExe as CoreEphemeralExe, EphemeralExeVersion as CoreEphemeralExeVersion, + EphemeralServer, EphemeralServerError as CoreEphemeralServerError, TemporalDevServerConfig, +}; + +type ActivityImplementers = HashMap>; + +/// Where to find the Temporal server executable used by a local workflow environment. +#[derive(Debug, Clone)] +#[non_exhaustive] +pub enum EphemeralExe { + /// Use an existing executable at this path. + ExistingPath(String), + /// Download and cache an executable when necessary. + CachedDownload { + /// Version to download. + version: EphemeralExeVersion, + /// Cache directory, or the operating system's temporary directory when absent. + dest_dir: Option, + /// Maximum cache age, or no expiration when absent. + ttl: Option, + }, +} + +impl EphemeralExe { + fn into_core(self) -> CoreEphemeralExe { + match self { + EphemeralExe::ExistingPath(path) => CoreEphemeralExe::ExistingPath(path), + EphemeralExe::CachedDownload { + version, + dest_dir, + ttl, + } => CoreEphemeralExe::CachedDownload { + version: version.into_core(), + dest_dir, + ttl, + }, + } + } +} + +impl Default for EphemeralExe { + fn default() -> Self { + EphemeralExe::CachedDownload { + version: EphemeralExeVersion::Default, + dest_dir: None, + ttl: Some(Duration::from_secs(60 * 60 * 24 * 15)), + } + } +} + +/// Version of a downloadable Temporal server executable. +#[derive(Debug, Clone)] +#[non_exhaustive] +pub enum EphemeralExeVersion { + /// Resolve the server version selected for this SDK release. + Default, + /// Download a specific server version. + Fixed(String), +} + +impl EphemeralExeVersion { + fn into_core(self) -> CoreEphemeralExeVersion { + match self { + EphemeralExeVersion::Default => CoreEphemeralExeVersion::SDKDefault { + sdk_name: "sdk-rust".to_owned(), + sdk_version: env!("CARGO_PKG_VERSION").to_owned(), + }, + EphemeralExeVersion::Fixed(version) => CoreEphemeralExeVersion::Fixed(version), + } + } +} + +/// Errors encountered while downloading, starting, or stopping a local Temporal server. +#[derive(Debug, thiserror::Error)] +#[error(transparent)] +pub struct EphemeralServerError(CoreEphemeralServerError); + +impl EphemeralServerError { + fn from_core(error: CoreEphemeralServerError) -> Self { + Self(error) + } +} + +/// Options for constructing [`ActivityInfo`] with defaults suitable for an activity test. +#[derive(bon::Builder)] +#[builder( + finish_fn(name = build_internal, vis = ""), + state_mod(vis = "pub"), + on(String, into) +)] +pub struct TestActivityInfoOptions { + #[builder(default = b"test".to_vec())] + task_token: Vec, + #[builder(required, default = Some("test".to_owned()))] + workflow_type: Option, + #[builder(default = "default".to_owned())] + namespace: String, + #[builder(required, default = Some("test".to_owned()))] + workflow_id: Option, + #[builder(required, default = Some("test-run".to_owned()))] + workflow_run_id: Option, + #[builder(default = "test".to_owned())] + activity_id: String, + #[builder(default = "unknown".to_owned())] + activity_type: String, + #[builder(default = "test".to_owned())] + task_queue: String, + heartbeat_timeout: Option, + #[builder(required, default = Some(SystemTime::UNIX_EPOCH))] + scheduled_time: Option, + #[builder(required, default = Some(SystemTime::UNIX_EPOCH))] + started_time: Option, + #[builder( + required, + default = SystemTime::UNIX_EPOCH.checked_add(Duration::from_secs(1)) + )] + deadline: Option, + #[builder(default = 1)] + attempt: u32, + #[builder(required, default = Some(SystemTime::UNIX_EPOCH))] + current_attempt_scheduled_time: Option, + retry_policy: Option, + #[builder(default)] + is_local: bool, + #[builder(default)] + priority: Priority, + activity_run_id: Option, +} + +impl TestActivityInfoOptionsBuilder { + /// Build activity information from these test options. + pub fn build(self) -> ActivityInfo { + self.build_internal().into() + } +} + +impl From for ActivityInfo { + fn from(options: TestActivityInfoOptions) -> Self { + Self { + task_token: options.task_token, + workflow_type: options.workflow_type, + namespace: options.namespace, + workflow_id: options.workflow_id, + workflow_run_id: options.workflow_run_id, + activity_id: options.activity_id, + activity_type: options.activity_type, + task_queue: options.task_queue, + heartbeat_timeout: options.heartbeat_timeout, + scheduled_time: options.scheduled_time, + started_time: options.started_time, + deadline: options.deadline, + attempt: options.attempt, + current_attempt_scheduled_time: options.current_attempt_scheduled_time, + retry_policy: options.retry_policy, + is_local: options.is_local, + priority: options.priority, + activity_run_id: options.activity_run_id, + } + } +} + +/// Environment for running activity code with a test [`ActivityContext`]. +#[derive(bon::Builder)] +#[builder( + start_fn(name = builder_internal, vis = ""), + state_mod(vis = "pub") +)] +pub struct ActivityEnvironment { + #[builder(field)] + heartbeat_callback: Option, + #[builder(field)] + heartbeat_details: Vec, + #[builder(field)] + implementers: ActivityImplementers, + #[builder(field = CancellationToken::new())] + cancellation_token: CancellationToken, + #[builder( + default, + getter(name = payload_converter_ref, vis = ""), + setters(option_fn(vis = "")) + )] + payload_converter: PayloadConverter, + #[builder(default = TestActivityInfoOptions::builder().build())] + info: ActivityInfo, + #[builder(default)] + headers: HashMap, + client: Option, +} + +impl ActivityEnvironmentBuilder { + /// Register all activities implemented by an instance. + pub fn register_activities(mut self, instance: AI) -> Self + where + AI: ActivityImplementer + Send + Sync + 'static, + { + let instance = Arc::new(instance); + let mut definitions = ActivityDefinitions::default(); + AI::register_all(instance.clone(), &mut definitions); + let instance: Arc = instance; + for activity_type in definitions.names() { + self.implementers.insert(activity_type, instance.clone()); + } + self + } + + /// Observe the typed details supplied to every heartbeat. + pub fn on_heartbeat(mut self, callback: F) -> Self + where + F: Fn(Box) + Send + Sync + 'static, + { + self.heartbeat_callback = Some(Arc::new(callback)); + self + } +} + +impl ActivityEnvironmentBuilder +where + S: activity_environment_builder::State, + S::PayloadConverter: activity_environment_builder::IsSet, +{ + /// Supply heartbeat details from an activity attempt. + /// + /// Accessible via [`ActivityContext::heartbeat_details`]. + pub fn heartbeat_details(mut self, details: T) -> Result + where + T: TemporalSerializable + 'static, + { + let payload_converter = self + .payload_converter_ref() + .expect("payload converter must be set in builder state"); + let context_data = SerializationContextData::Activity(ActivitySerializationContext::new()); + let context = SerializationContext::new(&context_data, payload_converter); + self.heartbeat_details = payload_converter.to_payloads(&context, &details)?; + Ok(self) + } +} + +impl ActivityEnvironment { + /// Construct an activity environment builder. + pub fn builder() -> ActivityEnvironmentBuilder { + Self::builder_internal() + } + + /// Construct an activity environment builder using the default payload converter. + pub fn builder_with_default() + -> ActivityEnvironmentBuilder { + Self::builder_internal().payload_converter(PayloadConverter::default()) + } + + /// Run an activity. + pub async fn run( + &self, + activity: A, + input: A::Input, + ) -> Result + where + A: ExecutableActivity, + { + let receiver = if A::REQUIRES_INSTANCE { + let activity_type = activity.name(); + let implementer = self + .implementers + .get(activity_type) + .cloned() + .and_then(|instance| Arc::downcast::(instance).ok()) + .ok_or_else(|| ActivityEnvironmentError::MissingImplementer { + activity_type: activity_type.to_owned(), + })?; + Some(implementer) + } else { + None + }; + let context = ActivityContext::new_for_test( + self.info.clone(), + self.headers.clone(), + self.payload_converter.clone(), + self.cancellation_token.clone(), + self.heartbeat_details.clone(), + self.client.clone(), + self.heartbeat_callback.clone(), + ); + A::execute(receiver, context, input) + .await + .map_err(ActivityEnvironmentError::Activity) + } + + /// Cancel activity contexts created by this environment. + pub fn cancel(&self) { + self.cancellation_token.cancel(); + } +} + +/// Errors produced while running an activity in a test environment. +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum ActivityEnvironmentError { + /// An instance activity was run without registering its implementer. + #[error("activity `{activity_type}` requires an instance in order to execute")] + MissingImplementer { + /// Activity type that could not be run. + activity_type: String, + }, + /// The activity returned an error. + #[error("activity execution failed: {0:?}")] + Activity(ActivityError), +} + +/// Temporal CLI output format for a local workflow environment. +#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, derive_more::Display)] +#[non_exhaustive] +pub enum DevServerLogFormat { + /// Human-readable text output. + #[default] + #[display("text")] + Text, + /// JSON output. + #[display("json")] + Json, +} + +/// Temporal CLI logging level for a local workflow environment. +#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, derive_more::Display)] +#[non_exhaustive] +pub enum DevServerLogLevel { + /// Debug and higher-severity messages. + #[display("debug")] + Debug, + /// Informational and higher-severity messages. + #[display("info")] + Info, + /// Warning and higher-severity messages. + #[default] + #[display("warn")] + Warn, + /// Error messages only. + #[display("error")] + Error, + /// Disable logging. + #[display("never")] + Never, +} + +/// Configuration for a local Temporal CLI dev server and its client. +#[derive(Debug, Clone, bon::Builder)] +#[builder(state_mod(vis = "pub"))] +#[non_exhaustive] +pub struct LocalWorkflowEnvironmentOptions { + /// Options used to create the namespace-bound client. + #[builder(default = ClientOptions::new("default").build())] + pub client_options: ClientOptions, + /// Existing or downloadable Temporal CLI executable. + #[builder(default)] + pub server_executable: EphemeralExe, + /// Fixed frontend port, or an OS-selected port when absent. + pub port: Option, + /// Whether to start the Temporal UI. + #[builder(default)] + pub ui: bool, + /// Fixed UI port, or the server default when absent. + pub ui_port: Option, + /// SQLite database path, or in-memory storage when absent. + pub database_filename: Option, + /// Dev server log format. + #[builder(default)] + pub log_format: DevServerLogFormat, + /// Dev server log level. + #[builder(default)] + pub log_level: DevServerLogLevel, + /// Additional arguments appended to the Temporal CLI invocation. + #[builder(default)] + pub extra_args: Vec, +} + +impl Default for LocalWorkflowEnvironmentOptions { + fn default() -> Self { + Self::builder().build() + } +} + +/// State for a workflow environment backed by an externally managed server. +#[derive(Debug)] +#[non_exhaustive] +pub struct ExternalServer { + _private: (), +} + +/// State for a workflow environment that starts a local dev server. +#[derive(Debug)] +#[non_exhaustive] +pub struct LocalServer { + server: EphemeralServer, +} + +/// Client environment for workflow tests, parameterized by server ownership. +#[derive(Debug)] +#[non_exhaustive] +pub struct WorkflowEnvironment { + client: Client, + state: S, +} + +impl WorkflowEnvironment { + /// Return the client used by workflow starters and workers in this environment. + pub fn client(&self) -> &Client { + &self.client + } +} + +impl WorkflowEnvironment { + /// Wrap a client connected to an externally managed Temporal server. + pub fn from_client(client: Client) -> Self { + Self { + client, + state: ExternalServer { _private: () }, + } + } +} + +impl WorkflowEnvironment { + /// Start a local Temporal CLI dev server and connect a client to it. + pub async fn start_local( + options: LocalWorkflowEnvironmentOptions, + ) -> Result { + let database_filename = options + .database_filename + .map(|path| { + path.into_os_string().into_string().map_err(|path| { + WorkflowEnvironmentError::InvalidDatabasePath { + path: PathBuf::from(path), + } + }) + }) + .transpose()?; + let server_config = TemporalDevServerConfig::builder() + .exe(options.server_executable.into_core()) + .namespace(options.client_options.namespace.clone()) + .maybe_port(options.port) + .ui(options.ui) + .maybe_ui_port(options.ui_port) + .maybe_db_filename(database_filename) + .log(( + options.log_format.to_string(), + options.log_level.to_string(), + )) + .extra_args(options.extra_args) + .build(); + let mut server = server_config + .start_server() + .await + .map_err(EphemeralServerError::from_core) + .map_err(WorkflowEnvironmentError::ServerStart)?; + let target = Url::parse(&format!("http://{}", server.target)) + .map_err(WorkflowEnvironmentError::InvalidServerTarget)?; + let connection_options = ConnectionOptions::new(target) + .identity("temporalio-sdk-testing".to_owned()) + .client_name("temporalio-sdk".to_owned()) + .client_version(env!("CARGO_PKG_VERSION").to_owned()) + .build(); + let client = match Client::connect(connection_options, options.client_options).await { + Ok(client) => client, + Err(connect) => { + return match server.shutdown().await { + Ok(()) => Err(WorkflowEnvironmentError::ClientConnect(connect)), + Err(shutdown) => Err(WorkflowEnvironmentError::ClientConnectAndShutdown { + connect: Box::new(connect), + shutdown: Box::new(EphemeralServerError::from_core(shutdown)), + }), + }; + } + }; + Ok(Self { + client, + state: LocalServer { server }, + }) + } + + /// Shut down the local server owned by this environment. + pub async fn shutdown(mut self) -> Result<(), WorkflowEnvironmentError> { + self.state + .server + .shutdown() + .await + .map_err(EphemeralServerError::from_core) + .map_err(WorkflowEnvironmentError::ServerShutdown) + } +} + +/// Errors produced while creating or shutting down a workflow test environment. +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum WorkflowEnvironmentError { + /// The local server could not be started. + #[error("failed to start local Temporal server: {0}")] + ServerStart(#[source] EphemeralServerError), + /// A client could not connect to the newly started server. + #[error("failed to connect client to local Temporal server: {0}")] + ClientConnect(#[source] ClientConnectError), + /// Client connection and subsequent server cleanup both failed. + #[error("failed to connect client ({connect}) and shut down local server ({shutdown})")] + ClientConnectAndShutdown { + /// Client connection failure. + connect: Box, + /// Server cleanup failure. + shutdown: Box, + }, + /// Explicit local server shutdown failed. + #[error("failed to shut down local Temporal server: {0}")] + ServerShutdown(#[source] EphemeralServerError), + /// The local server target could not be represented as a URL. + #[error("invalid local Temporal server target: {0}")] + InvalidServerTarget(#[source] url::ParseError), + /// The configured database path was not valid UTF-8 for the Temporal CLI. + #[error("local Temporal database path is not valid UTF-8: {}", path.display())] + InvalidDatabasePath { + /// Invalid database path. + path: PathBuf, + }, +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::Mutex; + use temporalio_common::data_converters::{MultiArgs2, MultiArgs3}; + use temporalio_macros::activities; + + struct TestActivities { + prefix: String, + } + + #[activities] + impl TestActivities { + #[activity] + async fn echo(_ctx: ActivityContext, value: String) -> Result { + Ok(value) + } + + #[activity] + async fn prefixed( + self: std::sync::Arc, + _ctx: ActivityContext, + value: String, + ) -> Result { + Ok(format!("{}{}", self.prefix, value)) + } + + #[activity] + async fn heartbeat(ctx: ActivityContext, increment: u32) -> Result { + let previous = ctx.heartbeat_details().deserialize::()?.unwrap_or(0); + ctx.record_heartbeat(previous + increment).await?; + Ok(previous) + } + + #[activity] + async fn cancellation_state(ctx: ActivityContext) -> Result { + Ok(ctx.is_cancelled()) + } + } + + struct StaticActivities; + + #[activities] + impl StaticActivities { + #[activity] + async fn echo(_ctx: ActivityContext, value: String) -> Result { + Ok(format!("static:{value}")) + } + } + + struct ActivityMacroShapes; + + #[activities] + impl ActivityMacroShapes { + #[activity] + fn sync(_ctx: ActivityContext, value: bool) -> Result { + Ok(value.to_string()) + } + + #[activity] + async fn no_input(_ctx: ActivityContext) -> Result { + Ok("no input".to_owned()) + } + + #[activity] + async fn async_no_return(_ctx: ActivityContext, _value: String) {} + + #[activity] + fn sync_no_return(_ctx: ActivityContext) {} + } + + struct MultiArgActivities; + + #[activities] + impl MultiArgActivities { + #[activity] + async fn two_args( + _ctx: ActivityContext, + first: String, + second: i32, + ) -> Result { + Ok(format!("{first}:{second}")) + } + + #[activity] + async fn three_args( + _ctx: ActivityContext, + first: String, + second: i32, + third: bool, + ) -> Result { + Ok(format!("{first}:{second}:{third}")) + } + + #[activity] + async fn instance_two_args( + self: Arc, + _ctx: ActivityContext, + first: String, + second: i32, + ) -> Result { + let _ = self; + Ok(format!("{first}:{second}")) + } + + #[activity] + fn sync_two_args( + _ctx: ActivityContext, + first: String, + second: i32, + ) -> Result { + Ok(format!("{first}:{second}")) + } + } + + #[tokio::test] + async fn runs_static_activities_without_instance() { + let env = ActivityEnvironment::builder().build(); + + assert_eq!( + env.run(StaticActivities::echo, "value".to_owned()) + .await + .unwrap(), + "static:value" + ); + } + + #[tokio::test] + async fn runs_activities_with_instance() { + let env = ActivityEnvironment::builder() + .register_activities(TestActivities { + prefix: "pre:".to_owned(), + }) + .build(); + + assert_eq!( + env.run(TestActivities::echo, "value".to_owned()) + .await + .unwrap(), + "value" + ); + assert_eq!( + env.run(TestActivities::prefixed, "value".to_owned()) + .await + .unwrap(), + "pre:value" + ); + } + + #[tokio::test] + async fn runs_sync_and_unit_output_activities() { + let env = ActivityEnvironment::builder().build(); + + assert_eq!( + env.run(ActivityMacroShapes::sync, true).await.unwrap(), + "true" + ); + assert_eq!( + env.run(ActivityMacroShapes::no_input, ()).await.unwrap(), + "no input" + ); + env.run(ActivityMacroShapes::async_no_return, "value".to_owned()) + .await + .unwrap(); + env.run(ActivityMacroShapes::sync_no_return, ()) + .await + .unwrap(); + } + + #[tokio::test] + async fn runs_multi_argument_activities() { + let env = ActivityEnvironment::builder() + .register_activities(MultiArgActivities) + .build(); + + assert_eq!( + env.run( + MultiArgActivities::two_args, + MultiArgs2("one".to_owned(), 2), + ) + .await + .unwrap(), + "one:2" + ); + assert_eq!( + env.run( + MultiArgActivities::three_args, + MultiArgs3("one".to_owned(), 2, true), + ) + .await + .unwrap(), + "one:2:true" + ); + assert_eq!( + env.run( + MultiArgActivities::instance_two_args, + MultiArgs2("one".to_owned(), 2), + ) + .await + .unwrap(), + "one:2" + ); + assert_eq!( + env.run( + MultiArgActivities::sync_two_args, + MultiArgs2("one".to_owned(), 2), + ) + .await + .unwrap(), + "one:2" + ); + } + + #[tokio::test] + async fn missing_instance_is_an_environment_error() { + let error = ActivityEnvironment::builder() + .build() + .run(TestActivities::prefixed, "value".to_owned()) + .await + .unwrap_err(); + + assert!(matches!( + error, + ActivityEnvironmentError::MissingImplementer { .. } + )); + } + + #[tokio::test] + async fn converts_previous_and_observes_typed_outbound_heartbeat_details() { + let heartbeats = Arc::new(Mutex::new(Vec::new())); + let env = ActivityEnvironment::builder_with_default() + .heartbeat_details(4_u32) + .unwrap() + .on_heartbeat({ + let heartbeats = heartbeats.clone(); + move |details| { + let details = details + .downcast::() + .expect("heartbeat details should retain their concrete type"); + heartbeats.lock().unwrap().push(*details); + } + }) + .build(); + + assert_eq!(env.run(TestActivities::heartbeat, 3).await.unwrap(), 4); + assert_eq!(heartbeats.lock().unwrap().pop(), Some(7)); + } + + #[tokio::test] + async fn cancel_affects_contexts_created_by_environment() { + let env = ActivityEnvironment::builder().build(); + env.cancel(); + + assert!( + env.run(TestActivities::cancellation_state, ()) + .await + .unwrap() + ); + } +} diff --git a/crates/sdk/src/workflow_executor.rs b/crates/sdk/src/workflow_executor.rs index b23565f09..1151deab5 100644 --- a/crates/sdk/src/workflow_executor.rs +++ b/crates/sdk/src/workflow_executor.rs @@ -10,7 +10,7 @@ use std::{ }, task::{Context, Poll, Wake, Waker}, }; -use temporalio_workflow::runtime::is_sdk_wake; +use temporalio_workflow::__private::sdk::is_sdk_wake; /// Persists across polls to accumulate non-SDK wake detection. Each poll creates a lightweight /// waker via [`WakeTracker::new_per_poll_waker`] that shares the detection flag but has the @@ -255,7 +255,7 @@ impl WorkflowExecutor { #[cfg(test)] mod tests { use super::*; - use temporalio_workflow::runtime::SdkWakeGuard; + use temporalio_workflow::WorkflowCancellationToken; use tokio::sync::oneshot; #[tokio::test] @@ -344,34 +344,7 @@ mod tests { } #[test] - fn sdk_wake_guard_nesting() { - assert!(!is_sdk_wake()); - - let guard1 = SdkWakeGuard::new(); - assert!(is_sdk_wake()); - - { - let _guard2 = SdkWakeGuard::new(); - assert!(is_sdk_wake()); - } - assert!(is_sdk_wake()); - - drop(guard1); - assert!(!is_sdk_wake()); - } - - #[test] - fn sdk_wake_guard_panic_safety() { - let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { - let _guard = SdkWakeGuard::new(); - panic!("test panic"); - })); - assert!(result.is_err()); - assert!(!is_sdk_wake()); - } - - #[test] - fn wake_tracker_detects_non_sdk_wake() { + fn wake_tracker_distinguishes_sdk_wakes() { let tracker = WakeTracker::new(); let noop = Waker::noop(); let waker = tracker.new_per_poll_waker(noop); @@ -379,23 +352,42 @@ mod tests { waker.wake_by_ref(); assert!(tracker.take_non_sdk_wake()); - let _guard = SdkWakeGuard::new(); - waker.wake_by_ref(); + // Create an SDK owned wake + let cancellation = WorkflowCancellationToken::new(); + let mut cancelled = std::pin::pin!(cancellation.cancelled()); + let mut cx = Context::from_waker(&waker); + assert!(cancelled.as_mut().poll(&mut cx).is_pending()); + + cancellation.cancel(); + assert!(!tracker.take_non_sdk_wake()); } + struct CrossThreadWake(Waker); + + impl Wake for CrossThreadWake { + fn wake(self: Arc) { + self.wake_by_ref(); + } + + fn wake_by_ref(self: &Arc) { + let waker = self.0.clone(); + std::thread::spawn(move || waker.wake()).join().unwrap(); + } + } #[test] fn wake_tracker_cross_thread_detection() { let tracker = WakeTracker::new(); let noop = Waker::noop(); - let waker = tracker.new_per_poll_waker(noop); + let tracked_waker = tracker.new_per_poll_waker(noop); + let cross_thread_waker = Waker::from(Arc::new(CrossThreadWake(tracked_waker))); - let _guard = SdkWakeGuard::new(); + let cancellation = WorkflowCancellationToken::new(); + let mut cancelled = std::pin::pin!(cancellation.cancelled()); + let mut cx = Context::from_waker(&cross_thread_waker); + assert!(cancelled.as_mut().poll(&mut cx).is_pending()); - let handle = std::thread::spawn(move || { - waker.wake_by_ref(); - }); - handle.join().unwrap(); + cancellation.cancel(); assert!(tracker.take_non_sdk_wake()); } diff --git a/crates/sdk/src/workflow_future.rs b/crates/sdk/src/workflow_future.rs index 963b8a52a..cbb960762 100644 --- a/crates/sdk/src/workflow_future.rs +++ b/crates/sdk/src/workflow_future.rs @@ -34,17 +34,13 @@ use temporalio_common::{ }, }; use temporalio_workflow::{ - PatchActivationCallback, - runtime::{ - guest::WorkflowInstance, - host::WorkflowHost, - model::{WorkflowResult, WorkflowTermination}, - types::{ - ActivationJobResult, ActivationResult, MainRoutineCompletion, RoutineCompletion, - RoutineId, RoutineKind, RoutinePendingState, RoutinePollResult, TerminalOutcome, - UpdateRoutineCompletion, WorkflowActivation, - }, + __private::sdk::{ + ActivationJobResult, ActivationResult, MainRoutineCompletion, RoutineCompletion, RoutineId, + RoutineKind, RoutinePendingState, RoutinePollResult, TerminalOutcome, + UpdateRoutineCompletion, WorkflowActivation, WorkflowHost, WorkflowInstance, }, + InternalPatchActivationCallback as PatchActivationCallback, WorkflowResult, + WorkflowTermination, }; use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel}; @@ -541,7 +537,7 @@ impl WorkflowFuture { self.active_routines = still_active; let main_poll_result = match self - .poll_guest_routine(temporalio_workflow::runtime::types::MAIN_ROUTINE_ID, cx) + .poll_guest_routine(temporalio_workflow::__private::sdk::MAIN_ROUTINE_ID, cx) { Ok(result) => result, Err(e) => { @@ -598,10 +594,10 @@ impl WorkflowFuture { ), ); } - TerminalOutcome::Cancelled => { + TerminalOutcome::Cancelled(details) => { self.host.push_command_variant( workflow_command::Variant::CancelWorkflowExecution( - CancelWorkflowExecution {}, + CancelWorkflowExecution { details }, ), ); } diff --git a/crates/sdk/src/workflow_registry.rs b/crates/sdk/src/workflow_registry.rs index 62e6b1f85..4b4b36406 100644 --- a/crates/sdk/src/workflow_registry.rs +++ b/crates/sdk/src/workflow_registry.rs @@ -5,22 +5,17 @@ use temporalio_common::{ WorkflowDefinition, data_converters::{ DataConverter, GenericPayloadConverter, PayloadConverter, SerializationContext, - SerializationContextData, + SerializationContextData, WorkflowSerializationContext, }, protos::{ coresdk::workflow_activation::InitializeWorkflow, temporal::api::common::v1::Payload, }, }; use temporalio_workflow::{ - BaseWorkflowContext, PatchActivationCallback, - runtime::{ - entry::WorkflowImplementation, - guest::WorkflowInstance, - host::WorkflowHost, - instance::{GuestWorkflowInstance, instantiate_workflow}, - types::{WorkflowDefinitionDescriptor, WorkflowInit}, - }, + __private::sdk::{GuestWorkflowInstance, WorkflowHost, WorkflowInit, WorkflowInstance}, + BaseWorkflowContext, InternalPatchActivationCallback as PatchActivationCallback, workflow_interceptors::WorkflowInterceptorConstructor, + workflows::{WorkflowDefinitionDescriptor, WorkflowImplementation}, }; /// Host-owned execution inputs used to instantiate a single workflow run. @@ -75,6 +70,8 @@ pub struct WorkflowDefinitions { } impl WorkflowDefinitions { + // Only used by Plugins so feature flagged to avoid dead code. + #[cfg(feature = "experimental")] pub(crate) fn extend(&mut self, other: &Self) -> Result<(), WorkflowRegistrationError> { for workflow in other.workflows.values() { self.insert_workflow(workflow.definition.clone(), workflow.factory.clone())?; @@ -98,7 +95,7 @@ impl WorkflowDefinitions { { let factory = Arc::new(move |input| { let (payloads, payload_converter, base_ctx) = workflow_input_parts(input); - instantiate_workflow::(payloads, payload_converter, base_ctx) + GuestWorkflowInstance::::instantiate(payloads, payload_converter, base_ctx) .context("Failed to instantiate native workflow") }); self.insert_workflow(W::definition(), factory)?; @@ -126,10 +123,9 @@ impl WorkflowDefinitions { let factory = Arc::new(move |input| { let (payloads, payload_converter, base_ctx) = workflow_input_parts(input); - let ser_ctx = SerializationContext { - data: &SerializationContextData::Workflow, - converter: &payload_converter, - }; + let context_data = + SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let ser_ctx = SerializationContext::new(&context_data, &payload_converter); let input: ::Input = payload_converter.from_payloads(&ser_ctx, payloads)?; diff --git a/crates/sdk/src/workflow_replayer.rs b/crates/sdk/src/workflow_replayer.rs index d84c4c436..f75fe8647 100644 --- a/crates/sdk/src/workflow_replayer.rs +++ b/crates/sdk/src/workflow_replayer.rs @@ -1,7 +1,8 @@ +#[cfg(feature = "experimental")] +use crate::plugins::WorkerPlugin; use crate::{ Worker, WorkerOptions, WorkerRunError, interceptors::{self, Next, WithWorkflowReplayWorkerInput, WorkerInterceptor}, - plugins::WorkerPlugin, runtime::WorkflowErrorType, workflow_interceptors::WorkflowInterceptorConstructor, workflow_registry::{WorkflowDefinitions, WorkflowRegistrationError}, @@ -12,7 +13,9 @@ use std::{ collections::{HashMap, HashSet}, sync::Arc, }; -use temporalio_client::{ClientOptions, PluginApplyError, WorkflowHistory}; +#[cfg(feature = "experimental")] +use temporalio_client::PluginApplyError; +use temporalio_client::{ClientOptions, WorkflowHistory, errors::WorkflowInteractionError}; use temporalio_common::{ WorkflowDefinition, data_converters::DataConverter, @@ -21,14 +24,14 @@ use temporalio_common::{ WorkflowActivation, remove_from_cache::EvictionReason, workflow_activation_job::Variant as ActivationVariant, }, - temporal::api::history::v1::History, + temporal::api::history::v1::{History, HistoryEvent}, }, }; use temporalio_sdk_core::{ init_replay_worker, replay::{HistoryForReplay, ReplayWorkerInput}, }; -use temporalio_workflow::{PatchActivationCallback, runtime::entry::WorkflowImplementation}; +use temporalio_workflow::workflows::WorkflowImplementation; #[cfg(feature = "wasm-workflows")] use crate::WasmWorkflowComponent; @@ -52,6 +55,7 @@ pub struct WorkflowReplayerOptions { pub(super) workflow_interceptor_constructors: Vec, #[builder(field)] + #[cfg(feature = "experimental")] pub(super) worker_plugins: Vec>, #[cfg(feature = "wasm-workflows")] @@ -81,21 +85,20 @@ pub struct WorkflowReplayerOptions { /// Whether to detect nondeterministic future usage in workflow code. #[builder(default = true)] pub detect_nondeterministic_futures: bool, - - /// Callback controlling first non-replay patch decisions. - pub patch_activation_callback: Option, } impl WorkflowReplayerOptionsBuilder { /// Register a worker plugin with this replayer. /// /// **Experimental:** This API may change or be removed. + #[cfg(feature = "experimental")] pub fn worker_plugin(mut self, plugin: P) -> Self { self.worker_plugins.push(Arc::new(plugin)); self } /// Append a worker interceptor used during replay. + #[cfg(feature = "experimental")] pub fn worker_interceptor(mut self, interceptor: I) -> Self { self.worker_interceptors.push(Arc::new(interceptor)); self @@ -160,6 +163,7 @@ impl WorkflowReplayerOptionsBuilder impl WorkflowReplayerOptions { /// Append a worker interceptor used during replay. + #[cfg(feature = "experimental")] pub fn worker_interceptor( &mut self, interceptor: I, @@ -258,12 +262,39 @@ pub enum WorkflowReplayFailure { }, } +/// Eagerly fetched workflow history returned after replay. +#[derive(Clone, Debug)] +pub struct ReplayHistory { + events: Vec, + /// Workflow ID when it is known. + workflow_id: Option, +} + +impl ReplayHistory { + fn new(events: Vec, workflow_id: Option) -> Self { + Self { + events, + workflow_id, + } + } + + /// The history events. + pub fn events(&self) -> &[HistoryEvent] { + &self.events + } + + /// The history events. + pub fn workflow_id(&self) -> Option<&str> { + self.workflow_id.as_deref() + } +} + /// Outcome of replaying one workflow history. #[derive(Clone, Debug)] #[non_exhaustive] pub struct WorkflowReplayResult { /// History supplied to the replayer. - pub history: WorkflowHistory, + pub history: ReplayHistory, /// Replay failure, or `None` when the workflow code is compatible with the history. pub replay_failure: Option, } @@ -272,6 +303,9 @@ pub struct WorkflowReplayResult { #[derive(Debug, thiserror::Error)] #[non_exhaustive] pub enum WorkflowReplayError { + /// Fetching a streamed workflow history failed. + #[error(transparent)] + History(#[from] WorkflowInteractionError), /// The replay worker could not be created or run. #[error(transparent)] Worker(#[from] WorkflowReplayWorkerError), @@ -285,6 +319,7 @@ pub enum WorkflowReplayError { #[non_exhaustive] pub enum WorkflowReplayWorkerError { /// A plugin failed while configuring replay options. + #[cfg(feature = "experimental")] #[error(transparent)] Plugin(#[from] PluginApplyError), /// No workflow definitions were registered after plugin configuration. @@ -314,7 +349,10 @@ pub struct WorkflowReplayer { impl WorkflowReplayer { /// Construct a replayer and apply its worker plugins. - pub fn new(mut options: WorkflowReplayerOptions) -> Result { + pub fn new(options: WorkflowReplayerOptions) -> Result { + #[cfg(feature = "experimental")] + let mut options = options; + #[cfg(feature = "experimental")] crate::plugins::apply_workflow_replayer_plugins(&mut options) .map_err(WorkflowReplayWorkerError::Plugin)?; if options.workflows.is_empty() { @@ -362,27 +400,26 @@ impl WorkflowReplayer { return Ok(Vec::new()); } - let mut results = histories - .into_iter() - .map(|history| WorkflowReplayResult { - history, + let mut results = Vec::with_capacity(histories.len()); + let mut core_histories = Vec::with_capacity(histories.len()); + for history in histories { + let workflow_id = history.workflow_id().map(str::to_owned); + let replay_workflow_id = workflow_id + .as_deref() + .unwrap_or(DEFAULT_REPLAY_WORKFLOW_ID) + .to_owned(); + let events = history.into_events().await?; + core_histories.push(HistoryForReplay::new( + History { + events: events.clone(), + }, + replay_workflow_id, + )); + results.push(WorkflowReplayResult { + history: ReplayHistory::new(events, workflow_id), replay_failure: None, - }) - .collect::>(); - let core_histories: Vec<_> = results - .iter() - .map(|result| { - HistoryForReplay::new( - History { - events: result.history.events().to_vec(), - }, - result - .history - .workflow_id() - .unwrap_or(DEFAULT_REPLAY_WORKFLOW_ID), - ) - }) - .collect(); + }); + } let recorded_outcomes = Arc::new(Mutex::new(Vec::new())); let observer = ReplayOutcomeInterceptor { @@ -424,7 +461,7 @@ impl WorkflowReplayer { ) .await { - let core_worker = worker.core_worker(); + let core_worker = worker.common.worker.clone(); core_worker.initiate_shutdown(); core_worker.shutdown().await; return Err(WorkflowReplayWorkerError::Run(source).into()); @@ -448,11 +485,12 @@ impl WorkflowReplayer { .with_workflow_interceptor_constructors( self.options.workflow_interceptor_constructors.clone(), ) - .with_worker_plugins(self.options.worker_plugins.clone()) .workflow_failure_errors(self.options.workflow_failure_errors.clone()) .workflow_types_to_failure_errors(self.options.workflow_types_to_failure_errors.clone()) - .detect_nondeterministic_futures(self.options.detect_nondeterministic_futures) - .maybe_patch_activation_callback(self.options.patch_activation_callback.clone()); + .detect_nondeterministic_futures(self.options.detect_nondeterministic_futures); + #[cfg(feature = "experimental")] + let worker_options = + worker_options.with_worker_plugins(self.options.worker_plugins.clone()); #[cfg(feature = "wasm-workflows")] let worker_options = worker_options .with_wasm_workflow_components(self.options.wasm_workflow_components.clone()); diff --git a/crates/sdk/src/workflow_wasm.rs b/crates/sdk/src/workflow_wasm.rs index a936fee69..94b6d30b8 100644 --- a/crates/sdk/src/workflow_wasm.rs +++ b/crates/sdk/src/workflow_wasm.rs @@ -6,17 +6,14 @@ use temporalio_common::protos::{ coresdk::workflow_commands::WorkflowCommand, temporal::api::failure::v1::Failure, }; use temporalio_workflow::{ - PatchActivationCaller, - runtime::{ - guest::WorkflowInstance, - host::WorkflowHost, - types::{ - ActivationJobResult, ActivationResult, MainRoutineCompletion, QueryResponse, - RoutineCompletion, RoutinePendingState, RoutinePollResult, StartedRoutine, TaskFailure, - TerminalOutcome, UpdateRoutineCompletion, UpdateRoutineKind, WorkflowActivation, - WorkflowDefinitionDescriptor, WorkflowFailure, - }, + __private::sdk::{ + ActivationJobResult, ActivationResult, MainRoutineCompletion, QueryResponse, + RoutineCompletion, RoutineKind, RoutinePendingState, RoutinePollResult, StartedRoutine, + TaskFailure, TerminalOutcome, UpdateRoutineCompletion, UpdateRoutineKind, + WorkflowActivation, WorkflowFailure, WorkflowHost, WorkflowInstance, }, + PatchActivationCaller, + workflows::{UpdateDefinitionDescriptor, WorkflowDefinitionDescriptor}, }; use wasmtime::{ Config, Engine, Store, @@ -175,11 +172,9 @@ impl CompiledWasmWorkflowModule { updates: def .updates .into_iter() - .map(|u| { - temporalio_workflow::runtime::types::UpdateDefinitionDescriptor { - name: u.name, - has_validator: u.has_validator, - } + .map(|u| UpdateDefinitionDescriptor { + name: u.name, + has_validator: u.has_validator, }) .collect(), }) @@ -278,20 +273,14 @@ impl WorkflowInstance for WasmWorkflowInstance { ActivationJobResult::StartedRoutine(StartedRoutine { routine_id: routine.routine_id, kind: match routine.kind { - wit_types::RoutineKind::Main => { - temporalio_workflow::runtime::types::RoutineKind::Main - } - wit_types::RoutineKind::Signal(name) => { - temporalio_workflow::runtime::types::RoutineKind::Signal(name) - } + wit_types::RoutineKind::Main => RoutineKind::Main, + wit_types::RoutineKind::Signal(name) => RoutineKind::Signal(name), wit_types::RoutineKind::Update(update) => { - temporalio_workflow::runtime::types::RoutineKind::Update( - UpdateRoutineKind { - name: update.name, - update_id: update.update_id, - protocol_instance_id: update.protocol_instance_id, - }, - ) + RoutineKind::Update(UpdateRoutineKind { + name: update.name, + update_id: update.update_id, + protocol_instance_id: update.protocol_instance_id, + }) } }, }) @@ -341,7 +330,9 @@ impl WorkflowInstance for WasmWorkflowInstance { wit_types::TerminalOutcome::Failed(failure) => { TerminalOutcome::Failed(convert_failure(failure)) } - wit_types::TerminalOutcome::Cancelled => TerminalOutcome::Cancelled, + wit_types::TerminalOutcome::Cancelled(details) => { + TerminalOutcome::Cancelled(details.map(decode_proto)) + } wit_types::TerminalOutcome::ContinueAsNew(req) => { TerminalOutcome::ContinueAsNew(Box::new(decode_proto(req))) } diff --git a/crates/workflow/Cargo.toml b/crates/workflow/Cargo.toml index 89bbdaa1c..e7a410296 100644 --- a/crates/workflow/Cargo.toml +++ b/crates/workflow/Cargo.toml @@ -1,15 +1,22 @@ [package] name = "temporalio-workflow" -version = "0.6.0" +version = "1.0.0" edition = "2024" +rust-version = "1.88.0" authors = ["Temporal Technologies Inc. "] license-file = { workspace = true } description = "Temporal Rust workflow authoring surface" homepage = "https://temporal.io/" -repository = "https://github.com/temporalio/sdk-core" +repository = "https://github.com/temporalio/sdk-rust" keywords = ["temporal", "workflow"] categories = ["development-tools"] +[package.metadata.docs.rs] +features = ["experimental"] + +[features] +experimental = [] + [dependencies] anyhow = "1.0" bon = { workspace = true } @@ -27,17 +34,21 @@ prost-types = { workspace = true } rand = { version = "0.10", default-features = false } rand_pcg = "0.10" serde = { version = "1.0", features = ["derive"] } +siphasher = "1.0" thiserror = "2" uuid = { version = "1.18", default-features = false } -wit-bindgen = { version = "0.57.1", default-features = false, features = ["macros", "std", "realloc", "bitflags"] } +wit-bindgen = { version = "0.61.1", default-features = false, features = ["macros", "std", "realloc", "bitflags"] } + +[target.'cfg(not(target_arch = "wasm32"))'.dependencies] +rand = { version = "0.10", default-features = false, features = ["thread_rng"] } [dependencies.temporalio-common-wasm] path = "../common-wasm" -version = "0.6" +version = "~1.0.0" [dependencies.temporalio-macros] path = "../macros" -version = "0.6" +version = "~1.0.0" [dev-dependencies] rstest = "0.26" diff --git a/crates/workflow/README.md b/crates/workflow/README.md new file mode 100644 index 000000000..297dcbbcc --- /dev/null +++ b/crates/workflow/README.md @@ -0,0 +1,8 @@ +# `temporalio-workflow` + +[![crates.io](https://img.shields.io/crates/v/temporalio-workflow.svg)](https://crates.io/crates/temporalio-workflow) +[![docs.rs](https://docs.rs/temporalio-workflow/badge.svg)](https://docs.rs/temporalio-workflow) + +Part of [Temporal](https://temporal.io)'s [Rust SDK](https://github.com/temporalio/sdk-rust). + +APIs and runtime support for authoring Temporal Workflows in native Rust and WASM components. diff --git a/crates/workflow/src/component.rs b/crates/workflow/src/component.rs index 6e3e28029..8e6026e98 100644 --- a/crates/workflow/src/component.rs +++ b/crates/workflow/src/component.rs @@ -2,12 +2,12 @@ //! //! Everything in this module is internal SDK/component glue. use crate::{ - BaseWorkflowContext, PatchActivationCallback, + BaseWorkflowContext, InternalPatchActivationCallback as PatchActivationCallback, runtime::{ entry::WorkflowImplementation, guest::WorkflowInstance as RuntimeWorkflowInstance, host::WorkflowHost, - instance::instantiate_workflow, + instance::GuestWorkflowInstance, types::{ ActivationJobResult, MainRoutineCompletion, RoutineCompletion, RoutinePendingState, TerminalOutcome, UpdateRoutineCompletion, WorkflowDefinitionDescriptor, @@ -24,6 +24,10 @@ use temporalio_common_wasm::{ protos::{coresdk::workflow_commands::WorkflowCommand, temporal::api::failure::v1::Failure}, }; +/// Generated component-model bindings named by the workflow export macro. +/// +/// This module must remain public because the export macro expands in the workflow author's crate. +#[doc(hidden)] pub mod bindings { wit_bindgen::generate!({ path: "wit", @@ -46,8 +50,11 @@ use self::bindings::{ temporal::workflow_runtime::{types as wit_types, workflow_host as wit_host}, }; +/// Connects the static workflow set emitted by `export_workflow_module!` to the component adapter. pub trait StaticWorkflowComponent { + /// Describes every workflow implementation exported by the component. fn list_workflows() -> Vec; + /// Instantiates the workflow selected by the host from the component's static workflow set. fn instantiate_workflow( workflow_type: &str, init: WorkflowInit, @@ -55,6 +62,7 @@ pub trait StaticWorkflowComponent { ) -> Result, WorkflowFailure>; } +/// Adapts a [`StaticWorkflowComponent`] to the guest interface generated from the workflow WIT. pub struct ExportedComponent(PhantomData); impl wit_guest::Guest for ExportedComponent { @@ -83,6 +91,7 @@ impl wit_guest::Guest for ExportedComponent { } } +/// Adapts one runtime workflow instance to the resource interface generated from the workflow WIT. pub struct ExportedWorkflowInstance(RefCell>); impl wit_guest::GuestWorkflowInstance for ExportedWorkflowInstance { @@ -154,8 +163,10 @@ impl wit_guest::GuestWorkflowInstance for ExportedWorkflowInstance { TerminalOutcome::Failed(failure) => { wit_types::TerminalOutcome::Failed(failure.encode_to_vec()) } - TerminalOutcome::Cancelled => { - wit_types::TerminalOutcome::Cancelled + TerminalOutcome::Cancelled(details) => { + wit_types::TerminalOutcome::Cancelled( + details.map(|details| details.encode_to_vec()), + ) } TerminalOutcome::ContinueAsNew(req) => { wit_types::TerminalOutcome::ContinueAsNew( @@ -210,6 +221,7 @@ impl wit_guest::GuestWorkflowInstance for ExportedWorkflowInstance { } } +/// Instantiates a generated workflow implementation for a component without interceptors. pub fn instantiate_component_workflow( init: WorkflowInit, host: Rc, @@ -220,6 +232,7 @@ where instantiate_component_workflow_with_interceptor_constructors::(init, host, Vec::new()) } +/// Instantiates a generated workflow implementation with component-local interceptor constructors. pub fn instantiate_component_workflow_with_interceptor_constructors( init: WorkflowInit, host: Rc, @@ -240,7 +253,7 @@ where Some(patch_activation_callback), interceptor_constructors, ); - instantiate_workflow::(args, payload_converter, base_ctx).map_err(|err| { + GuestWorkflowInstance::::instantiate(args, payload_converter, base_ctx).map_err(|err| { Box::new(Failure { message: format!("Workflow input deserialization failed: {err}"), ..Default::default() diff --git a/crates/workflow/src/lib.rs b/crates/workflow/src/lib.rs index 3d3d05333..4a59b5555 100644 --- a/crates/workflow/src/lib.rs +++ b/crates/workflow/src/lib.rs @@ -1,3 +1,4 @@ +#![cfg_attr(docsrs, feature(doc_cfg))] #![warn(missing_docs)] //! Temporal workflow authoring APIs and runtime glue. @@ -11,53 +12,87 @@ pub use temporalio_macros::{ #[doc(hidden)] pub mod __private { - pub use futures_util; + pub use futures_util::{FutureExt, future::LocalBoxFuture, join, select_biased}; + + pub mod macros { + pub use crate::{ + component::{ + __wit_export, ExportedComponent, StaticWorkflowComponent, bindings, + instantiate_component_workflow, + instantiate_component_workflow_with_interceptor_constructors, + }, + runtime::{ + entry::WorkflowImplementation, + guest::WorkflowInstance, + host::WorkflowHost, + types::{ + UpdateDefinitionDescriptor, WorkflowDefinitionDescriptor, WorkflowFailure, + WorkflowInit, + }, + }, + }; + } + + pub mod sdk { + pub use crate::runtime::{ + entry::WorkflowImplementation, + guest::WorkflowInstance, + host::WorkflowHost, + instance::GuestWorkflowInstance, + is_sdk_wake, + types::{ + ActivationJobResult, ActivationResult, MAIN_ROUTINE_ID, MainRoutineCompletion, + QueryResponse, RoutineCompletion, RoutineId, RoutineKind, RoutinePendingState, + RoutinePollResult, StartedRoutine, TaskFailure, TerminalOutcome, + UpdateRoutineCompletion, UpdateRoutineKind, WorkflowActivation, WorkflowFailure, + WorkflowInit, + }, + }; + } } mod cancellation; -#[doc(hidden)] -pub mod component; -mod memo; -#[doc(hidden)] -pub mod runtime; +mod component; +mod runtime; mod workflow_context; pub mod workflow_interceptors; pub mod workflows; pub use cancellation::{WorkflowCancellationError, WorkflowCancellationToken}; -pub use memo::{MemoValue, MemoValues}; -#[doc(hidden)] -pub use runtime::model::{CancellableID, UnblockEvent}; pub use runtime::model::{TimerResult, WorkflowResult, WorkflowTermination}; -#[doc(hidden)] -pub use runtime::{SdkWakeGuard, is_sdk_wake}; pub use temporalio_common_wasm::{ - Memo, RetryPolicy, + ActivityCloseTimeouts, Memo, MemoValue, MemoValues, RetryPolicy, error::{ - ActivityExecutionError, ChildWorkflowExecutionError, ChildWorkflowStartError, RetryState, - TimeoutType, WorkflowSignalError, + ActivityExecutionError, CancelExternalWorkflowError, ChildWorkflowExecutionError, + ChildWorkflowStartError, RetryState, TimeoutType, WorkflowSignalError, }, }; pub use workflow_context::{ - ActivityCancellationType, ActivityCloseTimeouts, ActivityOptions, BaseWorkflowContext, - CancellableFuture, CancellableFutureWithReason, ChildWorkflowCancellationType, - ChildWorkflowOptions, ContinueAsNewOptions, ContinueAsNewVersioningBehavior, - ExternalWorkflowHandle, LocalActivityOptions, NamespacedWorkflowInfo, - NexusOperationCancellationType, NexusOperationOptions, ParentClosePolicy, - SignalWorkflowOptions, StartChildWorkflowExecutionFailedCause, StartChildWorkflowOutput, - StartedChildWorkflow, StartedNexusOperation, SyncWorkflowContext, TimerOptions, - VersioningIntent, WaitConditionOptions, WorkflowContext, WorkflowContextView, - WorkflowIdReusePolicy, WorkflowRandomValue, + ActivityCancellationType, ActivityOptions, BaseWorkflowContext, CancellableFuture, + CancellableFutureWithReason, ChildWorkflowCancellationType, ChildWorkflowOptions, + ContinueAsNewOptions, ExternalWorkflowHandle, LocalActivityOptions, NamespacedWorkflowInfo, + ParentClosePolicy, SignalWorkflowOptions, StartChildWorkflowExecutionFailedCause, + StartChildWorkflowOutput, StartedChildWorkflow, SyncWorkflowContext, TimerOptions, + VersioningIntent, WaitConditionOptions, WorkflowContext, WorkflowContextFuture, + WorkflowContextKey, WorkflowContextView, WorkflowIdReusePolicy, WorkflowRandomStream, + WorkflowRandomValue, +}; +#[cfg(feature = "experimental")] +pub use workflow_context::{ + ContinueAsNewVersioningBehavior, NexusOperationCancellationType, NexusOperationOptions, + PatchActivationCallback, PatchActivationInput, StartedNexusOperation, }; #[doc(hidden)] -pub use workflow_context::{PatchActivationCallback, PatchActivationCaller}; +pub use workflow_context::{ + PatchActivationCallback as InternalPatchActivationCallback, PatchActivationCaller, +}; pub use workflows::{join, join_all, select}; #[macro_export] #[doc(hidden)] macro_rules! __temporal_select { ($($tokens:tt)*) => { - $crate::__private::futures_util::select_biased! { $($tokens)* } + $crate::__private::select_biased! { $($tokens)* } }; } @@ -65,7 +100,7 @@ macro_rules! __temporal_select { #[doc(hidden)] macro_rules! __temporal_join { ($($tokens:tt)*) => { - $crate::__private::futures_util::join!($($tokens)*) + $crate::__private::join!($($tokens)*) }; } @@ -73,8 +108,8 @@ macro_rules! __temporal_join { #[doc(hidden)] macro_rules! __temporalio_export_workflow_component { ($export_type:ident) => { - $crate::component::__wit_export!( - $export_type with_types_in $crate::component::bindings + $crate::__private::macros::__wit_export!( + $export_type with_types_in $crate::__private::macros::bindings ); }; } @@ -108,24 +143,24 @@ macro_rules! export_workflow_module { ] } - impl ::temporalio_workflow::component::StaticWorkflowComponent for __TemporalWorkflowModule { + impl $crate::__private::macros::StaticWorkflowComponent for __TemporalWorkflowModule { fn list_workflows( - ) -> ::std::vec::Vec<::temporalio_workflow::runtime::types::WorkflowDefinitionDescriptor> { - ::std::vec![$(<$workflow as ::temporalio_workflow::runtime::entry::WorkflowImplementation>::definition()),*] + ) -> ::std::vec::Vec<$crate::__private::macros::WorkflowDefinitionDescriptor> { + ::std::vec![$(<$workflow as $crate::__private::macros::WorkflowImplementation>::definition()),*] } fn instantiate_workflow( workflow_type: &str, - init: ::temporalio_workflow::runtime::types::WorkflowInit, - host: ::std::rc::Rc, + init: $crate::__private::macros::WorkflowInit, + host: ::std::rc::Rc, ) -> ::std::result::Result< - ::std::boxed::Box, - ::temporalio_workflow::runtime::types::WorkflowFailure, + ::std::boxed::Box, + $crate::__private::macros::WorkflowFailure, > { match workflow_type { $( - name if name == <$workflow as ::temporalio_workflow::runtime::entry::WorkflowImplementation>::name() => { - ::temporalio_workflow::component::instantiate_component_workflow_with_interceptor_constructors::<$workflow>( + name if name == <$workflow as $crate::__private::macros::WorkflowImplementation>::name() => { + $crate::__private::macros::instantiate_component_workflow_with_interceptor_constructors::<$workflow>( init, host, __temporal_workflow_interceptor_constructors(), @@ -146,7 +181,7 @@ macro_rules! export_workflow_module { } type __TemporalWorkflowComponentExport = - ::temporalio_workflow::component::ExportedComponent<__TemporalWorkflowModule>; + $crate::__private::macros::ExportedComponent<__TemporalWorkflowModule>; ::temporalio_workflow::__temporalio_export_workflow_component!( __TemporalWorkflowComponentExport diff --git a/crates/workflow/src/memo.rs b/crates/workflow/src/memo.rs deleted file mode 100644 index 4335c4790..000000000 --- a/crates/workflow/src/memo.rs +++ /dev/null @@ -1,124 +0,0 @@ -use std::{collections::BTreeMap, rc::Rc}; - -use temporalio_common_wasm::{ - data_converters::{ - GenericPayloadConverter, PayloadConversionError, PayloadConverter, SerializationContext, - SerializationContextData, TemporalSerializable, - }, - protos::temporal::api::common::v1::Payload, -}; - -trait SerializableMemoValue { - fn to_payload( - &self, - payload_converter: &PayloadConverter, - ) -> Result; -} - -impl SerializableMemoValue for T -where - T: TemporalSerializable + 'static, -{ - fn to_payload( - &self, - payload_converter: &PayloadConverter, - ) -> Result { - payload_converter.to_payload( - &SerializationContext { - data: &SerializationContextData::Workflow, - converter: payload_converter, - }, - self, - ) - } -} - -/// A typed value used in a workflow memo update. -#[derive(Clone, derive_more::Debug)] -#[non_exhaustive] -pub struct MemoValue { - #[debug(skip)] - value: Rc, -} - -impl MemoValue { - /// Create a memo value that will be serialized with the workflow's data converter. - pub fn new(value: T) -> Self { - Self { - value: Rc::new(value), - } - } - - pub(crate) fn to_payload( - &self, - payload_converter: &PayloadConverter, - ) -> Result { - self.value.to_payload(payload_converter) - } -} - -/// A complete set of memo values for a new workflow execution. -#[derive(Clone, Debug, Default)] -#[non_exhaustive] -pub struct MemoValues { - values: BTreeMap, -} - -impl MemoValues { - /// Create an empty set of memo values. - pub fn new() -> Self { - Self::default() - } - - /// Add or replace a memo value. - pub fn insert(&mut self, key: impl Into, value: T) -> &mut Self - where - T: TemporalSerializable + 'static, - { - self.values.insert(key.into(), MemoValue::new(value)); - self - } - - pub(crate) fn encode( - &self, - payload_converter: &PayloadConverter, - ) -> Result, PayloadConversionError> { - self.values - .iter() - .map(|(key, value)| { - value - .to_payload(payload_converter) - .map(|payload| (key.clone(), payload)) - }) - .collect() - } -} - -#[cfg(test)] -mod tests { - use super::*; - use temporalio_common_wasm::{Memo, protos::temporal::api::common::v1::Memo as ProtoMemo}; - - #[test] - fn memo_values_serialize_heterogeneous_values() { - let payload_converter = PayloadConverter::default(); - let mut values = MemoValues::new(); - values - .insert("count", 7_u32) - .insert("label", "hello".to_string()); - - let memo = Memo::from_raw( - Some(ProtoMemo { - fields: values.encode(&payload_converter).unwrap(), - }), - payload_converter, - SerializationContextData::Workflow, - ); - - assert_eq!(memo.get::("count").unwrap(), Some(7)); - assert_eq!( - memo.get::("label").unwrap(), - Some("hello".to_string()) - ); - } -} diff --git a/crates/workflow/src/runtime/entry.rs b/crates/workflow/src/runtime/entry.rs index f886ffc67..b39cc07e5 100644 --- a/crates/workflow/src/runtime/entry.rs +++ b/crates/workflow/src/runtime/entry.rs @@ -14,7 +14,7 @@ use temporalio_common_wasm::{ QueryDefinition, SignalDefinition, UpdateDefinition, WorkflowDefinition, data_converters::{ GenericPayloadConverter, PayloadConversionError, PayloadConverter, SerializationContext, - SerializationContextData, TemporalSerializable, + SerializationContextData, TemporalSerializable, WorkflowSerializationContext, }, protos::temporal::api::{ common::v1::{Payload, Payloads}, @@ -288,10 +288,8 @@ pub(crate) fn serialize_output( output: &O, converter: &PayloadConverter, ) -> Result { - let ctx = SerializationContext { - data: &SerializationContextData::Workflow, - converter, - }; + let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let ctx = SerializationContext::new(&context_data, converter); converter.to_payload(&ctx, output).map_err(Into::into) } diff --git a/crates/workflow/src/runtime/instance.rs b/crates/workflow/src/runtime/instance.rs index 60bb6f210..d5d1a443e 100644 --- a/crates/workflow/src/runtime/instance.rs +++ b/crates/workflow/src/runtime/instance.rs @@ -6,7 +6,10 @@ use crate::{ InterceptedFuturePollGuard, InterceptedFuturePollKind, InterceptedFutureStatus, entry::{WorkflowError, WorkflowImplementation}, guest::WorkflowInstance, - model::{TimerResult, UnblockEvent, WorkflowTermination}, + model::{ + CancelExternalWfFailure, SignalExternalWfFailure, TimerResult, UnblockEvent, + WorkflowTermination, + }, types::{ ActivationJobResult, ActivationResult, MAIN_ROUTINE_ID, MainRoutineCompletion, QueryResponse, RoutineCompletion, RoutineId, RoutineKind, RoutinePendingState, @@ -14,12 +17,13 @@ use crate::{ UpdateRoutineCompletion, UpdateRoutineKind, WorkflowActivation, WorkflowFailure, }, }, + workflow_context::HandlerExecutionGuard, workflow_interceptors::{ ExecuteWorkflowInput, ExecuteWorkflowResult, HandleQueryInput, HandleQueryResult, HandleSignalInput, HandleSignalResult, HandleUpdateInput, HandleUpdateResult, InitializeWorkflowInput, InitializeWorkflowOutput, SyncWorkflowInterceptorContext, ValidateUpdateInput, ValidateUpdateResult, WorkflowInterceptor, WorkflowInterceptorContext, - WorkflowInterceptorFuture, WorkflowNext, serialize_workflow_output, + WorkflowInterceptorFuture, WorkflowNext, WorkflowOutputValue, serialize_workflow_output, wrong_workflow_input_type, }, }; @@ -43,7 +47,7 @@ use temporalio_common_wasm::{ WorkflowDefinition, data_converters::{ GenericPayloadConverter, PayloadConversionError, PayloadConverter, SerializationContext, - SerializationContextData, + SerializationContextData, WorkflowSerializationContext, }, error::{ApplicationFailure, OutgoingError, OutgoingWorkflowError}, protos::{ @@ -58,6 +62,7 @@ use temporalio_common_wasm::{ }, }; +/// Owns the deterministic execution state for one native workflow instance. pub struct GuestWorkflowInstance { base_ctx: BaseWorkflowContext, ctx: WorkflowContext, @@ -81,6 +86,7 @@ enum GuestRoutine { struct InterceptedFuture { inner: Fuse>, status: InterceptedFutureStatus, + _handler_execution: Option, } impl InterceptedFuture { @@ -88,6 +94,19 @@ impl InterceptedFuture { Self { inner: inner.fuse(), status, + _handler_execution: None, + } + } + + fn with_handler_execution( + inner: LocalBoxFuture<'static, T>, + status: InterceptedFutureStatus, + handler_execution: HandlerExecutionGuard, + ) -> Self { + Self { + inner: inner.fuse(), + status, + _handler_execution: Some(handler_execution), } } @@ -292,6 +311,7 @@ fn intercepted_signal_future( base_ctx: BaseWorkflowContext, interceptors: Rc<[Arc]>, input: HandleSignalInput, + handler_execution: HandlerExecutionGuard, ) -> InterceptedFuture where W: WorkflowImplementation, @@ -310,7 +330,7 @@ where call_handle_signal(&interceptors, interceptor_ctx, input, next).await } .boxed_local(); - InterceptedFuture::new(future, status) + InterceptedFuture::with_handler_execution(future, status, handler_execution) } fn intercepted_update_future( @@ -318,6 +338,7 @@ fn intercepted_update_future( base_ctx: BaseWorkflowContext, interceptors: Rc<[Arc]>, input: HandleUpdateInput, + handler_execution: HandlerExecutionGuard, ) -> InterceptedFuture where W: WorkflowImplementation, @@ -336,13 +357,15 @@ where call_handle_update(&interceptors, interceptor_ctx, input, next).await } .boxed_local(); - InterceptedFuture::new(future, status) + InterceptedFuture::with_handler_execution(future, status, handler_execution) } impl GuestWorkflowInstance where ::Input: Send, { + /// Deserializes workflow input, runs initialization interceptors, and creates an executable + /// workflow instance. pub fn instantiate( payloads: Vec, converter: PayloadConverter, @@ -350,10 +373,8 @@ where ) -> Result, PayloadConversionError> { let view = base_ctx.view(); let interceptors = base_ctx.workflow_interceptors(); - let ser_ctx = SerializationContext { - data: &SerializationContextData::Workflow, - converter: &converter, - }; + let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let ser_ctx = SerializationContext::new(&context_data, &converter); let input = converter.from_payloads(&ser_ctx, payloads)?; let (init_input, run_input) = if W::INIT_TAKES_INPUT { (Some(input), None) @@ -392,6 +413,7 @@ where ))) } + /// Creates an executable instance around an already initialized workflow value. pub fn new_with_workflow( workflow: W, base_ctx: BaseWorkflowContext, @@ -436,10 +458,8 @@ where } let converter = PayloadConverter::default(); - let ctx = SerializationContext { - data: &SerializationContextData::Workflow, - converter: &converter, - }; + let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let ctx = SerializationContext::new(&context_data, &converter); QueryResponse { result: converter .to_payload( @@ -469,14 +489,14 @@ where } }; self.base_ctx.data_converter().to_failure( - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), OutgoingError::Workflow(outgoing), ) } fn message_to_failure(&self, message: String) -> Failure { self.base_ctx.data_converter().to_failure( - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), OutgoingError::Workflow(OutgoingWorkflowError::Application(Box::new( ApplicationFailure::new(message), ))), @@ -532,11 +552,13 @@ where let future = match W::decode_signal_input(&name, payloads, converter) { Ok(Some(input)) => { let input = HandleSignalInput::new(name.clone(), input, signal.headers); + let handler_execution = self.base_ctx.track_handler(); let mut future = intercepted_signal_future::( self.ctx.clone(), self.base_ctx.clone(), self.interceptors.clone(), input, + handler_execution, ); if let ConstructionPoll::Ready(result) = Self::poll_for_construction(&self.base_ctx, &mut future)? @@ -580,6 +602,7 @@ where None => return Ok(self.rejection_for_missing_update_handler(name)), }; + let mut handler_execution = None; if run_validator && has_validator { let payloads = Payloads { payloads: input.clone(), @@ -598,6 +621,8 @@ where }; let validation_input = ValidateUpdateInput::new(id.clone(), name.clone(), decoded_input, headers.clone()); + let guard = self.base_ctx.track_handler(); + let _read_only = self.base_ctx.enter_read_only(); let validation_ctx = SyncWorkflowInterceptorContext::new(self.base_ctx.clone()); let workflow_ctx = self.ctx.clone(); let validation_next = WorkflowNext::new(move |input: ValidateUpdateInput| { @@ -626,6 +651,7 @@ where ))); } } + handler_execution = Some(guard); } let payloads = Payloads { payloads: input }; @@ -633,11 +659,14 @@ where let future = match W::decode_update_input(&name, payloads, converter) { Ok(Some(input)) => { let input = HandleUpdateInput::new(id.clone(), name.clone(), input, headers); + let handler_execution = + handler_execution.unwrap_or_else(|| self.base_ctx.track_handler()); let mut future = intercepted_update_future::( self.ctx.clone(), self.base_ctx.clone(), self.interceptors.clone(), input, + handler_execution, ); if let ConstructionPoll::Ready(result) = Self::poll_for_construction(&self.base_ctx, &mut future)? @@ -700,6 +729,7 @@ where decoded_input, query.headers, ); + let _read_only = self.base_ctx.enter_read_only(); let interceptor_ctx = SyncWorkflowInterceptorContext::new(self.base_ctx.clone()); let workflow_ctx = self.ctx.clone(); let query_next = WorkflowNext::new(move |input: HandleQueryInput| { @@ -732,10 +762,22 @@ where UnblockEvent::WorkflowComplete(event.seq, Box::new(expect_resolution(event.result))) } ActivationVariant::ResolveSignalExternalWorkflow(event) => { - UnblockEvent::SignalExternal(event.seq, event.failure) + let cause = event.cause(); + UnblockEvent::SignalExternal( + event.seq, + event + .failure + .map(|failure| SignalExternalWfFailure { failure, cause }), + ) } ActivationVariant::ResolveRequestCancelExternalWorkflow(event) => { - UnblockEvent::CancelExternal(event.seq, event.failure) + let cause = event.cause(); + UnblockEvent::CancelExternal( + event.seq, + event + .failure + .map(|failure| CancelExternalWfFailure { failure, cause }), + ) } ActivationVariant::ResolveNexusOperationStart(event) => { UnblockEvent::NexusOperationStart( @@ -767,7 +809,28 @@ where match result { Ok(result) => Ok(TerminalOutcome::Completed(result)), Err(WorkflowTermination::ContinueAsNew(req)) => Ok(TerminalOutcome::ContinueAsNew(req)), - Err(WorkflowTermination::Cancelled) => Ok(TerminalOutcome::Cancelled), + Err(WorkflowTermination::Cancelled { details }) => { + let details = details + .map(|details| { + (&*details as &dyn WorkflowOutputValue) + .serialize_payloads(&SerializationContext::new( + &SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ), + self.ctx.payload_converter(), + )) + .map(|payloads| Payloads { payloads }) + }) + .transpose() + .map_err(|err| TaskFailure { + failure: Box::new(Failure { + message: format!("Workflow payload conversion failed: {err}"), + ..Default::default() + }), + force_cause: None, + })?; + Ok(TerminalOutcome::Cancelled(details)) + } Err(WorkflowTermination::Evicted) => { panic!("workflow instances must not explicitly return eviction") } @@ -781,12 +844,16 @@ where }) } Err(WorkflowTermination::Failed(err)) => { - if self.base_ctx.cancellation_token().is_cancelled() && err.as_cancelled().is_some() + if self.base_ctx.cancellation_token().is_cancelled() + && let Some(cancelled) = err.as_cancelled() { - return Ok(TerminalOutcome::Cancelled); + let details = cancelled.raw_details().map(|payloads| Payloads { + payloads: payloads.to_vec(), + }); + return Ok(TerminalOutcome::Cancelled(details)); } let failure = self.base_ctx.data_converter().to_failure( - &SerializationContextData::Workflow, + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), temporalio_common_wasm::error::OutgoingError::Workflow(err), ); Ok(TerminalOutcome::Failed(Box::new(failure))) @@ -1115,17 +1182,6 @@ where } } -pub fn instantiate_workflow( - payloads: Vec, - converter: PayloadConverter, - base_ctx: BaseWorkflowContext, -) -> Result, PayloadConversionError> -where - ::Input: Send, -{ - GuestWorkflowInstance::::instantiate(payloads, converter, base_ctx) -} - /// Attempts to turn caught panics into something printable fn panic_formatter(panic: Box) -> Box { _panic_formatter::<&str>(panic) @@ -1173,7 +1229,7 @@ mod tests { use std::{ cell::Cell, rc::Rc, - sync::atomic::{AtomicUsize, Ordering}, + sync::atomic::{AtomicU64, AtomicUsize, Ordering}, task::Waker, }; use temporalio_common_wasm::{ @@ -1434,13 +1490,19 @@ mod tests { fn interceptor_constructors_run_before_workflow_input_decoding() { let constructor_calls = Arc::new(AtomicUsize::new(0)); let execute_calls = Arc::new(AtomicUsize::new(0)); + let constructor_random = Arc::new(AtomicU64::new(0)); let constructor_calls_ref = constructor_calls.clone(); let execute_calls_ref = execute_calls.clone(); + let constructor_random_ref = constructor_random.clone(); let constructor = WorkflowInterceptorConstructor::new(move |ctx| { assert_eq!(ctx.namespace(), "default"); assert_eq!(ctx.task_queue(), "task-queue"); assert_eq!(ctx.run_id(), "run-id"); assert_eq!(ctx.workflow_type(), DecodeFailureWorkflow::name()); + constructor_random_ref.store( + ctx.random_stream("plugin").random::(), + Ordering::Relaxed, + ); constructor_calls_ref.fetch_add(1, Ordering::Relaxed); CountingExecuteInterceptor { calls: execute_calls_ref.clone(), @@ -1452,9 +1514,20 @@ mod tests { run_id: "run-id".to_string(), initialize_workflow: InitializeWorkflow { workflow_type: DecodeFailureWorkflow::name().to_string(), + randomness_seed: 42, ..Default::default() }, }; + let expected_base_ctx = BaseWorkflowContext::from_raw( + init.clone(), + DataConverter::default(), + Rc::new(NoopHost), + None, + Vec::new(), + ); + let expected_random = expected_base_ctx.random_stream("plugin"); + let expected_constructor_random = expected_random.random::(); + let expected_next_random = expected_random.random::(); let base_ctx = BaseWorkflowContext::from_raw( init, DataConverter::default(), @@ -1462,6 +1535,7 @@ mod tests { None, vec![constructor], ); + let next_random = base_ctx.random_stream("plugin").random::(); let result = GuestWorkflowInstance::::instantiate( vec![Payload::default()], @@ -1472,5 +1546,10 @@ mod tests { assert!(result.is_err()); assert_eq!(constructor_calls.load(Ordering::Relaxed), 1); assert_eq!(execute_calls.load(Ordering::Relaxed), 0); + assert_eq!( + constructor_random.load(Ordering::Relaxed), + expected_constructor_random + ); + assert_eq!(next_random, expected_next_random); } } diff --git a/crates/workflow/src/runtime/mod.rs b/crates/workflow/src/runtime/mod.rs index 5e49236e6..4e29e99a1 100644 --- a/crates/workflow/src/runtime/mod.rs +++ b/crates/workflow/src/runtime/mod.rs @@ -7,17 +7,18 @@ use crate::runtime::types::RoutinePendingState; use std::{ cell::{Cell, RefCell}, future::Future, + marker::PhantomData, pin::Pin, rc::Rc, task::{Context, Poll}, }; -pub mod entry; -pub mod guest; -pub mod host; -pub mod instance; -pub mod model; -pub mod types; +pub(crate) mod entry; +pub(crate) mod guest; +pub(crate) mod host; +pub(crate) mod instance; +pub(crate) mod model; +pub(crate) mod types; thread_local! { static SDK_WAKE_DEPTH: Cell = const { Cell::new(0) }; @@ -237,16 +238,23 @@ pub(crate) fn mark_intercepted_handler_ready() { } /// Guard that marks the current scope as an SDK-initiated wake source. -#[doc(hidden)] -pub struct SdkWakeGuard { - _priv: (), +pub(crate) struct SdkWakeGuard { + _not_send_or_sync: PhantomData>, } impl SdkWakeGuard { - #[doc(hidden)] - pub fn new() -> Self { + /// Enters an SDK wake scope until the returned guard is dropped. + pub(crate) fn new() -> Self { SDK_WAKE_DEPTH.with(|c| c.set(c.get() + 1)); - Self { _priv: () } + Self { + _not_send_or_sync: PhantomData, + } + } +} + +impl Default for SdkWakeGuard { + fn default() -> Self { + Self::new() } } @@ -256,7 +264,7 @@ impl Drop for SdkWakeGuard { } } -#[doc(hidden)] +/// Reports whether the current thread is inside an SDK-initiated wake scope. pub fn is_sdk_wake() -> bool { SDK_WAKE_DEPTH.with(|c| c.get() > 0) } @@ -273,3 +281,43 @@ impl Future for SdkGuardedFuture { Pin::new(&mut self.0).poll(cx) } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn sdk_wake_guard_nesting() { + assert!(!is_sdk_wake()); + + { + let _guard1 = SdkWakeGuard::new(); + assert!(is_sdk_wake()); + { + let _guard2 = SdkWakeGuard::new(); + assert!(is_sdk_wake()); + } + assert!(is_sdk_wake()); + } + assert!(!is_sdk_wake()); + } + + #[test] + fn sdk_wake_guard_panic_safety() { + let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { + let _guard = SdkWakeGuard::new(); + panic!("test panic"); + })); + assert!(result.is_err()); + assert!(!is_sdk_wake()); + } + + #[test] + fn sdk_wake_guard_is_thread_local() { + let _guard = SdkWakeGuard::new(); + assert!(is_sdk_wake()); + + let child_is_sdk_wake = std::thread::spawn(is_sdk_wake).join().unwrap(); + assert!(!child_is_sdk_wake); + } +} diff --git a/crates/workflow/src/runtime/model.rs b/crates/workflow/src/runtime/model.rs index 29b740c0a..55bd8d4b4 100644 --- a/crates/workflow/src/runtime/model.rs +++ b/crates/workflow/src/runtime/model.rs @@ -1,18 +1,22 @@ //! Runtime protocol and execution model types shared by workflow code and native hosts. +#[cfg(feature = "experimental")] +mod nexus; +#[cfg(feature = "experimental")] +pub(crate) use nexus::NexusStartResult; + use crate::{ WorkflowCancellationError, runtime::types::ContinueAsNewRequest, - workflow_context::{ - ChildWfCommon, NexusUnblockData, PendingChildWorkflow, StartedNexusOperation, - }, + workflow_context::{ChildWfCommon, PendingChildWorkflow}, + workflow_interceptors::WorkflowOutputValue, }; use temporalio_common_wasm::{ WorkflowDefinition, - data_converters::PayloadConversionError, + data_converters::{PayloadConversionError, TemporalSerializable}, error::{ - ActivityExecutionError, ApplicationFailure, ChildWorkflowExecutionError, - ChildWorkflowStartError, WorkflowSignalError, + ActivityExecutionError, ApplicationFailure, CancelExternalWorkflowError, + ChildWorkflowExecutionError, ChildWorkflowStartError, WorkflowSignalError, }, protos::{ coresdk::{ @@ -24,18 +28,25 @@ use temporalio_common_wasm::{ resolve_nexus_operation_start, }, }, - temporal::api::failure::v1::Failure, + temporal::api::{ + enums::v1::{ + CancelExternalWorkflowExecutionFailedCause, + SignalExternalWorkflowExecutionFailedCause, + }, + failure::v1::Failure, + }, }, }; +#[cfg_attr(not(feature = "experimental"), allow(dead_code))] #[derive(Debug)] -pub enum UnblockEvent { +pub(crate) enum UnblockEvent { Timer(u32, TimerResult), Activity(u32, Box), WorkflowStart(u32, Box), WorkflowComplete(u32, Box), - SignalExternal(u32, Option), - CancelExternal(u32, Option), + SignalExternal(u32, Option), + CancelExternal(u32, Option), NexusOperationStart(u32, Box), NexusOperationComplete(u32, Box), } @@ -51,15 +62,25 @@ pub enum TimerResult { /// Successful result of sending a signal to an external workflow #[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub struct SignalExternalOk; +pub(crate) struct SignalExternalOk; +#[derive(Debug)] +pub(crate) struct SignalExternalWfFailure { + pub(crate) failure: Failure, + pub(crate) cause: SignalExternalWorkflowExecutionFailedCause, +} /// Result of awaiting on sending a signal to an external workflow -pub type SignalExternalWfResult = Result; +pub(crate) type SignalExternalWfResult = Result; -/// Successful result of sending a cancel request to an external workflow +/// Distinguishes external cancellation resolutions from other command results. #[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub struct CancelExternalOk; -/// Result of awaiting on sending a cancel request to an external workflow -pub type CancelExternalWfResult = Result; +pub(crate) struct CancelExternalOk; +#[derive(Debug)] +pub(crate) struct CancelExternalWfFailure { + pub(crate) failure: Failure, + pub(crate) cause: CancelExternalWorkflowExecutionFailedCause, +} +/// Internal result delivered when an external cancellation command resolves. +pub(crate) type CancelExternalWfResult = Result; pub(crate) trait Unblockable { type OtherDat; @@ -141,55 +162,9 @@ impl Unblockable for CancelExternalWfResult { } } -pub(crate) type NexusStartResult = Result; - -impl Unblockable for NexusStartResult { - type OtherDat = NexusUnblockData; - - fn unblock(ue: UnblockEvent, od: Self::OtherDat) -> Self { - let NexusUnblockData { - result_future, - schedule_seq, - base_ctx, - } = od; - match ue { - UnblockEvent::NexusOperationStart(_, result) => match *result { - resolve_nexus_operation_start::Status::OperationToken(op_token) => { - Ok(StartedNexusOperation { - operation_token: Some(op_token), - result_future, - schedule_seq, - base_ctx, - }) - } - resolve_nexus_operation_start::Status::StartedSync(_) => { - Ok(StartedNexusOperation { - operation_token: None, - result_future, - schedule_seq, - base_ctx, - }) - } - resolve_nexus_operation_start::Status::Failed(f) => Err(f), - }, - _ => panic!("Invalid unblock event for nexus operation"), - } - } -} - -impl Unblockable for NexusOperationResult { - type OtherDat = (); - - fn unblock(ue: UnblockEvent, _: Self::OtherDat) -> Self { - match ue { - UnblockEvent::NexusOperationComplete(_, result) => *result, - _ => panic!("Invalid unblock event for nexus operation complete"), - } - } -} - +#[cfg_attr(not(feature = "experimental"), allow(dead_code))] #[derive(Debug, Clone)] -pub enum CancellableID { +pub(crate) enum CancellableID { Timer(u32), Activity(u32), LocalActivity(u32), @@ -218,19 +193,44 @@ pub type WorkflowResult = Result; /// the current Workflow Task so it can be retried. /// /// Wrap an error in an [`ApplicationFailure`] to explicitly fail the Workflow Execution. -#[derive(Debug, thiserror::Error)] +#[derive(derive_more::Debug, thiserror::Error)] pub enum WorkflowTermination { + /// The Workflow Execution was cancelled, optionally with user-supplied details. #[error("Workflow cancelled")] - Cancelled, + Cancelled { + /// Optional cancellation details. + #[debug(skip)] + details: Option>, + }, + /// The workflow was evicted and must stop without producing a completion command. #[error("Workflow evicted from cache")] Evicted, + /// The workflow requested a new run with the supplied command attributes. #[error("Continue as new")] ContinueAsNew(Box), + /// The Workflow Execution failed with an error already converted for outbound handling. #[error("Workflow failed: {0}")] Failed(#[source] temporalio_common_wasm::error::OutgoingWorkflowError), } impl WorkflowTermination { + /// Construct a cancelled workflow termination without details. + pub fn cancelled() -> Self { + Self::Cancelled { details: None } + } + + /// Construct a cancelled workflow termination with details that will be converted using the + /// active payload converter. + pub fn cancelled_with_details(details: T) -> Self + where + T: TemporalSerializable + Send + Sync + 'static, + { + Self::Cancelled { + details: Some(Box::new(details)), + } + } + + /// Constructs a termination that asks the worker to continue the workflow as a new run. pub fn continue_as_new(can: ContinueAsNewRequest) -> Self { Self::ContinueAsNew(Box::new(can)) } @@ -243,7 +243,7 @@ impl WorkflowTermination { impl From for WorkflowTermination { fn from(_value: WorkflowCancellationError) -> Self { - Self::Cancelled + Self::cancelled() } } @@ -277,6 +277,12 @@ impl From for WorkflowTermination { } } +impl From for WorkflowTermination { + fn from(value: CancelExternalWorkflowError) -> Self { + Self::Failed(value.into()) + } +} + impl From for WorkflowTermination { fn from(value: ChildWorkflowStartError) -> Self { Self::Failed(value.into()) @@ -299,6 +305,7 @@ mod tests { #[case::child_start(ChildWorkflowStartError::Serialization(conversion_error()))] #[case::child_execution(ChildWorkflowExecutionError::Serialization(conversion_error()))] #[case::signal(WorkflowSignalError::Serialization(conversion_error()))] + #[case::cancel_external(CancelExternalWorkflowError::Serialization(conversion_error()))] fn conversion_error_is_preserved_in_workflow_termination>( #[case] error: T, ) { diff --git a/crates/workflow/src/runtime/model/nexus.rs b/crates/workflow/src/runtime/model/nexus.rs new file mode 100644 index 000000000..31801f483 --- /dev/null +++ b/crates/workflow/src/runtime/model/nexus.rs @@ -0,0 +1,49 @@ +use super::*; +use crate::workflow_context::{NexusUnblockData, StartedNexusOperation}; + +pub(crate) type NexusStartResult = Result; + +impl Unblockable for NexusStartResult { + type OtherDat = NexusUnblockData; + + fn unblock(ue: UnblockEvent, od: Self::OtherDat) -> Self { + let NexusUnblockData { + result_future, + schedule_seq, + base_ctx, + } = od; + match ue { + UnblockEvent::NexusOperationStart(_, result) => match *result { + resolve_nexus_operation_start::Status::OperationToken(op_token) => { + Ok(StartedNexusOperation { + operation_token: Some(op_token), + result_future, + schedule_seq, + base_ctx, + }) + } + resolve_nexus_operation_start::Status::StartedSync(_) => { + Ok(StartedNexusOperation { + operation_token: None, + result_future, + schedule_seq, + base_ctx, + }) + } + resolve_nexus_operation_start::Status::Failed(f) => Err(f), + }, + _ => panic!("Invalid unblock event for nexus operation"), + } + } +} + +impl Unblockable for NexusOperationResult { + type OtherDat = (); + + fn unblock(ue: UnblockEvent, _: Self::OtherDat) -> Self { + match ue { + UnblockEvent::NexusOperationComplete(_, result) => *result, + _ => panic!("Invalid unblock event for nexus operation complete"), + } + } +} diff --git a/crates/workflow/src/runtime/types.rs b/crates/workflow/src/runtime/types.rs index 4d8cb172d..7fe00b366 100644 --- a/crates/workflow/src/runtime/types.rs +++ b/crates/workflow/src/runtime/types.rs @@ -7,115 +7,180 @@ use temporalio_common_wasm::protos::{ workflow_activation::{InitializeWorkflow, WorkflowActivation as CoreWorkflowActivation}, workflow_commands::ContinueAsNewWorkflowExecution, }, - temporal::api::{common::v1::Payload, failure::v1::Failure}, + temporal::api::{ + common::v1::{Payload, Payloads}, + failure::v1::Failure, + }, }; +/// Host-provided state required to construct one workflow execution. #[derive(Clone, Debug, PartialEq)] pub struct WorkflowInit { + /// Namespace used when workflow code constructs namespaced commands. pub namespace: String, + /// Task queue exposed through workflow information. pub task_queue: String, + /// Run ID used to seed deterministic workflow state. pub run_id: String, + /// Initialization activation job containing workflow metadata and input. pub initialize_workflow: InitializeWorkflow, } +/// Static metadata a host needs before choosing and instantiating a workflow implementation. #[derive(Clone, Debug, PartialEq, Eq)] pub struct WorkflowDefinitionDescriptor { + /// Workflow type registered with the worker. pub workflow_type: String, + /// Whether initialization must invoke a user-defined `#[init]` method. pub has_init: bool, + /// Whether workflow input is consumed by `#[init]` instead of `#[run]`. pub init_takes_input: bool, + /// Signal names accepted by the workflow implementation. pub signals: Vec, + /// Query names accepted by the workflow implementation. pub queries: Vec, + /// Update definitions accepted by the workflow implementation. pub updates: Vec, } +/// Static metadata needed to route an update before constructing its handler future. #[derive(Clone, Debug, PartialEq, Eq)] pub struct UpdateDefinitionDescriptor { + /// Update name registered by the workflow implementation. pub name: String, + /// Whether the update has a validator that must run before its handler. pub has_validator: bool, } +/// Encoded query result returned directly while applying an activation. #[derive(Clone, Debug, PartialEq)] pub struct QueryResponse { + /// Successful payload or failure produced by the query handler. pub result: Result, } +/// Identifier assigned by the workflow runtime to a pollable routine. pub type RoutineId = u64; +/// Reserved routine identifier for the workflow's main run method. pub const MAIN_ROUTINE_ID: RoutineId = 0; +/// Activation representation shared by native and component workflow backends. pub type WorkflowActivation = CoreWorkflowActivation; +/// Identifies which workflow handler owns a runtime routine. #[derive(Clone, Debug, PartialEq)] pub enum RoutineKind { + /// The workflow's main run method. Main, + /// A signal handler, identified by signal name. Signal(String), + /// An update handler and its protocol routing metadata. Update(UpdateRoutineKind), } +/// Routing metadata required to complete an update routine through the update protocol. #[derive(Clone, Debug, PartialEq, Eq)] pub struct UpdateRoutineKind { + /// Registered update name. pub name: String, + /// User-visible update ID. pub update_id: String, + /// Protocol instance receiving the update response. pub protocol_instance_id: String, } +/// Describes a handler routine created while applying an activation. #[derive(Clone, Debug, PartialEq)] pub struct StartedRoutine { + /// Runtime-assigned identifier used for subsequent polls. pub routine_id: RoutineId, + /// Handler category and routing metadata for the new routine. pub kind: RoutineKind, } +/// Result produced synchronously while applying one activation job. #[derive(Clone, Debug, PartialEq)] pub enum ActivationJobResult { + /// The job produced no host-visible result. None, + /// The job started a routine that the host must poll. StartedRoutine(StartedRoutine), + /// A query completed without creating a persistent routine. QueryResponse(Box), + /// An update validator rejected the update before its handler started. UpdateRejected(WorkflowFailure), } +/// Results produced while applying all jobs in one activation. #[derive(Clone, Debug, PartialEq)] pub struct ActivationResult { + /// One result for each activation job, preserving activation order. pub job_results: Vec, } -pub type ContinueAsNewRequest = ContinueAsNewWorkflowExecution; +/// Command attributes used when a workflow continues as a new run. +pub(crate) type ContinueAsNewRequest = ContinueAsNewWorkflowExecution; +/// Workflow Task failure requested by the main workflow routine. #[derive(Clone, Debug, PartialEq)] pub struct TaskFailure { + /// Failure returned to Core for the current Workflow Task. pub failure: WorkflowFailure, + /// Optional server failure cause override used for failures such as nondeterminism. pub force_cause: Option, } +/// Terminal command requested when the main workflow routine finishes. #[derive(Clone, Debug, PartialEq)] pub enum TerminalOutcome { + /// Complete the Workflow Execution with the encoded result. Completed(Payload), + /// Fail the Workflow Execution with the encoded failure. Failed(WorkflowFailure), - Cancelled, + /// Cancel the Workflow Execution with optional encoded details. + Cancelled(Option), + /// Continue the Workflow Execution as a new run. ContinueAsNew(Box), } +/// Completion state returned when polling the main workflow routine. #[derive(Clone, Debug, PartialEq)] pub enum MainRoutineCompletion { + /// The main routine is intentionally blocked until a later activation. Blocked, + /// The current Workflow Task must fail without terminating the Workflow Execution. TaskFailed(TaskFailure), + /// The Workflow Execution reached a terminal or continue-as-new outcome. Terminal(Box), } +/// Completion state returned when polling an update handler routine. #[derive(Clone, Debug, PartialEq)] pub enum UpdateRoutineCompletion { + /// The update handler completed successfully. Completed { + /// Protocol instance receiving the successful response. protocol_instance_id: String, + /// Encoded update result. result: Payload, }, + /// The update handler failed after being accepted. Rejected { + /// Protocol instance receiving the failure response. protocol_instance_id: String, + /// Encoded handler failure. failure: WorkflowFailure, }, } +/// Completion state for any pollable workflow routine. #[derive(Clone, Debug, PartialEq)] pub enum RoutineCompletion { + /// Completion from the main workflow routine. Main(MainRoutineCompletion), + /// Completion from a signal handler. Signal(Result<(), WorkflowFailure>), + /// Completion from an update handler. Update(UpdateRoutineCompletion), } @@ -137,11 +202,16 @@ pub enum RoutinePendingState { InterceptorWithActivation, } +/// Outcome of polling one workflow routine. #[derive(Clone, Debug, PartialEq)] pub struct RoutinePollResult { + /// Completion emitted when the routine finished during this poll. pub completion: Option, + /// Whether polling advanced runtime state even if the routine remains pending. pub made_progress: bool, + /// Why an intercepted routine remains pending, when interceptor tracking applies. pub pending_state: Option, } +/// Failure representation shared across native and component workflow runtime boundaries. pub type WorkflowFailure = Box; diff --git a/crates/workflow/src/workflow_context.rs b/crates/workflow/src/workflow_context.rs index 27ad3cf0b..383f6c492 100644 --- a/crates/workflow/src/workflow_context.rs +++ b/crates/workflow/src/workflow_context.rs @@ -1,54 +1,61 @@ +#[cfg(feature = "experimental")] +mod nexus; mod options; mod view; +#[cfg(feature = "experimental")] +pub(crate) use nexus::NexusUnblockData; +#[cfg(feature = "experimental")] +pub use nexus::StartedNexusOperation; pub use options::{ - ActivityCancellationType, ActivityCloseTimeouts, ActivityOptions, - ChildWorkflowCancellationType, ChildWorkflowOptions, ContinueAsNewOptions, - ContinueAsNewVersioningBehavior, LocalActivityOptions, NexusOperationCancellationType, - NexusOperationOptions, ParentClosePolicy, SignalWorkflowOptions, TimerOptions, - VersioningIntent, WaitConditionOptions, WorkflowIdReusePolicy, + ActivityCancellationType, ActivityOptions, ChildWorkflowCancellationType, ChildWorkflowOptions, + ContinueAsNewOptions, LocalActivityOptions, ParentClosePolicy, SignalWorkflowOptions, + TimerOptions, VersioningIntent, WaitConditionOptions, WorkflowIdReusePolicy, }; -pub use temporalio_common_wasm::protos::coresdk::child_workflow::StartChildWorkflowExecutionFailedCause; +#[cfg(feature = "experimental")] +pub use options::{ + ContinueAsNewVersioningBehavior, NexusOperationCancellationType, NexusOperationOptions, +}; +pub use temporalio_common_wasm::error::StartChildWorkflowExecutionFailedCause; pub use view::{NamespacedWorkflowInfo, WorkflowContextView}; use crate::{ MemoValue, WorkflowCancellationError, WorkflowCancellationToken, runtime::{ - SdkGuardedFuture, SdkWakeGuard, + SdkWakeGuard, entry::WorkflowImplementation, host::WorkflowHost, mark_intercepted_future_activation, model::{ - CancelExternalWfResult, CancellableID, NexusStartResult, SignalExternalWfResult, - TimerResult, UnblockEvent, Unblockable, WorkflowTermination, + CancelExternalWfResult, CancellableID, SignalExternalWfResult, TimerResult, + UnblockEvent, Unblockable, WorkflowTermination, }, types::WorkflowInit, }, workflow_interceptors::{ - CancelExternalWorkflowInput, CancellableWorkflowOutboundFuture, - ChildWorkflowOutboundResult, ContinueAsNewInput, ScheduleActivityInput, - ScheduleLocalActivityInput, SignalWorkflowInput, SignalWorkflowResult, - SignalWorkflowTarget, StartChildWorkflowInput, StartChildWorkflowResult, - StartNexusOperationInput, StartTimerInput, WorkflowCancellationHandle, WorkflowInterceptor, + CancelExternalWorkflowInput, CancelExternalWorkflowResult, + CancellableWorkflowOutboundFuture, ChildWorkflowOutboundResult, ContinueAsNewInput, + ScheduleActivityInput, ScheduleLocalActivityInput, SignalWorkflowInput, + SignalWorkflowResult, SignalWorkflowTarget, StartChildWorkflowInput, + StartChildWorkflowResult, StartTimerInput, WorkflowCancellationHandle, WorkflowInterceptor, WorkflowInterceptorConstructor, WorkflowInterceptorContext, WorkflowNext, WorkflowOutboundFuture, WorkflowOutboundValue, call_cancel_external_workflow, call_continue_as_new, call_schedule_activity, call_schedule_local_activity, - call_signal_workflow, call_start_child_workflow, call_start_nexus_operation, - call_start_timer, + call_signal_workflow, call_start_child_workflow, call_start_timer, }, }; use futures_channel::oneshot; -use futures_util::{ - FutureExt, - future::{FusedFuture, Shared}, - task::Context, -}; +use futures_util::{FutureExt, future::FusedFuture, task::Context}; use rand::SeedableRng; use rand_pcg::Pcg64Mcg; +use siphasher::sip::SipHasher13; use std::{ + any::{Any, TypeId}, cell::{Cell, RefCell}, collections::{HashMap, HashSet}, + fmt, future::{self, Future}, + hash::Hasher, marker::PhantomData, pin::Pin, rc::Rc, @@ -62,10 +69,11 @@ use std::{ use temporalio_common_wasm::{ ActivityDefinition, Memo, SignalDefinition, WorkflowDefinition, data_converters::{ - ActivityExecutionDecodeHint, ChildWorkflowExecutionDecodeHint, - ChildWorkflowStartDecodeHint, DataConverter, GenericPayloadConverter, - PayloadConversionError, PayloadConverter, SerializationContext, SerializationContextData, - TemporalDeserializable, WorkflowSignalDecodeHint, + ActivityExecutionDecodeHint, CancelExternalWorkflowDecodeHint, + ChildWorkflowExecutionDecodeHint, ChildWorkflowStartDecodeHint, DataConverter, + GenericPayloadConverter, PayloadConversionError, PayloadConverter, SerializationContext, + SerializationContextData, TemporalDeserializable, WorkflowSerializationContext, + WorkflowSignalDecodeHint, }, error::{ ActivityExecutionError, ChildWorkflowExecutionError, ChildWorkflowStartError, @@ -74,9 +82,12 @@ use temporalio_common_wasm::{ protos::{ coresdk::{ activity_result::{ActivityResolution, Cancellation, activity_resolution}, - child_workflow::{ChildWorkflowResult, child_workflow_result}, + child_workflow::{ + ChildWorkflowResult, + StartChildWorkflowExecutionFailedCause as ProtoStartChildCause, + child_workflow_result, + }, common::NamespacedWorkflowExecution, - nexus::NexusOperationResult, workflow_activation::{ InitializeWorkflow, WorkflowActivation as CoreWorkflowActivation, resolve_child_workflow_execution_start::Status as ChildWorkflowStartStatus, @@ -139,6 +150,118 @@ macro_rules! impl_random_value { impl_random_value!(u8, u16, u32, u64, u128, i8, i16, i32, i64, i128, f32, f64); +/// A pseudo-random stream private to a stable caller-supplied name. +/// +/// Obtain a stream with [`WorkflowContext::random_stream`], +/// [`SyncWorkflowContext::random_stream`], or +/// [`crate::workflow_interceptors::WorkflowInterceptorContext::random_stream`]. Looking up the +/// same name again continues the same stream, while different names and the context's default +/// [`WorkflowContext::random`] stream do not consume one another. Clones of this value refer to the +/// same named stream. +/// +/// Draws advance workflow state without recording individual values in history, so replaying code +/// must draw from a given name in the same order. Adding or removing draws from one name does not +/// change any other name. +/// +/// Workflow reset replays the original sequence through the reset point. When Core supplies the +/// reset run's new randomness seed, all named streams start new sequences for work after that +/// point. Continue-as-new creates a new workflow run and independently seeds all streams. +#[derive(Clone)] +pub struct WorkflowRandomStream { + source: WorkflowRandomStreamSource, + name: String, +} + +#[derive(Clone)] +enum WorkflowRandomStreamSource { + Workflow(Rc>), + System(Rc>), +} + +impl WorkflowRandomStream { + /// Generates the next pseudo-random value from this named stream. + /// + /// This generator is not cryptographically secure. + pub fn random(&self) -> T + where + T: WorkflowRandomValue, + { + match &self.source { + WorkflowRandomStreamSource::Workflow(random) => { + random.borrow_mut().named_random(&self.name) + } + WorkflowRandomStreamSource::System(random) => { + ::sample(&mut random.borrow_mut()) + } + } + } + + /// Returns the stable name associated with this stream. + pub fn name(&self) -> &str { + &self.name + } +} + +fn system_random_stream_source() -> WorkflowRandomStreamSource { + #[cfg(not(target_arch = "wasm32"))] + let seed = rand::random(); + #[cfg(target_arch = "wasm32")] + let seed = { + // wasm32-unknown-unknown has no system entropy source by default, so RandomState uses the + // standard library's allocation-address fallback and varies its keys between constructions. + // This stream is only used when replay safety is not required; the important property here + // is that generating incidental identifiers does not consume workflow randomness. + let mut hasher = std::hash::BuildHasher::build_hasher(&std::hash::RandomState::new()); + std::hash::Hasher::write(&mut hasher, b"temporal-rust-system-random-stream"); + std::hash::Hasher::finish(&hasher) + }; + + WorkflowRandomStreamSource::System(Rc::new(RefCell::new(Pcg64Mcg::seed_from_u64(seed)))) +} + +fn named_random_seed(randomness_seed: u64, name: &str) -> u64 { + // The fixed second key provides domain separation and is part of replay compatibility. + let second_key = randomness_seed ^ u64::from_be_bytes(*b"temporal"); + let mut hasher = SipHasher13::new_with_keys(randomness_seed, second_key); + hasher.write(b"temporal-rust-workflow-random-stream\0"); + hasher.write(name.as_bytes()); + hasher.finish() +} + +#[derive(Clone, Debug)] +pub(super) struct WorkflowRandomState { + random: Pcg64Mcg, + randomness_seed: u64, + named_random: HashMap, +} + +impl WorkflowRandomState { + fn new(randomness_seed: u64) -> Self { + Self { + random: Pcg64Mcg::seed_from_u64(randomness_seed), + randomness_seed, + named_random: HashMap::new(), + } + } + + fn random(&mut self) -> T { + ::sample(&mut self.random) + } + + fn named_random(&mut self, name: &str) -> T { + let random = self.named_random.entry(name.to_owned()).or_insert_with(|| { + Pcg64Mcg::seed_from_u64(named_random_seed(self.randomness_seed, name)) + }); + ::sample(random) + } + + fn reseed(&mut self, randomness_seed: u64) { + self.random = Pcg64Mcg::seed_from_u64(randomness_seed); + self.randomness_seed = randomness_seed; + self.named_random.clear(); + } +} + /// Non-generic base context containing all workflow execution infrastructure. /// /// This is used internally by futures and commands that don't need typed workflow state. @@ -147,6 +270,106 @@ pub struct BaseWorkflowContext { inner: Rc, } +/// A typed key for values stored in the current workflow execution context. +/// +/// Implement this trait on a dedicated marker type shared by the workflow and its interceptors. +/// The marker type itself is the key, so different markers can store the same value type without +/// colliding. +/// +/// # Scope and propagation +/// +/// A scope inherits the values active when it is created. A nested scope shadows only its selected +/// key. Plain child futures polled inside a scope see that scope, while separately scoped +/// concurrent futures retain their own snapshots. Independently scheduled signal and update +/// handlers start without another routine's values; an inbound interceptor or the handler itself +/// can establish a handler-local scope. +/// +/// Values remain installed only while scoped workflow code is being polled. The SDK restores the +/// prior snapshot on suspension, completion, cancellation by dropping the future, and panic. This +/// prevents a value from leaking to another routine sharing the workflow's single-threaded +/// executor, or to another workflow execution. Cache eviction drops all values; replay recreates +/// them by executing the same deterministic scope calls. +/// +/// Storage is in-memory and local to one workflow run. Cross-boundary propagation is explicit: +/// outbound interceptors read values and write headers for activities, local activities, child +/// workflows, signals, Nexus operations, or continue-as-new, and inbound interceptors decode those +/// headers and establish a new scope. +/// +/// ``` +/// use temporalio_workflow::WorkflowContextKey; +/// +/// struct RequestId; +/// +/// impl WorkflowContextKey for RequestId { +/// type Value = String; +/// } +/// ``` +pub trait WorkflowContextKey: 'static { + /// Value stored under this key. + type Value: 'static; +} + +type WorkflowContextValues = Rc>>; + +#[derive(Clone, Default)] +pub(super) struct WorkflowContextValueStore { + current: Rc>, +} + +impl fmt::Debug for WorkflowContextValueStore { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("WorkflowContextValueStore") + .finish_non_exhaustive() + } +} + +impl WorkflowContextValueStore { + pub(super) fn context_value(&self) -> Option> { + self.current + .borrow() + .get(&TypeId::of::()) + .cloned() + .and_then(|value| value.downcast().ok()) + } +} + +/// A future that installs workflow context values while polling its inner future. +/// +/// Create this with [`WorkflowContext::with_context_value`] or +/// [`WorkflowInterceptorContext::with_context_value`](crate::workflow_interceptors::WorkflowInterceptorContext::with_context_value). +/// Values survive suspension and are isolated from concurrently polled workflow futures. +#[must_use = "futures do nothing unless polled"] +pub struct WorkflowContextFuture { + base: BaseWorkflowContext, + values: WorkflowContextValues, + inner: Pin>, +} + +impl Future for WorkflowContextFuture { + type Output = F::Output; + + fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { + let this = self.get_mut(); + let _guard = this.base.install_context_values(this.values.clone()); + this.inner.as_mut().poll(cx) + } +} + +struct WorkflowContextRestoreGuard { + base: BaseWorkflowContext, + previous: Option, +} + +impl Drop for WorkflowContextRestoreGuard { + fn drop(&mut self) { + self.base + .inner + .context_values + .current + .replace(self.previous.take().expect("context is restored once")); + } +} + /// Input provided to a worker's patch activation callback. #[derive(Clone, Debug)] #[non_exhaustive] @@ -186,6 +409,8 @@ impl PatchActivationCaller { run_id, init, payload_converter, + false, + None, ), } } @@ -227,14 +452,17 @@ impl BaseWorkflowContext { activation: &CoreWorkflowActivation, is_replaying_history_events: bool, ) { - let mut shared = self.inner.shared.borrow_mut(); - shared.activation = activation.clone(); - shared.is_replaying_history_events = is_replaying_history_events; - if let Some(seed) = activation.jobs.iter().find_map(|job| match &job.variant { - Some(ActivationVariant::UpdateRandomSeed(attrs)) => Some(attrs.randomness_seed), - _ => None, - }) { - shared.random = Pcg64Mcg::seed_from_u64(seed); + let new_seed = { + let mut shared = self.inner.shared.borrow_mut(); + shared.activation = activation.clone(); + shared.is_replaying_history_events = is_replaying_history_events; + activation.jobs.iter().find_map(|job| match &job.variant { + Some(ActivationVariant::UpdateRandomSeed(attrs)) => Some(attrs.randomness_seed), + _ => None, + }) + }; + if let Some(seed) = new_seed { + self.inner.random.borrow_mut().reseed(seed); } } @@ -242,8 +470,14 @@ impl BaseWorkflowContext { where T: WorkflowRandomValue, { - let random = &mut self.inner.shared.borrow_mut().random; - ::sample(random) + self.inner.random.borrow_mut().random() + } + + pub(crate) fn random_stream(&self, name: impl Into) -> WorkflowRandomStream { + WorkflowRandomStream { + source: WorkflowRandomStreamSource::Workflow(self.inner.random.clone()), + name: name.into(), + } } fn uuid4(&self) -> String { @@ -317,6 +551,18 @@ impl BaseWorkflowContext { self.inner.shared.borrow().is_replaying_history_events } + fn requires_replay_safety(&self) -> bool { + self.inner.requires_replay_safety.get() + } + + pub(crate) fn enter_read_only(&self) -> ReadOnlyGuard { + let previous = self.inner.requires_replay_safety.replace(false); + ReadOnlyGuard { + base: self.clone(), + previous, + } + } + /// Returns the payload converter used by the worker running this workflow. pub fn payload_converter(&self) -> &PayloadConverter { self.inner.data_converter.payload_converter() @@ -378,7 +624,10 @@ impl BaseWorkflowContext { self.inner.run_id.clone(), initial_information, self.inner.data_converter.payload_converter().clone(), + self.requires_replay_safety(), + Some(self.inner.random.clone()), ) + .with_context_values(self.inner.context_values.clone()) } } @@ -481,14 +730,48 @@ struct WorkflowContextInner { cancellation_token: WorkflowCancellationToken, cancelled_operations: RefCell>, shared: RefCell, + random: Rc>, seq_nums: RefCell, data_converter: DataConverter, patch_activation_callback: Option, state_mutated: Cell, + active_handlers: Cell, + requires_replay_safety: Cell, + condition_wakers: RefCell>, current_waker: RefCell>, + context_values: WorkflowContextValueStore, workflow_interceptors: Rc<[Arc]>, } +pub(crate) struct HandlerExecutionGuard { + base: BaseWorkflowContext, +} + +pub(crate) struct ReadOnlyGuard { + base: BaseWorkflowContext, + previous: bool, +} + +impl Drop for ReadOnlyGuard { + fn drop(&mut self) { + self.base.inner.requires_replay_safety.set(self.previous); + } +} + +impl Drop for HandlerExecutionGuard { + fn drop(&mut self) { + let active_handlers = self.base.inner.active_handlers.get(); + debug_assert!(active_handlers > 0, "handler execution count underflow"); + self.base + .inner + .active_handlers + .set(active_handlers.saturating_sub(1)); + if active_handlers <= 1 { + self.base.wake_condition_waiters(); + } + } +} + /// Identical to [`CancellableID`], but only containing command type and seq number, omitting any reason. #[derive(Eq, Hash, PartialEq)] enum CancellableSeqNum { @@ -545,10 +828,6 @@ pub struct WorkflowContext { sync: SyncWorkflowContext, /// The workflow instance workflow_state: Rc>, - /// Wakers registered by `wait_condition` futures. Drained and woken on - /// every `state_mut` call so that waker-based combinators (e.g. - /// `FuturesOrdered`) re-poll the condition after state changes. - condition_wakers: Rc>>, } impl Clone for WorkflowContext { @@ -556,7 +835,6 @@ impl Clone for WorkflowContext { Self { sync: self.sync.clone(), workflow_state: self.workflow_state.clone(), - condition_wakers: self.condition_wakers.clone(), } } } @@ -577,13 +855,20 @@ impl BaseWorkflowContext { run_id, initialize_workflow, } = init; + let random = Rc::new(RefCell::new(WorkflowRandomState::new( + initialize_workflow.randomness_seed, + ))); + let context_values = WorkflowContextValueStore::default(); let view = WorkflowContextView::new( namespace, task_queue, run_id, initialize_workflow, data_converter.payload_converter().clone(), - ); + true, + Some(random.clone()), + ) + .with_context_values(context_values.clone()); let workflow_interceptors = workflow_interceptor_constructors .into_iter() .map(|constructor| constructor.construct(&view)) @@ -596,7 +881,6 @@ impl BaseWorkflowContext { task_queue, run_id, shared: RefCell::new(WorkflowContextSharedData { - random: Pcg64Mcg::seed_from_u64(init_workflow_job.randomness_seed), memo: init_workflow_job.memo.clone().unwrap_or_default(), search_attributes: init_workflow_job .search_attributes @@ -608,6 +892,7 @@ impl BaseWorkflowContext { current_details: Default::default(), notified_patches: Default::default(), }), + random, initial_information: init_workflow_job, runtime: WorkflowRuntimeState::new(host), cancellation_token: WorkflowCancellationToken::new(), @@ -618,17 +903,62 @@ impl BaseWorkflowContext { next_child_workflow_sequence_number: 1, next_cancel_external_wf_sequence_number: 1, next_signal_external_wf_sequence_number: 1, + #[cfg(feature = "experimental")] next_nexus_op_sequence_number: 1, }), data_converter, patch_activation_callback, state_mutated: Cell::new(false), + active_handlers: Cell::new(0), + requires_replay_safety: Cell::new(true), + condition_wakers: Default::default(), current_waker: RefCell::new(None), + context_values, workflow_interceptors, }), } } + pub(crate) fn context_value(&self) -> Option> { + self.inner.context_values.context_value::() + } + + fn context_values_with(&self, value: K::Value) -> WorkflowContextValues { + let mut values = self.inner.context_values.current.borrow().as_ref().clone(); + values.insert(TypeId::of::(), Rc::new(value)); + Rc::new(values) + } + + pub(crate) fn with_context_value( + &self, + value: K::Value, + future: F, + ) -> WorkflowContextFuture { + WorkflowContextFuture { + base: self.clone(), + values: self.context_values_with::(value), + inner: Box::pin(future), + } + } + + pub(crate) fn with_context_value_sync( + &self, + value: K::Value, + f: impl FnOnce() -> R, + ) -> R { + let values = self.context_values_with::(value); + let _guard = self.install_context_values(values); + f() + } + + fn install_context_values(&self, values: WorkflowContextValues) -> WorkflowContextRestoreGuard { + let previous = self.inner.context_values.current.replace(values); + WorkflowContextRestoreGuard { + base: self.clone(), + previous: Some(previous), + } + } + pub(crate) fn workflow_interceptors(&self) -> Rc<[Arc]> { self.inner.workflow_interceptors.clone() } @@ -644,6 +974,24 @@ impl BaseWorkflowContext { self.inner.state_mutated.set(true); } + pub(crate) fn all_handlers_finished(&self) -> bool { + self.inner.active_handlers.get() == 0 + } + + pub(crate) fn track_handler(&self) -> HandlerExecutionGuard { + self.inner + .active_handlers + .set(self.inner.active_handlers.get() + 1); + HandlerExecutionGuard { base: self.clone() } + } + + fn wake_condition_waiters(&self) { + let _guard = SdkWakeGuard::new(); + for waker in self.inner.condition_wakers.borrow_mut().drain(..) { + waker.wake(); + } + } + pub(crate) fn take_runtime_progress(&self) -> bool { self.inner.runtime.take_progress() } @@ -823,10 +1171,9 @@ impl BaseWorkflowContext { } }; let payload_converter = base_ctx.inner.data_converter.payload_converter(); - let ctx = SerializationContext { - data: &SerializationContextData::Workflow, - converter: payload_converter, - }; + let context_data = + SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let ctx = SerializationContext::new(&context_data, payload_converter); match payload_converter.to_payloads(&ctx, &input) { Ok(payloads) => { let cancellation_token = opts @@ -919,10 +1266,9 @@ impl BaseWorkflowContext { } }; let payload_converter = base_ctx.inner.data_converter.payload_converter(); - let ctx = SerializationContext { - data: &SerializationContextData::Workflow, - converter: payload_converter, - }; + let context_data = + SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let ctx = SerializationContext::new(&context_data, payload_converter); match payload_converter.to_payloads(&ctx, &input) { Ok(payloads) => { let cancellation_token = opts @@ -1003,10 +1349,9 @@ impl BaseWorkflowContext { } }; let payload_converter = base_ctx.inner.data_converter.payload_converter(); - let ctx = SerializationContext { - data: &SerializationContextData::Workflow, - converter: payload_converter, - }; + let context_data = + SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let ctx = SerializationContext::new(&context_data, payload_converter); let payloads = match payload_converter.to_payloads(&ctx, &input) { Ok(payloads) => payloads, Err(err) => { @@ -1157,10 +1502,9 @@ impl BaseWorkflowContext { } }; let payload_converter = base_ctx.data_converter().payload_converter(); - let ctx = SerializationContext { - data: &SerializationContextData::Workflow, - converter: payload_converter, - }; + let context_data = + SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let ctx = SerializationContext::new(&context_data, payload_converter); let payloads = match payload_converter.to_payloads(&ctx, &input) { Ok(payloads) => payloads, Err(err) => { @@ -1236,7 +1580,7 @@ impl BaseWorkflowContext { fn cancel_external_workflow( &self, input: CancelExternalWorkflowInput, - ) -> WorkflowOutboundFuture { + ) -> WorkflowOutboundFuture { let base_ctx = self.clone(); let next = WorkflowNext::new(move |input: CancelExternalWorkflowInput| { let seq = base_ctx @@ -1244,7 +1588,7 @@ impl BaseWorkflowContext { .seq_nums .borrow_mut() .next_cancel_external_wf_seq(); - let (cmd, unblocker) = WFCommandFut::new(); + let (cmd, unblocker) = WFCommandFut::::new(); base_ctx .inner .runtime @@ -1263,7 +1607,21 @@ impl BaseWorkflowContext { ) .into(), ); - WorkflowOutboundFuture::new(cmd) + let data_converter = base_ctx.data_converter().clone(); + WorkflowOutboundFuture::new(async move { + match cmd.await { + Ok(_) => Ok(()), + Err(error) => { + let context = + SerializationContextData::Workflow(WorkflowSerializationContext::new()); + Err(data_converter.to_error( + &context, + error.failure, + CancelExternalWorkflowDecodeHint::new(error.cause), + )?) + } + } + }) }); let interceptors = self.inner.workflow_interceptors.clone(); let future = call_cancel_external_workflow( @@ -1274,64 +1632,29 @@ impl BaseWorkflowContext { ); self.prepare_outbound_future(future) } +} + +impl SyncWorkflowContext { + /// Return the value associated with key type `K` in the current workflow context scope. + /// + /// The returned [`Rc`] makes lookup inexpensive without requiring stored values to implement + /// [`Clone`]. Values exist only in memory for this workflow run and are rebuilt during replay. + pub fn context_value(&self) -> Option> { + self.base.context_value::() + } - pub(crate) fn start_nexus_operation( + /// Run synchronous workflow code with `value` installed for key type `K`. + /// + /// Nested calls inherit other current values and shadow the same key. The previous context is + /// restored when `f` returns or unwinds. + pub fn with_context_value_sync( &self, - opts: NexusOperationOptions, - ) -> impl CancellableFuture { - let input = StartNexusOperationInput::new(opts); - let base_ctx = self.clone(); - let next = WorkflowNext::new(move |input: StartNexusOperationInput| { - let mut opts = input.into_options(); - let cancellation_token = opts - .cancellation_token - .take() - .unwrap_or_else(|| base_ctx.cancellation_token()); - let seq = base_ctx.inner.seq_nums.borrow_mut().next_nexus_op_seq(); - let (result_future, unblocker) = - CancellableWFCommandFut::new(CancellableID::NexusOp(seq), base_ctx.clone()); - base_ctx - .inner - .runtime - .register_unblocker(PendingCommandId::NexusOpComplete(seq), unblocker); - base_ctx - .inner - .runtime - .host - .push_command(opts.into_command(seq)); - let result_future = CancellableWorkflowOutboundFuture::new( - result_future, - base_ctx.cancellation_handle(CancellableID::NexusOp(seq)), - ) - .with_cancellation_token(cancellation_token) - .shared(); - let (cmd, unblocker) = CancellableWFCommandFut::new_with_dat( - CancellableID::NexusOp(seq), - NexusUnblockData { - result_future: result_future.clone(), - schedule_seq: seq, - base_ctx: base_ctx.clone(), - }, - base_ctx.clone(), - ); - base_ctx - .inner - .runtime - .register_unblocker(PendingCommandId::NexusOpStart(seq), unblocker); - cancellable_outbound(cmd) - }); - let interceptors = self.inner.workflow_interceptors.clone(); - let future = call_start_nexus_operation( - interceptors, - WorkflowInterceptorContext::new(self.clone()), - input, - next, - ); - self.prepare_cancellable_outbound_future(future) + value: K::Value, + f: impl FnOnce() -> R, + ) -> R { + self.base.with_context_value_sync::(value, f) } -} -impl SyncWorkflowContext { /// Return the workflow's unique identifier pub fn workflow_id(&self) -> &str { &self.base.inner.initial_information.workflow_id @@ -1392,7 +1715,7 @@ impl SyncWorkflowContext { Memo::from_raw( Some(self.base.inner.shared.borrow().memo.clone()), self.payload_converter().clone(), - SerializationContextData::Workflow, + SerializationContextData::Workflow(WorkflowSerializationContext::new()), ) } @@ -1414,6 +1737,25 @@ impl SyncWorkflowContext { self.base.uuid4() } + /// Returns the deterministic pseudo-random stream associated with `name`. + /// + /// Repeated lookup of the same name continues the prior stream. Different names are isolated + /// from one another and from [`Self::random`]. Keep the name stable across workflow replays. + /// + /// # Example + /// + /// ```no_run + /// # use temporalio_workflow::{SyncWorkflowContext, WorkflowRandomStream}; + /// # fn choose(ctx: &SyncWorkflowContext) { + /// let stream: WorkflowRandomStream = ctx.random_stream("example.com/orders/tiebreaker"); + /// let choice = stream.random::(); + /// # let _ = choice; + /// # } + /// ``` + pub fn random_stream(&self, name: impl Into) -> WorkflowRandomStream { + self.base.random_stream(name) + } + /// Returns true if the current workflow task is happening under replay pub fn is_replaying(&self) -> bool { self.base.inner.shared.borrow().activation.is_replaying @@ -1424,6 +1766,13 @@ impl SyncWorkflowContext { self.base.inner.shared.borrow().is_replaying_history_events } + /// Returns whether all currently dispatched signal and update handlers have finished. + /// + /// This includes the current handler invocation, if any, and all inbound interceptor work. + pub fn all_handlers_finished(&self) -> bool { + self.base.all_handlers_finished() + } + /// Returns true if the server suggests this workflow should continue-as-new pub fn continue_as_new_suggested(&self) -> bool { self.base @@ -1437,6 +1786,7 @@ impl SyncWorkflowContext { /// Returns true if the workflow's target worker deployment version changed. /// /// This experimental signal is intended for workers using worker deployment versioning. + #[cfg(feature = "experimental")] pub fn target_worker_deployment_version_changed(&self) -> bool { self.base .inner @@ -1501,10 +1851,9 @@ impl SyncWorkflowContext { Err(_) => return Err(outbound_type_error("continue-as-new input").into()), }; let pc = base_ctx.data_converter().payload_converter(); - let ctx = SerializationContext { - data: &SerializationContextData::Workflow, - converter: pc, - }; + let context_data = + SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let ctx = SerializationContext::new(&context_data, pc); let arguments = pc .to_payloads(&ctx, &*input) .map_err(WorkflowTermination::from)?; @@ -1647,6 +1996,7 @@ impl SyncWorkflowContext { let res = if deprecated || replaying || notified { !replaying || notified } else if let Some(callback) = &self.base.inner.patch_activation_callback { + let _read_only = self.base.enter_read_only(); callback(PatchActivationInput { workflow_info: self.base.view(), patch_id: patch_id.to_string(), @@ -1727,17 +2077,20 @@ impl SyncWorkflowContext { where K: Into, { + let payload_converter = self.payload_converter(); + let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new()); + let context = SerializationContext::new(&context_data, payload_converter); let mut fields = HashMap::new(); let mut local_updates = Vec::new(); for (key, value) in updates { let key = key.into(); let (command_payload, local_payload) = match value { Some(value) => { - let payload = value.to_payload(self.payload_converter())?; + let payload = payload_converter.to_payload(&context, &value)?; (payload.clone(), Some(payload)) } None => ( - MemoValue::new(()).to_payload(self.payload_converter())?, + payload_converter.to_payload(&context, &MemoValue::new(()))?, None, ), }; @@ -1781,14 +2134,6 @@ impl SyncWorkflowContext { self.base.inner.runtime.set_forced_wft_failure(with.into()); } - /// Start a nexus operation - pub fn start_nexus_operation( - &self, - opts: NexusOperationOptions, - ) -> impl CancellableFuture { - self.base.start_nexus_operation(opts) - } - /// Create a read-only view of this context. pub(crate) fn view(&self) -> WorkflowContextView { self.base.view() @@ -1805,7 +2150,6 @@ impl WorkflowContext { _phantom: PhantomData, }, workflow_state, - condition_wakers: Rc::new(RefCell::new(Vec::new())), } } @@ -1818,7 +2162,6 @@ impl WorkflowContext { _phantom: PhantomData, }, workflow_state: self.workflow_state.clone(), - condition_wakers: self.condition_wakers.clone(), } } @@ -1834,6 +2177,37 @@ impl WorkflowContext { // --- Delegated methods from SyncWorkflowContext --- + /// Return the value associated with key type `K` in the current workflow context scope. + pub fn context_value(&self) -> Option> { + self.sync.context_value::() + } + + /// Poll `future` with `value` installed for key type `K`. + /// + /// The scope captures the context active when this method is called. Nested scopes inherit + /// other values and shadow the same key. Context is restored after every poll, including when + /// the future completes or panics, so concurrent workflow branches and handlers cannot observe + /// one another's scoped values. + /// + /// Context values are runtime-only. They are not recorded in history or automatically placed + /// in command headers; outbound interceptors can read them and propagate selected values. + pub fn with_context_value( + &self, + value: K::Value, + future: F, + ) -> WorkflowContextFuture { + self.sync.base.with_context_value::(value, future) + } + + /// Run synchronous workflow code with `value` installed for key type `K`. + pub fn with_context_value_sync( + &self, + value: K::Value, + f: impl FnOnce() -> R, + ) -> R { + self.sync.with_context_value_sync::(value, f) + } + /// Return the workflow's unique identifier pub fn workflow_id(&self) -> &str { self.sync.workflow_id() @@ -1898,6 +2272,13 @@ impl WorkflowContext { self.sync.uuid4() } + /// Returns the deterministic pseudo-random stream associated with `name`. + /// + /// See [`SyncWorkflowContext::random_stream`]. + pub fn random_stream(&self, name: impl Into) -> WorkflowRandomStream { + self.sync.random_stream(name) + } + /// Returns true if the current workflow task is happening under replay pub fn is_replaying(&self) -> bool { self.sync.is_replaying() @@ -1908,6 +2289,28 @@ impl WorkflowContext { self.sync.is_replaying_history_events() } + /// Returns whether all currently dispatched signal and update handlers have finished. + /// + /// Consider waiting on this condition before completing or continuing as new so in-progress + /// handlers are not interrupted. Use a cloned context in [`Self::wait_condition`]: + /// + /// ```rust + /// # use temporalio_workflow::{WorkflowContext, WorkflowResult}; + /// # struct MyWorkflow; + /// # async fn wait_for_handlers(ctx: &mut WorkflowContext) -> WorkflowResult<()> { + /// let wait_condition_ctx = ctx.clone(); + /// ctx.wait_condition(move |_| wait_condition_ctx.all_handlers_finished()) + /// .await?; + /// # Ok(()) + /// # } + /// ``` + /// + /// The check includes inbound interceptor work and the current handler invocation, if any. + /// It does not prevent future signal or update handlers from starting. + pub fn all_handlers_finished(&self) -> bool { + self.sync.all_handlers_finished() + } + /// Returns true if the server suggests this workflow should continue-as-new pub fn continue_as_new_suggested(&self) -> bool { self.sync.continue_as_new_suggested() @@ -1916,6 +2319,7 @@ impl WorkflowContext { /// Returns true if the workflow's target worker deployment version changed. /// /// This experimental signal is intended for workers using worker deployment versioning. + #[cfg(feature = "experimental")] pub fn target_worker_deployment_version_changed(&self) -> bool { self.sync.target_worker_deployment_version_changed() } @@ -2093,14 +2497,6 @@ impl WorkflowContext { self.sync.force_task_fail(with) } - /// Start a nexus operation - pub fn start_nexus_operation( - &self, - opts: NexusOperationOptions, - ) -> impl CancellableFuture { - self.sync.start_nexus_operation(opts) - } - /// Access workflow state immutably via closure. /// /// The borrow is scoped to the closure and cannot escape, preventing @@ -2119,10 +2515,7 @@ impl WorkflowContext { /// `FuturesOrdered`) re-poll them on the next pass. pub fn state_mut(&self, f: impl FnOnce(&mut W) -> R) -> R { let result = f(&mut *self.workflow_state.borrow_mut()); - let _guard = SdkWakeGuard::new(); - for waker in self.condition_wakers.borrow_mut().drain(..) { - waker.wake(); - } + self.sync.base.wake_condition_waiters(); self.sync.base.set_state_mutated(); result } @@ -2173,7 +2566,12 @@ impl WorkflowContext { } else if cancelled.as_mut().poll(cx).is_ready() { Poll::Ready(Err(WorkflowCancellationError::new(token.reason()))) } else { - self.condition_wakers.borrow_mut().push(cx.waker().clone()); + self.sync + .base + .inner + .condition_wakers + .borrow_mut() + .push(cx.waker().clone()); Poll::Pending } }) @@ -2187,6 +2585,7 @@ struct WfCtxProtectedDat { next_child_workflow_sequence_number: u32, next_cancel_external_wf_sequence_number: u32, next_signal_external_wf_sequence_number: u32, + #[cfg(feature = "experimental")] next_nexus_op_sequence_number: u32, } @@ -2216,11 +2615,6 @@ impl WfCtxProtectedDat { self.next_signal_external_wf_sequence_number += 1; seq } - fn next_nexus_op_seq(&mut self) -> u32 { - let seq = self.next_nexus_op_sequence_number; - self.next_nexus_op_sequence_number += 1; - seq - } } #[derive(Clone, Debug)] @@ -2233,7 +2627,6 @@ struct WorkflowContextSharedData { memo: ProtoMemo, is_replaying_history_events: bool, search_attributes: ProtoSearchAttributes, - random: Pcg64Mcg, /// Current details string, surfaced via the workflow metadata query. current_details: String, } @@ -2525,6 +2918,8 @@ impl Future for LATimerBackoffFut { .expect("duration converts ok"), cancellation_token: Some(self.cancellation_token.clone()), summary: None, + #[cfg(feature = "experimental")] + event_group_markers: self.la_opts.event_group_markers.clone(), }); self.timer_fut = Some(Box::pin(timer_f)); self.next_attempt = b.attempt; @@ -2596,61 +2991,73 @@ where fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { let this = self.get_mut(); - let poll = match this { - ActivityFut::Errored { error, .. } => { - Poll::Ready(Err(*error.take().expect("polled after completion"))) - } - ActivityFut::Running { - inner, - data_converter, - .. - } => match Pin::new(inner).poll(cx) { - Poll::Pending => Poll::Pending, - Poll::Ready(resolution) => Poll::Ready({ - let status = resolution.status.ok_or_else(|| { - data_converter - .to_error( - &SerializationContextData::Workflow, - Failure { - message: "Activity completed without a status".to_string(), - ..Default::default() - }, - ActivityExecutionDecodeHint { cancelled: false }, - ) - .expect("synthetic activity failure should decode") - })?; - - match status { - activity_resolution::Status::Completed(success) => { - let payload = success.result.unwrap_or_default(); - let ctx = SerializationContext { - data: &SerializationContextData::Workflow, - converter: data_converter.payload_converter(), - }; + let poll = + match this { + ActivityFut::Errored { error, .. } => { + Poll::Ready(Err(*error.take().expect("polled after completion"))) + } + ActivityFut::Running { + inner, + data_converter, + .. + } => match Pin::new(inner).poll(cx) { + Poll::Pending => Poll::Pending, + Poll::Ready(resolution) => Poll::Ready({ + let status = resolution.status.ok_or_else(|| { data_converter - .payload_converter() - .from_payload::(&ctx, payload) - .map_err(ActivityExecutionError::Serialization) + .to_error( + &SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ), + Failure { + message: "Activity completed without a status".to_string(), + ..Default::default() + }, + ActivityExecutionDecodeHint::new(false), + ) + .expect("synthetic activity failure should decode") + })?; + + match status { + activity_resolution::Status::Completed(success) => { + let payload = success.result.unwrap_or_default(); + let context_data = SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ); + let ctx = SerializationContext::new( + &context_data, + data_converter.payload_converter(), + ); + data_converter + .payload_converter() + .from_payload::(&ctx, payload) + .map_err(ActivityExecutionError::Serialization) + } + activity_resolution::Status::Failed(f) => Err(data_converter + .to_error( + &SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ), + f.failure.unwrap_or_default(), + ActivityExecutionDecodeHint::new(false), + )?), + activity_resolution::Status::Cancelled(c) => Err(data_converter + .to_error( + &SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ), + c.failure.unwrap_or_default(), + ActivityExecutionDecodeHint::new(true), + )?), + activity_resolution::Status::Backoff(_) => { + panic!("DoBackoff should be handled by LATimerBackoffFut") + } } - activity_resolution::Status::Failed(f) => Err(data_converter.to_error( - &SerializationContextData::Workflow, - f.failure.unwrap_or_default(), - ActivityExecutionDecodeHint { cancelled: false }, - )?), - activity_resolution::Status::Cancelled(c) => Err(data_converter.to_error( - &SerializationContextData::Workflow, - c.failure.unwrap_or_default(), - ActivityExecutionDecodeHint { cancelled: true }, - )?), - activity_resolution::Status::Backoff(_) => { - panic!("DoBackoff should be handled by LATimerBackoffFut") - } - } - }), - }, - ActivityFut::Terminated => panic!("polled after termination"), - }; - if poll.is_ready() { + }), + }, + ActivityFut::Terminated => panic!("polled after termination"), + }; + if poll.is_ready() { *this = ActivityFut::Terminated; } poll @@ -2781,38 +3188,49 @@ where let status = result.status.ok_or_else(|| { data_converter .to_error( - &SerializationContextData::Workflow, + &SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ), Failure { message: "Child workflow completed without a status" .to_string(), ..Default::default() }, - ChildWorkflowExecutionDecodeHint, + ChildWorkflowExecutionDecodeHint::default(), ) .expect("synthetic child workflow failure should decode") })?; match status { child_workflow_result::Status::Completed(success) => { let payloads = success.result.into_iter().collect(); - let ctx = SerializationContext { - data: &SerializationContextData::Workflow, - converter: data_converter.payload_converter(), - }; + let context_data = SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ); + let ctx = SerializationContext::new( + &context_data, + data_converter.payload_converter(), + ); data_converter .payload_converter() .from_payloads::(&ctx, payloads) .map_err(ChildWorkflowExecutionError::Serialization) } - child_workflow_result::Status::Failed(f) => Err(data_converter.to_error( - &SerializationContextData::Workflow, - f.failure.unwrap_or_default(), - ChildWorkflowExecutionDecodeHint, - )?), + child_workflow_result::Status::Failed(f) => { + Err(data_converter.to_error( + &SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ), + f.failure.unwrap_or_default(), + ChildWorkflowExecutionDecodeHint::default(), + )?) + } child_workflow_result::Status::Cancelled(c) => Err(data_converter .to_error( - &SerializationContextData::Workflow, + &SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ), c.failure.unwrap_or_default(), - ChildWorkflowExecutionDecodeHint, + ChildWorkflowExecutionDecodeHint::default(), )?), } }), @@ -2896,49 +3314,63 @@ where ChildWorkflowStartFut::Errored { error, .. } => { Poll::Ready(Err(*error.take().expect("polled after completion"))) } - ChildWorkflowStartFut::Running(inner) => match Pin::new(inner).poll(cx) { - Poll::Pending => Poll::Pending, - Poll::Ready(pending) => Poll::Ready(match pending.status { - ChildWorkflowStartStatus::Succeeded(s) => { - let ChildWfCommon { - workflow_id, - child_seq, - result_future, - base_ctx, - } = pending.common; - Ok(StartChildWorkflowOutput { - run_id: s.run_id, - result_future, - workflow_id, - child_seq, - base_ctx, - }) - } - ChildWorkflowStartStatus::Failed(f) => { - let mut result_future = pending.common.result_future; - result_future.unregister_cancellation(); - Err(ChildWorkflowStartError::StartFailed { - workflow_id: f.workflow_id, - workflow_type: f.workflow_type, - cause: StartChildWorkflowExecutionFailedCause::try_from(f.cause) - .unwrap_or(StartChildWorkflowExecutionFailedCause::Unspecified), - }) - } - ChildWorkflowStartStatus::Cancelled(c) => { - let ChildWfCommon { - mut result_future, - base_ctx, - .. - } = pending.common; - result_future.unregister_cancellation(); - Err(base_ctx.data_converter().to_error( - &SerializationContextData::Workflow, - c.failure.unwrap_or_default(), - ChildWorkflowStartDecodeHint, - )?) - } - }), - }, + ChildWorkflowStartFut::Running(inner) => { + match Pin::new(inner).poll(cx) { + Poll::Pending => Poll::Pending, + Poll::Ready(pending) => Poll::Ready(match pending.status { + ChildWorkflowStartStatus::Succeeded(s) => { + let ChildWfCommon { + workflow_id, + child_seq, + result_future, + base_ctx, + } = pending.common; + Ok(StartChildWorkflowOutput { + run_id: s.run_id, + result_future, + workflow_id, + child_seq, + base_ctx, + }) + } + ChildWorkflowStartStatus::Failed(f) => { + let mut result_future = pending.common.result_future; + result_future.unregister_cancellation(); + Err(ChildWorkflowStartError::StartFailed { + workflow_id: f.workflow_id, + workflow_type: f.workflow_type, + cause: match f.cause { + cause if cause == ProtoStartChildCause::Unspecified as i32 => { + StartChildWorkflowExecutionFailedCause::Unspecified + } + cause + if cause + == ProtoStartChildCause::WorkflowAlreadyExists as i32 => + { + StartChildWorkflowExecutionFailedCause::WorkflowAlreadyExists + } + _ => StartChildWorkflowExecutionFailedCause::Unknown, + }, + }) + } + ChildWorkflowStartStatus::Cancelled(c) => { + let ChildWfCommon { + mut result_future, + base_ctx, + .. + } = pending.common; + result_future.unregister_cancellation(); + Err(base_ctx.data_converter().to_error( + &SerializationContextData::Workflow( + WorkflowSerializationContext::new(), + ), + c.failure.unwrap_or_default(), + ChildWorkflowStartDecodeHint::default(), + )?) + } + }), + } + } ChildWorkflowStartFut::Terminated => panic!("polled after termination"), }; if poll.is_ready() { @@ -3008,10 +3440,10 @@ where } => match Pin::new(inner).poll(cx) { Poll::Pending => Poll::Pending, Poll::Ready(Ok(_)) => Poll::Ready(Ok(())), - Poll::Ready(Err(failure)) => Poll::Ready(Err(data_converter.to_error( - &SerializationContextData::Workflow, - failure, - WorkflowSignalDecodeHint, + Poll::Ready(Err(error)) => Poll::Ready(Err(data_converter.to_error( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + error.failure, + WorkflowSignalDecodeHint::new(error.cause), )?)), }, SignalChildFut::Terminated => panic!("polled after termination"), @@ -3145,7 +3577,7 @@ impl ExternalWorkflowHandle { pub fn cancel( &self, reason: Option, - ) -> impl FusedFuture { + ) -> impl FusedFuture { self.base_ctx .cancel_external_workflow(CancelExternalWorkflowInput { workflow_id: self.workflow_id.clone(), @@ -3155,41 +3587,6 @@ impl ExternalWorkflowHandle { } } -#[derive(derive_more::Debug)] -#[debug("StartedNexusOperation{{ operation_token: {operation_token:?} }}")] -/// Handle to a started Nexus operation. -pub struct StartedNexusOperation { - /// The operation token, if the operation started asynchronously - pub operation_token: Option, - #[debug(skip)] - pub(crate) result_future: Shared>, - pub(crate) schedule_seq: u32, - #[debug(skip)] - pub(crate) base_ctx: BaseWorkflowContext, -} - -pub(crate) struct NexusUnblockData { - pub(crate) result_future: Shared>, - pub(crate) schedule_seq: u32, - pub(crate) base_ctx: BaseWorkflowContext, -} - -impl StartedNexusOperation { - /// Wait for the operation result. - pub async fn result(&self) -> NexusOperationResult { - // The result future is a `Shared`; poll it inside an `SdkWakeGuard` (via - // `SdkGuardedFuture`) so its internal waker machinery isn't mistaken for a non-SDK wake on - // replay (which would fail the workflow task with TMPRL1100). - SdkGuardedFuture(self.result_future.clone()).await - } - - /// Request cancellation of the operation. - pub fn cancel(&self) { - self.base_ctx - .cancel(CancellableID::NexusOp(self.schedule_seq)); - } -} - #[cfg(test)] mod tests { use super::*; @@ -3204,21 +3601,17 @@ mod tests { time::Duration, }; use temporalio_common_wasm::{ - RetryPolicy, data_converters::{TemporalDeserializable, TemporalSerializable}, error::OutgoingWorkflowError, protos::{ coresdk::{ AsJsonPayloadExt, FromJsonPayloadExt, common::VersioningIntent as ProtoVersioningIntent, - workflow_activation::{ - ResolveChildWorkflowExecutionStartSuccess, UpdateRandomSeed, - WorkflowActivationJob, resolve_nexus_operation_start, - }, + workflow_activation::{UpdateRandomSeed, WorkflowActivationJob}, workflow_commands::WorkflowCommand, }, temporal::api::{ - common::v1::{Payload, RetryPolicy as ProtoRetryPolicy}, + common::v1::Payload, enums::v1::ContinueAsNewVersioningBehavior as ProtoContinueAsNewVersioningBehavior, }, }, @@ -3292,17 +3685,6 @@ mod tests { } } - struct TestActivity; - - impl ActivityDefinition for TestActivity { - type Input = (); - type Output = (); - - fn name(&self) -> &str { - "test_activity" - } - } - fn test_context() -> WorkflowContext { test_context_with_seed(0) } @@ -3424,278 +3806,319 @@ mod tests { assert_eq!(timer.seq, 1); } - #[test] - fn custom_token_cancels_command_backed_operations() { - let host = Rc::new(RecordingHost::default()); - let init = WorkflowInit { - namespace: "default".to_string(), - task_queue: "task-queue".to_string(), - run_id: "run-id".to_string(), - initialize_workflow: InitializeWorkflow { - workflow_type: TestWorkflow.name().to_string(), - ..Default::default() + #[cfg(feature = "experimental")] + mod experimental_operation_tests { + use super::*; + use temporalio_common_wasm::protos::{ + coresdk::workflow_activation::{ + ResolveChildWorkflowExecutionStartSuccess, resolve_nexus_operation_start, }, + temporal::api::sdk::v1::{EventGroupMarker, event_group_marker}, }; - let base = BaseWorkflowContext::from_raw( - init, - DataConverter::default(), - host.clone(), - None, - Vec::new(), - ); - let token = WorkflowCancellationToken::new(); - let timer = base.timer(TimerOptions { - duration: Duration::from_secs(1), - cancellation_token: Some(token.clone()), - summary: None, - }); + struct TestActivity; - let mut activity_options = ActivityOptions::start_to_close_timeout(Duration::from_secs(1)); - activity_options.cancellation_token = Some(token.clone()); - let activity = base.execute_activity(TestActivity, (), activity_options); + impl ActivityDefinition for TestActivity { + type Input = (); + type Output = (); - let mut local_activity_options = LocalActivityOptions { - schedule_to_close_timeout: Some(Duration::from_secs(1)), - ..Default::default() - }; - local_activity_options.cancellation_token = Some(token.clone()); - let local_activity = base.execute_local_activity(TestActivity, (), local_activity_options); + fn name(&self) -> &str { + "test_activity" + } + } - let child_options = ChildWorkflowOptions { - cancellation_token: Some(token.clone()), - ..Default::default() - }; - let child = base.start_child_workflow(TestWorkflow::run, 1, child_options); + #[test] + fn custom_token_cancels_command_backed_operations() { + let host = Rc::new(RecordingHost::default()); + let init = WorkflowInit { + namespace: "default".to_string(), + task_queue: "task-queue".to_string(), + run_id: "run-id".to_string(), + initialize_workflow: InitializeWorkflow { + workflow_type: TestWorkflow.name().to_string(), + ..Default::default() + }, + }; + let base = BaseWorkflowContext::from_raw( + init, + DataConverter::default(), + host.clone(), + None, + Vec::new(), + ); + let token = WorkflowCancellationToken::new(); - let signal = base.external_workflow("external", None).signal( - TestWorkflow::test_signal, - "input".to_string(), - SignalWorkflowOptions::builder() - .cancellation_token(token.clone()) - .build(), - ); + let timer = base.timer(TimerOptions { + duration: Duration::from_secs(1), + cancellation_token: Some(token.clone()), + summary: None, + event_group_markers: vec![], + }); - let nexus_options = NexusOperationOptions::builder() - .endpoint("endpoint") - .service("service") - .operation("operation") - .cancellation_token(token.clone()) - .build(); - let nexus = base.start_nexus_operation(nexus_options); - - token.cancel_with_reason("group cancelled"); - timer.cancel(); - activity.cancel(); - local_activity.cancel(); - child.cancel_with_reason("explicit cancellation".to_string()); - signal.cancel(); - nexus.cancel(); + let mut activity_options = + ActivityOptions::start_to_close_timeout(Duration::from_secs(1)); + activity_options.cancellation_token = Some(token.clone()); + let activity = base.execute_activity(TestActivity, (), activity_options); - let commands = host.commands.borrow(); - assert_eq!( - commands - .iter() - .filter(|command| matches!( - &command.variant, - Some(workflow_command::Variant::CancelTimer(_)) - )) - .count(), - 1 - ); - assert_eq!( - commands - .iter() - .filter(|command| matches!( - &command.variant, - Some(workflow_command::Variant::RequestCancelActivity(_)) - )) - .count(), - 1 - ); - assert_eq!( - commands - .iter() - .filter(|command| matches!( - &command.variant, - Some(workflow_command::Variant::RequestCancelLocalActivity(_)) - )) - .count(), - 1 - ); - let child_cancellations = commands - .iter() - .filter_map(|command| match &command.variant { - Some(workflow_command::Variant::CancelChildWorkflowExecution(cancel)) => { - Some(cancel) - } - _ => None, - }) - .collect::>(); - assert_eq!(child_cancellations.len(), 1); - assert_eq!(child_cancellations[0].reason, "group cancelled"); - assert_eq!( - commands - .iter() - .filter(|command| matches!( - &command.variant, - Some(workflow_command::Variant::CancelSignalWorkflow(_)) - )) - .count(), - 1 - ); - assert_eq!( - commands - .iter() - .filter(|command| matches!( - &command.variant, - Some(workflow_command::Variant::RequestCancelNexusOperation(_)) - )) - .count(), - 1 - ); - } + let mut local_activity_options = LocalActivityOptions { + schedule_to_close_timeout: Some(Duration::from_secs(1)), + ..Default::default() + }; + local_activity_options.cancellation_token = Some(token.clone()); + let local_activity = + base.execute_local_activity(TestActivity, (), local_activity_options); - #[test] - fn child_and_nexus_tokens_remain_active_after_start() { - let host = Rc::new(RecordingHost::default()); - let init = WorkflowInit { - namespace: "default".to_string(), - task_queue: "task-queue".to_string(), - run_id: "run-id".to_string(), - initialize_workflow: InitializeWorkflow { - workflow_type: TestWorkflow.name().to_string(), + let child_options = ChildWorkflowOptions { + cancellation_token: Some(token.clone()), ..Default::default() - }, - }; - let base = BaseWorkflowContext::from_raw( - init, - DataConverter::default(), - host.clone(), - None, - Vec::new(), - ); + }; + let child = base.start_child_workflow(TestWorkflow::run, 1, child_options); - let child_token = WorkflowCancellationToken::new(); - let child_options = ChildWorkflowOptions { - cancellation_token: Some(child_token.clone()), - ..Default::default() - }; - let child = base.start_child_workflow(TestWorkflow::run, 1, child_options); - base.unblock(UnblockEvent::WorkflowStart( - 1, - Box::new(ChildWorkflowStartStatus::Succeeded( - ResolveChildWorkflowExecutionStartSuccess { - run_id: "child-run".to_string(), - }, - )), - )) - .unwrap(); - let started_child = child - .now_or_never() - .expect("child start should resolve") - .unwrap(); - child_token.cancel(); - started_child.cancel("explicit cancellation".to_string()); - - let nexus_token = WorkflowCancellationToken::new(); - let nexus_options = NexusOperationOptions::builder() - .endpoint("endpoint") - .service("service") - .operation("operation") - .cancellation_token(nexus_token.clone()) - .build(); - let nexus = base.start_nexus_operation(nexus_options); - base.unblock(UnblockEvent::NexusOperationStart( - 1, - Box::new(resolve_nexus_operation_start::Status::OperationToken( - "operation-token".to_string(), - )), - )) - .unwrap(); - let started_nexus = nexus - .now_or_never() - .expect("Nexus start should resolve") - .unwrap(); - nexus_token.cancel(); - started_nexus.cancel(); + let signal = base.external_workflow("external", None).signal( + TestWorkflow::test_signal, + "input".to_string(), + SignalWorkflowOptions::builder() + .cancellation_token(token.clone()) + .build(), + ); - let commands = host.commands.borrow(); - assert_eq!( - commands - .iter() - .filter(|command| matches!( - &command.variant, - Some(workflow_command::Variant::CancelChildWorkflowExecution(_)) - )) - .count(), - 1 - ); - assert_eq!( - commands + let nexus_options = NexusOperationOptions::builder() + .endpoint("endpoint") + .service("service") + .operation("operation") + .cancellation_token(token.clone()) + .build(); + let nexus = base.start_nexus_operation(nexus_options); + + token.cancel_with_reason("group cancelled"); + timer.cancel(); + activity.cancel(); + local_activity.cancel(); + child.cancel_with_reason("explicit cancellation".to_string()); + signal.cancel(); + nexus.cancel(); + + let commands = host.commands.borrow(); + assert_eq!( + commands + .iter() + .filter(|command| matches!( + &command.variant, + Some(workflow_command::Variant::CancelTimer(_)) + )) + .count(), + 1 + ); + assert_eq!( + commands + .iter() + .filter(|command| matches!( + &command.variant, + Some(workflow_command::Variant::RequestCancelActivity(_)) + )) + .count(), + 1 + ); + assert_eq!( + commands + .iter() + .filter(|command| matches!( + &command.variant, + Some(workflow_command::Variant::RequestCancelLocalActivity(_)) + )) + .count(), + 1 + ); + let child_cancellations = commands .iter() - .filter(|command| matches!( - &command.variant, - Some(workflow_command::Variant::RequestCancelNexusOperation(_)) - )) - .count(), - 1 - ); - } + .filter_map(|command| match &command.variant { + Some(workflow_command::Variant::CancelChildWorkflowExecution(cancel)) => { + Some(cancel) + } + _ => None, + }) + .collect::>(); + assert_eq!(child_cancellations.len(), 1); + assert_eq!(child_cancellations[0].reason, "group cancelled"); + assert_eq!( + commands + .iter() + .filter(|command| matches!( + &command.variant, + Some(workflow_command::Variant::CancelSignalWorkflow(_)) + )) + .count(), + 1 + ); + assert_eq!( + commands + .iter() + .filter(|command| matches!( + &command.variant, + Some(workflow_command::Variant::RequestCancelNexusOperation(_)) + )) + .count(), + 1 + ); + } - #[test] - fn local_activity_token_cancels_retry_backoff_timer() { - let host = Rc::new(RecordingHost::default()); - let init = WorkflowInit { - namespace: "default".to_string(), - task_queue: "task-queue".to_string(), - run_id: "run-id".to_string(), - initialize_workflow: InitializeWorkflow { - workflow_type: TestWorkflow.name().to_string(), + #[test] + fn child_and_nexus_tokens_remain_active_after_start() { + let host = Rc::new(RecordingHost::default()); + let init = WorkflowInit { + namespace: "default".to_string(), + task_queue: "task-queue".to_string(), + run_id: "run-id".to_string(), + initialize_workflow: InitializeWorkflow { + workflow_type: TestWorkflow.name().to_string(), + ..Default::default() + }, + }; + let base = BaseWorkflowContext::from_raw( + init, + DataConverter::default(), + host.clone(), + None, + Vec::new(), + ); + + let child_token = WorkflowCancellationToken::new(); + let child_options = ChildWorkflowOptions { + cancellation_token: Some(child_token.clone()), ..Default::default() - }, - }; - let base = BaseWorkflowContext::from_raw( - init, - DataConverter::default(), - host.clone(), - None, - Vec::new(), - ); - let token = WorkflowCancellationToken::new(); - let mut options = LocalActivityOptions { - schedule_to_close_timeout: Some(Duration::from_secs(10)), - ..Default::default() - }; - options.cancellation_token = Some(token.clone()); - let activity = base.execute_local_activity(TestActivity, (), options); - futures_util::pin_mut!(activity); - base.unblock(UnblockEvent::Activity( - 1, - Box::new(ActivityResolution { - status: Some(activity_resolution::Status::Backoff( - temporalio_common_wasm::protos::coresdk::activity_result::DoBackoff { - attempt: 2, - backoff_duration: Some(Duration::from_secs(5).try_into().unwrap()), - original_schedule_time: None, + }; + let child = base.start_child_workflow(TestWorkflow::run, 1, child_options); + base.unblock(UnblockEvent::WorkflowStart( + 1, + Box::new(ChildWorkflowStartStatus::Succeeded( + ResolveChildWorkflowExecutionStartSuccess { + run_id: "child-run".to_string(), }, )), - }), - )) - .unwrap(); + )) + .unwrap(); + let started_child = child + .now_or_never() + .expect("child start should resolve") + .unwrap(); + child_token.cancel(); + started_child.cancel("explicit cancellation".to_string()); + + let nexus_token = WorkflowCancellationToken::new(); + let nexus_options = NexusOperationOptions::builder() + .endpoint("endpoint") + .service("service") + .operation("operation") + .cancellation_token(nexus_token.clone()) + .build(); + let nexus = base.start_nexus_operation(nexus_options); + base.unblock(UnblockEvent::NexusOperationStart( + 1, + Box::new(resolve_nexus_operation_start::Status::OperationToken( + "operation-token".to_string(), + )), + )) + .unwrap(); + let started_nexus = nexus + .now_or_never() + .expect("Nexus start should resolve") + .unwrap(); + nexus_token.cancel(); + started_nexus.cancel(); + + let commands = host.commands.borrow(); + assert_eq!( + commands + .iter() + .filter(|command| matches!( + &command.variant, + Some(workflow_command::Variant::CancelChildWorkflowExecution(_)) + )) + .count(), + 1 + ); + assert_eq!( + commands + .iter() + .filter(|command| matches!( + &command.variant, + Some(workflow_command::Variant::RequestCancelNexusOperation(_)) + )) + .count(), + 1 + ); + } - assert!(activity.as_mut().now_or_never().is_none()); - token.cancel(); + #[test] + fn local_activity_token_cancels_retry_backoff_timer() { + let host = Rc::new(RecordingHost::default()); + let init = WorkflowInit { + namespace: "default".to_string(), + task_queue: "task-queue".to_string(), + run_id: "run-id".to_string(), + initialize_workflow: InitializeWorkflow { + workflow_type: TestWorkflow.name().to_string(), + ..Default::default() + }, + }; + let base = BaseWorkflowContext::from_raw( + init, + DataConverter::default(), + host.clone(), + None, + Vec::new(), + ); + let token = WorkflowCancellationToken::new(); + let marker = EventGroupMarker { + variant: Some(event_group_marker::Variant::Label( + event_group_marker::Label { + id: "la-group".to_string(), + label: Some("la-group".as_json_payload().unwrap()), + }, + )), + }; + let mut options = LocalActivityOptions { + schedule_to_close_timeout: Some(Duration::from_secs(10)), + event_group_markers: vec![marker.clone()], + ..Default::default() + }; + options.cancellation_token = Some(token.clone()); + let activity = base.execute_local_activity(TestActivity, (), options); + futures_util::pin_mut!(activity); + base.unblock(UnblockEvent::Activity( + 1, + Box::new(ActivityResolution { + status: Some(activity_resolution::Status::Backoff( + temporalio_common_wasm::protos::coresdk::activity_result::DoBackoff { + attempt: 2, + backoff_duration: Some(Duration::from_secs(5).try_into().unwrap()), + original_schedule_time: None, + }, + )), + }), + )) + .unwrap(); - let commands = host.commands.borrow(); - assert!(commands.iter().any(|command| matches!( - &command.variant, - Some(workflow_command::Variant::StartTimer(_)) - ))); - assert!(commands.iter().any(|command| matches!( - &command.variant, - Some(workflow_command::Variant::CancelTimer(_)) - ))); + assert!(activity.as_mut().now_or_never().is_none()); + token.cancel(); + + let commands = host.commands.borrow(); + assert!(commands.iter().any(|command| matches!( + &command.variant, + Some(workflow_command::Variant::CancelTimer(_)) + ))); + + let start_timer = commands + .iter() + .find(|command| { + matches!( + &command.variant, + Some(workflow_command::Variant::StartTimer(_)) + ) + }) + .expect("backoff StartTimer is issued"); + assert_eq!(start_timer.event_group_markers, [marker]); + } } #[test] @@ -3705,8 +4128,16 @@ mod tests { let callback_calls = calls.clone(); let callback_input = input.clone(); let callback: PatchActivationCallback = Arc::new(move |value| { + assert!(matches!( + value.workflow_info.random_stream("plugin").source, + WorkflowRandomStreamSource::System(_) + )); callback_calls.fetch_add(1, AtomicOrdering::Relaxed); - *callback_input.lock().unwrap() = Some(value); + *callback_input.lock().unwrap() = Some(( + value.workflow_info.workflow_id().to_string(), + value.workflow_info.run_id().to_string(), + value.patch_id, + )); true }); let (_, ctx, commands) = patch_test_context(Some(callback)); @@ -3717,9 +4148,9 @@ mod tests { assert_eq!(commands.borrow().len(), 1); let input = input.lock().unwrap(); let input = input.as_ref().unwrap(); - assert_eq!(input.workflow_info.workflow_id(), "workflow-id"); - assert_eq!(input.workflow_info.run_id(), "run-id"); - assert_eq!(input.patch_id, "my-patch"); + assert_eq!(input.0, "workflow-id"); + assert_eq!(input.1, "run-id"); + assert_eq!(input.2, "my-patch"); } #[test] @@ -3809,48 +4240,327 @@ mod tests { assert_eq!(ctx.random::(), expected); } - struct MutatingRemainingOutboundInterceptor; + #[test] + fn named_random_lookup_continues_the_same_stream() { + let ctx = test_context_with_seed(42); + let first_lookup = ctx.random_stream("orders"); + let first = first_lookup.random::(); + let second = ctx.random_stream("orders").random::(); - impl WorkflowInterceptor for MutatingRemainingOutboundInterceptor { - fn signal_workflow( - &self, - _ctx: WorkflowInterceptorContext, - mut input: SignalWorkflowInput, - next: WorkflowNext< - 'static, - SignalWorkflowInput, - CancellableWorkflowOutboundFuture, - >, - ) -> CancellableWorkflowOutboundFuture { - *input.signal_name_mut() = "mutated-signal".to_string(); - *input.input_mut::().unwrap() = "mutated-input".to_string(); - *input.target_mut() = SignalWorkflowTarget::External { - namespace: "mutated-namespace".to_string(), - workflow_id: "mutated-workflow".to_string(), - run_id: Some("mutated-run".to_string()), - }; - input - .headers_mut() - .insert("signal-header".to_string(), Payload::default()); - next.run(input) + let expected = test_context_with_seed(42).random_stream("orders"); + assert_eq!(first, expected.random::()); + assert_eq!(second, expected.random::()); + } + + #[test] + fn named_random_sequence_is_stable() { + let stream = test_context_with_seed(42).random_stream("example.com/orders"); + + // Changing seed derivation or the generator would break existing workflow replays. + assert_eq!(stream.random::(), 18_054_372_068_998_079_507); + } + + #[test] + fn named_random_streams_are_isolated() { + let ctx = test_context_with_seed(42); + let alpha = ctx.random_stream("alpha"); + let first_alpha = alpha.random::(); + let _ = ctx.random_stream("beta").random::(); + let second_alpha = alpha.random::(); + + let expected_ctx = test_context_with_seed(42); + let expected_alpha = expected_ctx.random_stream("alpha"); + assert_eq!(first_alpha, expected_alpha.random::()); + assert_eq!(second_alpha, expected_alpha.random::()); + assert_ne!( + test_context_with_seed(42) + .random_stream("alpha") + .random::(), + test_context_with_seed(42) + .random_stream("beta") + .random::() + ); + } + + #[test] + fn named_random_does_not_advance_default_randomness() { + let ctx = test_context_with_seed(42); + let first = ctx.random::(); + let _ = ctx.random_stream("plugin").random::(); + let second = ctx.random::(); + + let expected = test_context_with_seed(42); + assert_eq!(first, expected.random::()); + assert_eq!(second, expected.random::()); + } + + #[test] + fn interceptor_context_shares_named_random_stream_state() { + let ctx = test_context_with_seed(42); + let first = ctx.random_stream("plugin").random::(); + let interceptor_ctx = + crate::workflow_interceptors::WorkflowInterceptorContext::new(ctx.sync.base.clone()); + let second = interceptor_ctx.random_stream("plugin").random::(); + + let expected = test_context_with_seed(42).random_stream("plugin"); + assert_eq!(first, expected.random::()); + assert_eq!(second, expected.random::()); + } + + #[test] + fn replay_safe_context_view_shares_workflow_randomness() { + let ctx = test_context_with_seed(42); + let first = ctx.sync.base.view().random_stream("plugin").random::(); + let second = ctx.random_stream("plugin").random::(); + + let expected = test_context_with_seed(42).random_stream("plugin"); + assert_eq!(first, expected.random::()); + assert_eq!(second, expected.random::()); + } + + #[test] + fn read_only_context_view_does_not_advance_workflow_randomness() { + let ctx = test_context_with_seed(42); + let expected = test_context_with_seed(42) + .random_stream("plugin") + .random::(); + + { + let _read_only = ctx.sync.base.enter_read_only(); + let _ = ctx.sync.base.view().random_stream("plugin").random::(); } - fn cancel_external_workflow( - &self, - _ctx: WorkflowInterceptorContext, - mut input: CancelExternalWorkflowInput, - next: WorkflowNext< - 'static, - CancelExternalWorkflowInput, - WorkflowOutboundFuture, - >, - ) -> WorkflowOutboundFuture { - input.workflow_id = "mutated-cancel-workflow".to_string(); - input.run_id = Some("mutated-cancel-run".to_string()); - input.reason = Some("mutated-reason".to_string()); - next.run(input) + assert_eq!(ctx.random_stream("plugin").random::(), expected); + } + + #[test] + fn nested_read_only_scopes_restore_replay_safety() { + let ctx = test_context_with_seed(42); + assert!(ctx.sync.base.requires_replay_safety()); + + { + let _outer = ctx.sync.base.enter_read_only(); + assert!(!ctx.sync.base.requires_replay_safety()); + { + let _inner = ctx.sync.base.enter_read_only(); + assert!(!ctx.sync.base.requires_replay_safety()); + } + assert!(!ctx.sync.base.requires_replay_safety()); + } + + assert!(ctx.sync.base.requires_replay_safety()); + } + + #[test] + fn named_random_streams_are_reseeded_by_activation() { + let ctx = test_context_with_seed(123); + let stream = ctx.random_stream("orders"); + let _ = stream.random::(); + let activation = CoreWorkflowActivation { + jobs: vec![WorkflowActivationJob { + variant: Some(ActivationVariant::UpdateRandomSeed(UpdateRandomSeed { + randomness_seed: 456, + })), + }], + ..Default::default() + }; + + ctx.sync.base.apply_activation_context(&activation, false); + + let expected = test_context_with_seed(456).random_stream("orders"); + assert_eq!(stream.random::(), expected.random::()); + } + + #[cfg(feature = "experimental")] + mod experimental_interceptor_tests { + use super::*; + use crate::workflow_interceptors::StartNexusOperationInput; + + struct MutatingRemainingOutboundInterceptor; + + impl WorkflowInterceptor for MutatingRemainingOutboundInterceptor { + fn signal_workflow( + &self, + _ctx: WorkflowInterceptorContext, + mut input: SignalWorkflowInput, + next: WorkflowNext< + 'static, + SignalWorkflowInput, + CancellableWorkflowOutboundFuture, + >, + ) -> CancellableWorkflowOutboundFuture { + *input.signal_name_mut() = "mutated-signal".to_string(); + *input.input_mut::().unwrap() = "mutated-input".to_string(); + *input.target_mut() = SignalWorkflowTarget::External { + namespace: "mutated-namespace".to_string(), + workflow_id: "mutated-workflow".to_string(), + run_id: Some("mutated-run".to_string()), + }; + input + .headers_mut() + .insert("signal-header".to_string(), Payload::default()); + next.run(input) + } + + fn cancel_external_workflow( + &self, + _ctx: WorkflowInterceptorContext, + mut input: CancelExternalWorkflowInput, + next: WorkflowNext< + 'static, + CancelExternalWorkflowInput, + WorkflowOutboundFuture, + >, + ) -> WorkflowOutboundFuture { + input.workflow_id = "mutated-cancel-workflow".to_string(); + input.run_id = Some("mutated-cancel-run".to_string()); + input.reason = Some("mutated-reason".to_string()); + next.run(input) + } + + fn continue_as_new( + &self, + _ctx: crate::workflow_interceptors::SyncWorkflowInterceptorContext, + mut input: ContinueAsNewInput, + next: WorkflowNext< + 'static, + ContinueAsNewInput, + crate::workflow_interceptors::ContinueAsNewResult, + >, + ) -> crate::workflow_interceptors::ContinueAsNewResult { + *input.input_mut::().unwrap() = 42; + input.options_mut().workflow_type = Some("mutated-workflow-type".to_string()); + input.headers_mut().insert( + "continue-header".to_string(), + Payload::from(b"continue-header-value".as_slice()), + ); + next.run(input) + } + + fn start_nexus_operation( + &self, + _ctx: WorkflowInterceptorContext, + mut input: StartNexusOperationInput, + next: WorkflowNext< + 'static, + StartNexusOperationInput, + CancellableWorkflowOutboundFuture< + crate::workflow_interceptors::StartNexusOperationResult, + >, + >, + ) -> CancellableWorkflowOutboundFuture< + crate::workflow_interceptors::StartNexusOperationResult, + > { + input.options_mut().endpoint = "mutated-endpoint".to_string(); + input.options_mut().service = "mutated-service".to_string(); + input.options_mut().operation = "mutated-operation".to_string(); + next.run(input) + } + } + + #[test] + fn outbound_interceptors_mutate_signal_cancel_continue_as_new_and_nexus() { + let host = Rc::new(RecordingHost::default()); + let init = InitializeWorkflow { + workflow_type: TestWorkflow.name().to_string(), + ..Default::default() + }; + let init = WorkflowInit { + namespace: "default".to_string(), + task_queue: "task-queue".to_string(), + run_id: "run-id".to_string(), + initialize_workflow: init, + }; + let base = BaseWorkflowContext::from_raw( + init, + DataConverter::default(), + host.clone(), + None, + vec![WorkflowInterceptorConstructor::new(|_| { + MutatingRemainingOutboundInterceptor + })], + ); + let ctx = WorkflowContext::from_base(base, Rc::new(RefCell::new(TestWorkflow))); + + let signal = ctx + .external_workflow("original-workflow", Some("original-run".to_string())) + .signal( + TestWorkflow::test_signal, + "original-input".to_string(), + Default::default(), + ); + let cancel_target = + ctx.external_workflow("cancel-workflow", Some("cancel-run".to_string())); + let cancel = cancel_target.cancel(Some("original-reason".to_string())); + let termination = ctx + .continue_as_new(7, ContinueAsNewOptions::default()) + .expect_err("continue_as_new should terminate the workflow"); + let sync_ctx = ctx.sync_context(); + let nexus = sync_ctx.start_nexus_operation( + NexusOperationOptions::builder() + .endpoint("original-endpoint") + .service("original-service") + .operation("original-operation") + .build(), + ); + drop((signal, cancel, nexus)); + + let WorkflowTermination::ContinueAsNew(continue_as_new) = termination else { + panic!("expected continue-as-new termination") + }; + assert_eq!(continue_as_new.workflow_type, "mutated-workflow-type"); + assert_eq!( + continue_as_new.arguments, + vec![42u8.as_json_payload().unwrap()] + ); + assert!(continue_as_new.headers.contains_key("continue-header")); + + let commands = host.commands.borrow(); + assert_eq!(commands.len(), 3); + let Some(workflow_command::Variant::SignalExternalWorkflowExecution(signal)) = + &commands[0].variant + else { + panic!("expected signal command") + }; + assert_eq!(signal.signal_name, "mutated-signal"); + assert_eq!( + signal.args, + vec!["mutated-input".to_string().as_json_payload().unwrap()] + ); + assert!(signal.headers.contains_key("signal-header")); + let Some(signal_external_workflow_execution::Target::WorkflowExecution(target)) = + &signal.target + else { + panic!("expected external workflow signal target") + }; + assert_eq!(target.namespace, "mutated-namespace"); + assert_eq!(target.workflow_id, "mutated-workflow"); + assert_eq!(target.run_id, "mutated-run"); + + let Some(workflow_command::Variant::RequestCancelExternalWorkflowExecution(cancel)) = + &commands[1].variant + else { + panic!("expected external cancellation command") + }; + let target = cancel.workflow_execution.as_ref().unwrap(); + assert_eq!(target.workflow_id, "mutated-cancel-workflow"); + assert_eq!(target.run_id, "mutated-cancel-run"); + assert_eq!(cancel.reason, "mutated-reason"); + + let Some(workflow_command::Variant::ScheduleNexusOperation(nexus)) = + &commands[2].variant + else { + panic!("expected Nexus operation command") + }; + assert_eq!(nexus.endpoint, "mutated-endpoint"); + assert_eq!(nexus.service, "mutated-service"); + assert_eq!(nexus.operation, "mutated-operation"); } + } + + struct HeaderAddingContinueAsNewInterceptor; + impl WorkflowInterceptor for HeaderAddingContinueAsNewInterceptor { fn continue_as_new( &self, _ctx: crate::workflow_interceptors::SyncWorkflowInterceptorContext, @@ -3861,132 +4571,12 @@ mod tests { crate::workflow_interceptors::ContinueAsNewResult, >, ) -> crate::workflow_interceptors::ContinueAsNewResult { - *input.input_mut::().unwrap() = 42; - input.options_mut().workflow_type = Some("mutated-workflow-type".to_string()); input.headers_mut().insert( "continue-header".to_string(), Payload::from(b"continue-header-value".as_slice()), ); next.run(input) } - - fn start_nexus_operation( - &self, - _ctx: WorkflowInterceptorContext, - mut input: StartNexusOperationInput, - next: WorkflowNext< - 'static, - StartNexusOperationInput, - CancellableWorkflowOutboundFuture< - crate::workflow_interceptors::StartNexusOperationResult, - >, - >, - ) -> CancellableWorkflowOutboundFuture< - crate::workflow_interceptors::StartNexusOperationResult, - > { - input.options_mut().endpoint = "mutated-endpoint".to_string(); - input.options_mut().service = "mutated-service".to_string(); - input.options_mut().operation = "mutated-operation".to_string(); - next.run(input) - } - } - - #[test] - fn outbound_interceptors_mutate_signal_cancel_continue_as_new_and_nexus() { - let host = Rc::new(RecordingHost::default()); - let init = InitializeWorkflow { - workflow_type: TestWorkflow.name().to_string(), - ..Default::default() - }; - let init = WorkflowInit { - namespace: "default".to_string(), - task_queue: "task-queue".to_string(), - run_id: "run-id".to_string(), - initialize_workflow: init, - }; - let base = BaseWorkflowContext::from_raw( - init, - DataConverter::default(), - host.clone(), - None, - vec![WorkflowInterceptorConstructor::new(|_| { - MutatingRemainingOutboundInterceptor - })], - ); - let ctx = WorkflowContext::from_base(base, Rc::new(RefCell::new(TestWorkflow))); - - let signal = ctx - .external_workflow("original-workflow", Some("original-run".to_string())) - .signal( - TestWorkflow::test_signal, - "original-input".to_string(), - Default::default(), - ); - let cancel_target = - ctx.external_workflow("cancel-workflow", Some("cancel-run".to_string())); - let cancel = cancel_target.cancel(Some("original-reason".to_string())); - let termination = ctx - .continue_as_new(7, ContinueAsNewOptions::default()) - .expect_err("continue_as_new should terminate the workflow"); - let sync_ctx = ctx.sync_context(); - let nexus = sync_ctx.start_nexus_operation( - NexusOperationOptions::builder() - .endpoint("original-endpoint") - .service("original-service") - .operation("original-operation") - .build(), - ); - drop((signal, cancel, nexus)); - - let WorkflowTermination::ContinueAsNew(continue_as_new) = termination else { - panic!("expected continue-as-new termination") - }; - assert_eq!(continue_as_new.workflow_type, "mutated-workflow-type"); - assert_eq!( - continue_as_new.arguments, - vec![42u8.as_json_payload().unwrap()] - ); - assert!(continue_as_new.headers.contains_key("continue-header")); - - let commands = host.commands.borrow(); - assert_eq!(commands.len(), 3); - let Some(workflow_command::Variant::SignalExternalWorkflowExecution(signal)) = - &commands[0].variant - else { - panic!("expected signal command") - }; - assert_eq!(signal.signal_name, "mutated-signal"); - assert_eq!( - signal.args, - vec!["mutated-input".to_string().as_json_payload().unwrap()] - ); - assert!(signal.headers.contains_key("signal-header")); - let Some(signal_external_workflow_execution::Target::WorkflowExecution(target)) = - &signal.target - else { - panic!("expected external workflow signal target") - }; - assert_eq!(target.namespace, "mutated-namespace"); - assert_eq!(target.workflow_id, "mutated-workflow"); - assert_eq!(target.run_id, "mutated-run"); - - let Some(workflow_command::Variant::RequestCancelExternalWorkflowExecution(cancel)) = - &commands[1].variant - else { - panic!("expected external cancellation command") - }; - let target = cancel.workflow_execution.as_ref().unwrap(); - assert_eq!(target.workflow_id, "mutated-cancel-workflow"); - assert_eq!(target.run_id, "mutated-cancel-run"); - assert_eq!(cancel.reason, "mutated-reason"); - - let Some(workflow_command::Variant::ScheduleNexusOperation(nexus)) = &commands[2].variant - else { - panic!("expected Nexus operation command") - }; - assert_eq!(nexus.endpoint, "mutated-endpoint"); - assert_eq!(nexus.service, "mutated-service"); - assert_eq!(nexus.operation, "mutated-operation"); } #[test] @@ -4007,7 +4597,7 @@ mod tests { Rc::new(NoopHost), None, vec![WorkflowInterceptorConstructor::new(|_| { - MutatingRemainingOutboundInterceptor + HeaderAddingContinueAsNewInterceptor })], ); let ctx = WorkflowContext::from_base(base, Rc::new(RefCell::new(TestWorkflow))); @@ -4073,72 +4663,105 @@ mod tests { ); } - #[test] - fn sync_workflow_context_continue_as_new_applies_options() { - let ctx = test_context(); - let sync = ctx.sync_context(); - let mut memo = MemoValues::new(); - memo.insert("memo-key", "memo-value".to_string()); - let mut proto_search_attributes = ProtoSearchAttributes::default(); - proto_search_attributes.indexed_fields.insert( - "CustomKeywordField".to_string(), - Payload::from(b"value".as_slice()), - ); - let search_attributes = SearchAttributes::from_proto(&proto_search_attributes); - - let termination = sync - .continue_as_new( - 11, - ContinueAsNewOptions { - workflow_type: Some("next-workflow".to_string()), - task_queue: Some("next-task-queue".to_string()), - run_timeout: Some(Duration::from_secs(10)), - task_timeout: Some(Duration::from_secs(3)), - backoff_start_interval: Some(Duration::from_secs(4)), - memo: Some(memo.clone()), - search_attributes: Some(search_attributes.clone()), - retry_policy: Some(RetryPolicy::builder().maximum_attempts(5).build()), - versioning_intent: Some(ProtoVersioningIntent::Compatible.into()), - initial_versioning_behavior: Some( - ContinueAsNewVersioningBehavior::UseRampingVersion, - ), - }, - ) - .expect_err("continue_as_new should terminate the workflow"); - assert!( - matches!(termination, WorkflowTermination::ContinueAsNew(_)), - "expected continue-as-new termination, got {termination:?}" - ); - let WorkflowTermination::ContinueAsNew(cmd) = termination else { - unreachable!() + #[cfg(feature = "experimental")] + mod experimental_continue_as_new_tests { + use super::*; + use temporalio_common_wasm::{ + RetryPolicy, protos::temporal::api::common::v1::RetryPolicy as ProtoRetryPolicy, }; - assert_eq!( - *cmd, - crate::runtime::types::ContinueAsNewRequest { - workflow_type: "next-workflow".to_string(), - task_queue: "next-task-queue".to_string(), - arguments: vec![11u8.as_json_payload().unwrap()], - workflow_run_timeout: Some(Duration::from_secs(10).try_into().unwrap()), - workflow_task_timeout: Some(Duration::from_secs(3).try_into().unwrap()), - backoff_start_interval: Some(Duration::from_secs(4).try_into().unwrap()), - memo: HashMap::from([( - "memo-key".to_string(), - "memo-value".as_json_payload().unwrap(), - )]), - headers: HashMap::new(), - search_attributes: Some(proto_search_attributes), - retry_policy: Some(ProtoRetryPolicy { - initial_interval: Some(Duration::from_secs(1).try_into().unwrap()), - backoff_coefficient: 2.0, - maximum_attempts: 5, - ..Default::default() - }), - versioning_intent: ProtoVersioningIntent::Compatible.into(), - initial_versioning_behavior: ProtoContinueAsNewVersioningBehavior::UseRampingVersion - as i32, - } - ); + #[test] + fn sync_workflow_context_continue_as_new_applies_options() { + let ctx = test_context(); + let sync = ctx.sync_context(); + let mut memo = MemoValues::new(); + memo.insert("memo-key", "memo-value".to_string()); + let mut proto_search_attributes = ProtoSearchAttributes::default(); + proto_search_attributes.indexed_fields.insert( + "CustomKeywordField".to_string(), + Payload::from(b"value".as_slice()), + ); + let search_attributes = SearchAttributes::from_proto(&proto_search_attributes); + + let termination = sync + .continue_as_new( + 11, + ContinueAsNewOptions { + workflow_type: Some("next-workflow".to_string()), + task_queue: Some("next-task-queue".to_string()), + run_timeout: Some(Duration::from_secs(10)), + task_timeout: Some(Duration::from_secs(3)), + backoff_start_interval: Some(Duration::from_secs(4)), + memo: Some(memo.clone()), + search_attributes: Some(search_attributes.clone()), + retry_policy: Some(RetryPolicy::builder().maximum_attempts(5).build()), + versioning_intent: Some(ProtoVersioningIntent::Compatible.into()), + initial_versioning_behavior: Some( + ContinueAsNewVersioningBehavior::UseRampingVersion, + ), + }, + ) + .expect_err("continue_as_new should terminate the workflow"); + assert!( + matches!(termination, WorkflowTermination::ContinueAsNew(_)), + "expected continue-as-new termination, got {termination:?}" + ); + let WorkflowTermination::ContinueAsNew(cmd) = termination else { + unreachable!() + }; + + assert_eq!( + *cmd, + crate::runtime::types::ContinueAsNewRequest { + workflow_type: "next-workflow".to_string(), + task_queue: "next-task-queue".to_string(), + arguments: vec![11u8.as_json_payload().unwrap()], + workflow_run_timeout: Some(Duration::from_secs(10).try_into().unwrap()), + workflow_task_timeout: Some(Duration::from_secs(3).try_into().unwrap()), + backoff_start_interval: Some(Duration::from_secs(4).try_into().unwrap()), + memo: HashMap::from([( + "memo-key".to_string(), + "memo-value".as_json_payload().unwrap(), + )]), + headers: HashMap::new(), + search_attributes: Some(proto_search_attributes), + retry_policy: Some(ProtoRetryPolicy { + initial_interval: Some(Duration::from_secs(1).try_into().unwrap()), + backoff_coefficient: 2.0, + maximum_attempts: 5, + ..Default::default() + }), + versioning_intent: ProtoVersioningIntent::Compatible.into(), + initial_versioning_behavior: + ProtoContinueAsNewVersioningBehavior::UseRampingVersion as i32, + } + ); + } + + #[test] + fn workflow_context_continue_as_new_applies_auto_upgrade_versioning_behavior() { + let ctx = test_context(); + + let termination = ctx + .continue_as_new( + 13, + ContinueAsNewOptions { + initial_versioning_behavior: Some( + ContinueAsNewVersioningBehavior::AutoUpgrade, + ), + ..Default::default() + }, + ) + .expect_err("continue_as_new should terminate the workflow"); + let WorkflowTermination::ContinueAsNew(cmd) = termination else { + unreachable!() + }; + + assert_eq!( + cmd.initial_versioning_behavior, + ProtoContinueAsNewVersioningBehavior::AutoUpgrade as i32 + ); + } } #[test] @@ -4165,29 +4788,6 @@ mod tests { ); } - #[test] - fn workflow_context_continue_as_new_applies_auto_upgrade_versioning_behavior() { - let ctx = test_context(); - - let termination = ctx - .continue_as_new( - 13, - ContinueAsNewOptions { - initial_versioning_behavior: Some(ContinueAsNewVersioningBehavior::AutoUpgrade), - ..Default::default() - }, - ) - .expect_err("continue_as_new should terminate the workflow"); - let WorkflowTermination::ContinueAsNew(cmd) = termination else { - unreachable!() - }; - - assert_eq!( - cmd.initial_versioning_behavior, - ProtoContinueAsNewVersioningBehavior::AutoUpgrade as i32 - ); - } - #[test] fn continue_as_new_preserves_input_serialization_errors() { #[derive(Debug)] @@ -4357,7 +4957,15 @@ mod tests { }; let fields = &command.upserted_memo.as_ref().unwrap().fields; let payload_converter = PayloadConverter::default(); - let removal_payload = MemoValue::new(()).to_payload(&payload_converter).unwrap(); + let removal_payload = payload_converter + .to_payload( + &SerializationContext::new( + &SerializationContextData::Workflow(WorkflowSerializationContext::new()), + &payload_converter, + ), + &MemoValue::new(()), + ) + .unwrap(); assert_eq!(fields.get("old"), Some(&removal_payload)); assert_eq!( u32::from_json_payload(fields.get("new").unwrap()).unwrap(), @@ -4527,4 +5135,124 @@ mod tests { assert_eq!(info.raw(), &expected); assert_eq!(info.into_raw(), expected); } + + #[test] + fn async_context_values_survive_suspension_and_isolate_concurrent_branches() { + struct Label; + + impl WorkflowContextKey for Label { + type Value = &'static str; + } + + let ctx = test_context(); + let first_poll = Rc::new(Cell::new(true)); + let second_poll = Rc::new(Cell::new(true)); + let first_ctx = ctx.clone(); + let first_poll_in_future = first_poll.clone(); + let first = ctx.with_context_value::( + "first", + future::poll_fn(move |_| { + assert_eq!( + first_ctx.context_value::