diff --git a/temporal-sdk/src/main/java/io/temporal/internal/replay/ReplayWorkflowRunTaskHandler.java b/temporal-sdk/src/main/java/io/temporal/internal/replay/ReplayWorkflowRunTaskHandler.java index 1d0d09a692..3b0894d5d9 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/replay/ReplayWorkflowRunTaskHandler.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/replay/ReplayWorkflowRunTaskHandler.java @@ -175,7 +175,9 @@ public WorkflowTaskResult handleWorkflowTask( .setForceWorkflowTask( localActivityTaskCount > 0 && !context.isWorkflowMethodCompleted()) .setNonfirstLocalActivityAttempts(localActivityMeteringHelper.getNonfirstAttempts()) - .setSdkFlags(newSdkFlags); + .setSdkFlags(newSdkFlags) + .setParentWorkflowExecution(context.getParentWorkflowExecution()) + .setContinuedAsNew(context.getContinuedExecutionRunId().isPresent()); if (workflowStateMachines.sdkNameToWrite() != null) { result.setWriteSdkName(workflowStateMachines.sdkNameToWrite()); } diff --git a/temporal-sdk/src/main/java/io/temporal/internal/replay/ReplayWorkflowTaskHandler.java b/temporal-sdk/src/main/java/io/temporal/internal/replay/ReplayWorkflowTaskHandler.java index f5b7cb0d29..a4800a7f7f 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/replay/ReplayWorkflowTaskHandler.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/replay/ReplayWorkflowTaskHandler.java @@ -23,6 +23,7 @@ import io.temporal.common.converter.DataConverter; import io.temporal.internal.common.ProtobufTimeUtils; import io.temporal.internal.common.WorkflowExecutionUtils; +import io.temporal.internal.payload.storage.ExternalStorageRunner; import io.temporal.internal.worker.*; import io.temporal.payload.context.WorkflowSerializationContext; import io.temporal.serviceclient.MetricsTag; @@ -34,6 +35,7 @@ import java.time.Duration; import java.util.List; import java.util.Objects; +import java.util.concurrent.CancellationException; import java.util.concurrent.atomic.AtomicBoolean; import java.util.stream.Collectors; import org.slf4j.Logger; @@ -89,12 +91,19 @@ private Result handleWorkflowTaskWithQuery( boolean useCache = stickyTaskQueue != null; try { + workflowTask = retrieveStoredPayloads(workflowTask); workflowRunTaskHandler = getOrCreateWorkflowExecutor(useCache, workflowTask, metricsScope, createdNew); logWorkflowTaskToBeProcessed(workflowTask, createdNew); ServiceWorkflowHistoryIterator historyIterator = - new ServiceWorkflowHistoryIterator(service, namespace, workflowTask, metricsScope); + new ServiceWorkflowHistoryIterator( + service, + namespace, + workflowTask, + metricsScope, + options.getExternalStorageRunner(), + options.getStorageCancellation()); boolean finalCommand; Result result; @@ -132,7 +141,7 @@ private Result handleWorkflowTaskWithQuery( } return result; - } catch (InterruptedException e) { + } catch (InterruptedException | CancellationException e) { throw e; } catch (Throwable e) { // Note here that the executor might not be in the cache, even when the caching is on. In that @@ -170,6 +179,18 @@ private Result handleWorkflowTaskWithQuery( } } + private PollWorkflowTaskQueueResponse.Builder retrieveStoredPayloads( + PollWorkflowTaskQueueResponse.Builder workflowTask) { + ExternalStorageRunner externalStorageRunner = options.getExternalStorageRunner(); + if (externalStorageRunner == null) { + ExternalStorageRunner.throwIfContainsReference(workflowTask.build()); + return workflowTask; + } + return externalStorageRunner + .retrieve(workflowTask.build(), options.getStorageCancellation()) + .toBuilder(); + } + private Result createCompletedWFTRequest( String workflowType, PollWorkflowTaskQueueResponseOrBuilder workflowTask, @@ -253,7 +274,8 @@ private Result createCompletedWFTRequest( null, result.isFinalCommand(), eventIdSetHandle, - result.getApplyPostCompletionMetrics()); + result.getApplyPostCompletionMetrics(), + result.isContinuedAsNew() ? null : result.getParentWorkflowExecution()); } private Result failureToWFTResult( @@ -395,6 +417,13 @@ private WorkflowRunTaskHandler createStatefulHandler( .blockingStub() .withOption(METRICS_TAGS_CALL_OPTIONS_KEY, metricsScope) .getWorkflowExecutionHistory(getHistoryRequest); + ExternalStorageRunner externalStorageRunner = options.getExternalStorageRunner(); + if (externalStorageRunner == null) { + ExternalStorageRunner.throwIfContainsReference(getHistoryResponse); + } else { + getHistoryResponse = + externalStorageRunner.retrieve(getHistoryResponse, options.getStorageCancellation()); + } workflowTask .setHistory(getHistoryResponse.getHistory()) .setNextPageToken(getHistoryResponse.getNextPageToken()); diff --git a/temporal-sdk/src/main/java/io/temporal/internal/replay/ServiceWorkflowHistoryIterator.java b/temporal-sdk/src/main/java/io/temporal/internal/replay/ServiceWorkflowHistoryIterator.java index 229b66186e..8c9974a5ef 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/replay/ServiceWorkflowHistoryIterator.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/replay/ServiceWorkflowHistoryIterator.java @@ -12,12 +12,16 @@ import io.temporal.api.workflowservice.v1.GetWorkflowExecutionHistoryRequest; import io.temporal.api.workflowservice.v1.GetWorkflowExecutionHistoryResponse; import io.temporal.api.workflowservice.v1.PollWorkflowTaskQueueResponseOrBuilder; +import io.temporal.common.CancellationToken; +import io.temporal.internal.payload.storage.ExternalStorageRunner; import io.temporal.internal.retryer.GrpcRetryer; import io.temporal.serviceclient.RpcRetryOptions; import io.temporal.serviceclient.WorkflowServiceStubs; import java.time.Duration; import java.util.Iterator; import java.util.NoSuchElementException; +import java.util.concurrent.CancellationException; +import javax.annotation.Nullable; /** Supports iteration over history while loading new pages through calls to the service. */ class ServiceWorkflowHistoryIterator implements WorkflowHistoryIterator { @@ -29,6 +33,8 @@ class ServiceWorkflowHistoryIterator implements WorkflowHistoryIterator { private final Scope metricsScope; private final PollWorkflowTaskQueueResponseOrBuilder task; private final GrpcRetryer grpcRetryer; + private final @Nullable ExternalStorageRunner externalStorageRunner; + private final CancellationToken storageCancellation; private Deadline deadline; private Iterator current; ByteString nextPageToken; @@ -38,10 +44,22 @@ class ServiceWorkflowHistoryIterator implements WorkflowHistoryIterator { String namespace, PollWorkflowTaskQueueResponseOrBuilder task, Scope metricsScope) { + this(service, namespace, task, metricsScope, null, CancellationToken.none()); + } + + ServiceWorkflowHistoryIterator( + WorkflowServiceStubs service, + String namespace, + PollWorkflowTaskQueueResponseOrBuilder task, + Scope metricsScope, + @Nullable ExternalStorageRunner externalStorageRunner, + CancellationToken storageCancellation) { + this.storageCancellation = storageCancellation; this.service = service; this.namespace = namespace; this.task = task; this.metricsScope = metricsScope; + this.externalStorageRunner = externalStorageRunner; // TODO Refactor WorkflowHistoryIteratorTest or WorkflowHistoryIterator to remove this check. // `service == null` shouldn't be allowed as it's needed for a normal functioning of this // class. @@ -64,7 +82,13 @@ public boolean hasNext() { // true. GetWorkflowExecutionHistoryResponse response = queryWorkflowExecutionHistory(); - current = response.getHistory().getEventsList().iterator(); + History history = response.getHistory(); + if (externalStorageRunner == null) { + ExternalStorageRunner.throwIfContainsReference(history); + } else { + history = externalStorageRunner.retrieve(history, storageCancellation); + } + current = history.getEventsList().iterator(); nextPageToken = response.getNextPageToken(); // Server can return an empty page, but a valid nextPageToken that contains // more events. diff --git a/temporal-sdk/src/main/java/io/temporal/internal/replay/WorkflowTaskResult.java b/temporal-sdk/src/main/java/io/temporal/internal/replay/WorkflowTaskResult.java index 9433b99f93..db740c8d33 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/replay/WorkflowTaskResult.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/replay/WorkflowTaskResult.java @@ -1,12 +1,14 @@ package io.temporal.internal.replay; import io.temporal.api.command.v1.Command; +import io.temporal.api.common.v1.WorkflowExecution; import io.temporal.api.protocol.v1.Message; import io.temporal.api.query.v1.WorkflowQueryResult; import io.temporal.common.VersioningBehavior; import java.util.Collections; import java.util.List; import java.util.Map; +import javax.annotation.Nullable; public final class WorkflowTaskResult { @@ -26,6 +28,8 @@ public static final class Builder { private String writeSdkVersion; private VersioningBehavior versioningBehavior; private Runnable applyPostCompletionMetrics; + private @Nullable WorkflowExecution parentWorkflowExecution; + private boolean continuedAsNew; public Builder setCommands(List commands) { this.commands = commands; @@ -77,6 +81,16 @@ public Builder setVersioningBehavior(VersioningBehavior versioningBehavior) { return this; } + public Builder setParentWorkflowExecution(@Nullable WorkflowExecution parentWorkflowExecution) { + this.parentWorkflowExecution = parentWorkflowExecution; + return this; + } + + public Builder setContinuedAsNew(boolean continuedAsNew) { + this.continuedAsNew = continuedAsNew; + return this; + } + public Builder setApplyPostCompletionMetrics(Runnable applyPostCompletionMetrics) { this.applyPostCompletionMetrics = applyPostCompletionMetrics; return this; @@ -94,7 +108,9 @@ public WorkflowTaskResult build() { writeSdkName, writeSdkVersion, versioningBehavior == null ? VersioningBehavior.UNSPECIFIED : versioningBehavior, - applyPostCompletionMetrics); + applyPostCompletionMetrics, + parentWorkflowExecution, + continuedAsNew); } } @@ -109,6 +125,8 @@ public WorkflowTaskResult build() { private final String writeSdkVersion; private final VersioningBehavior versioningBehavior; private final Runnable applyPostCompletionMetrics; + private final @Nullable WorkflowExecution parentWorkflowExecution; + private final boolean continuedAsNew; private WorkflowTaskResult( List commands, @@ -121,7 +139,9 @@ private WorkflowTaskResult( String writeSdkName, String writeSdkVersion, VersioningBehavior versioningBehavior, - Runnable applyPostCompletionMetrics) { + Runnable applyPostCompletionMetrics, + @Nullable WorkflowExecution parentWorkflowExecution, + boolean continuedAsNew) { this.commands = commands; this.messages = messages; this.nonfirstLocalActivityAttempts = nonfirstLocalActivityAttempts; @@ -136,6 +156,19 @@ private WorkflowTaskResult( this.writeSdkVersion = writeSdkVersion; this.versioningBehavior = versioningBehavior; this.applyPostCompletionMetrics = applyPostCompletionMetrics; + this.parentWorkflowExecution = parentWorkflowExecution; + this.continuedAsNew = continuedAsNew; + } + + /** The workflow that started this one as a child, or {@code null} if it has no parent. */ + @Nullable + public WorkflowExecution getParentWorkflowExecution() { + return parentWorkflowExecution; + } + + /** Whether this run was created by a continue-as-new rather than started directly. */ + public boolean isContinuedAsNew() { + return continuedAsNew; } public List getCommands() { diff --git a/temporal-sdk/src/main/java/io/temporal/internal/worker/BasePoller.java b/temporal-sdk/src/main/java/io/temporal/internal/worker/BasePoller.java index a8a77d680f..5c91e79619 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/worker/BasePoller.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/worker/BasePoller.java @@ -9,6 +9,7 @@ import java.time.Duration; import java.util.Objects; import java.util.concurrent.*; +import java.util.concurrent.CancellationException; import java.util.concurrent.atomic.AtomicReference; import org.slf4j.Logger; import org.slf4j.LoggerFactory; @@ -176,6 +177,8 @@ static boolean shouldIgnoreDuringShutdown(Throwable ex) { ex instanceof RejectedExecutionException // if the worker thread gets InterruptedException - it's normal during shutdown || ex instanceof InterruptedException + || ex instanceof CancellationException + || ex.getCause() instanceof CancellationException // if we get wrapped InterruptedException like what PollTask or GRPC clients do with // setting Thread.interrupted() on - it's normal during shutdown too. See PollTask // javadoc. diff --git a/temporal-sdk/src/main/java/io/temporal/internal/worker/SingleWorkerOptions.java b/temporal-sdk/src/main/java/io/temporal/internal/worker/SingleWorkerOptions.java index a34e55d904..aff380e3ef 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/worker/SingleWorkerOptions.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/worker/SingleWorkerOptions.java @@ -3,6 +3,7 @@ import com.uber.m3.tally.NoopScope; import com.uber.m3.tally.Scope; import io.temporal.api.common.v1.WorkerVersionStamp; +import io.temporal.common.CancellationToken; import io.temporal.common.context.ContextPropagator; import io.temporal.common.converter.DataConverter; import io.temporal.common.converter.GlobalDataConverter; @@ -12,6 +13,7 @@ import io.temporal.worker.WorkerDeploymentOptions; import java.time.Duration; import java.util.List; +import java.util.concurrent.CancellationException; import javax.annotation.Nullable; public final class SingleWorkerOptions { @@ -48,6 +50,7 @@ public static final class Builder { private String workerControlTaskQueue; private PreferredVersionProvider preferredVersionProvider; private @Nullable ExternalStorageRunner externalStorageRunner; + private CancellationToken storageCancellation = CancellationToken.none(); private Builder() {} @@ -77,6 +80,7 @@ private Builder(SingleWorkerOptions options) { this.workerControlTaskQueue = options.getWorkerControlTaskQueue(); this.preferredVersionProvider = options.getPreferredVersionProvider(); this.externalStorageRunner = options.getExternalStorageRunner(); + this.storageCancellation = options.getStorageCancellation(); } public Builder setIdentity(String identity) { @@ -189,6 +193,13 @@ public Builder setPreferredVersionProvider(PreferredVersionProvider preferredVer return this; } + /** Cancelled when this worker stops, to abandon its in-flight external storage work. */ + public Builder setStorageCancellation( + CancellationToken storageCancellation) { + this.storageCancellation = storageCancellation; + return this; + } + public Builder setExternalStorageRunner(@Nullable ExternalStorageRunner externalStorageRunner) { this.externalStorageRunner = externalStorageRunner; return this; @@ -237,7 +248,8 @@ public SingleWorkerOptions build() { this.allowActivityHeartbeatDuringShutdown, this.workerControlTaskQueue, this.preferredVersionProvider, - this.externalStorageRunner); + this.externalStorageRunner, + this.storageCancellation); } } @@ -263,6 +275,7 @@ public SingleWorkerOptions build() { private final String workerControlTaskQueue; private final PreferredVersionProvider preferredVersionProvider; private final @Nullable ExternalStorageRunner externalStorageRunner; + private final CancellationToken storageCancellation; private SingleWorkerOptions( String identity, @@ -286,7 +299,8 @@ private SingleWorkerOptions( boolean allowActivityHeartbeatDuringShutdown, String workerControlTaskQueue, PreferredVersionProvider preferredVersionProvider, - @Nullable ExternalStorageRunner externalStorageRunner) { + @Nullable ExternalStorageRunner externalStorageRunner, + CancellationToken storageCancellation) { this.identity = identity; this.binaryChecksum = binaryChecksum; this.buildId = buildId; @@ -309,6 +323,7 @@ private SingleWorkerOptions( this.workerControlTaskQueue = workerControlTaskQueue; this.preferredVersionProvider = preferredVersionProvider; this.externalStorageRunner = externalStorageRunner; + this.storageCancellation = storageCancellation; } public String getIdentity() { @@ -411,6 +426,10 @@ public ExternalStorageRunner getExternalStorageRunner() { return externalStorageRunner; } + public CancellationToken getStorageCancellation() { + return storageCancellation; + } + public WorkerVersioningOptions getWorkerVersioningOptions() { return new WorkerVersioningOptions( this.getBuildId(), this.isUsingBuildIdForVersioning(), this.getDeploymentOptions()); diff --git a/temporal-sdk/src/main/java/io/temporal/internal/worker/SyncWorkflowWorker.java b/temporal-sdk/src/main/java/io/temporal/internal/worker/SyncWorkflowWorker.java index be128a5e62..c86ff4ad77 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/worker/SyncWorkflowWorker.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/worker/SyncWorkflowWorker.java @@ -10,6 +10,7 @@ import io.temporal.internal.activity.ActivityExecutionContextFactory; import io.temporal.internal.activity.ActivityTaskHandlerImpl; import io.temporal.internal.activity.LocalActivityExecutionContextFactoryImpl; +import io.temporal.internal.concurrent.structured.CancelSource; import io.temporal.internal.replay.ReplayWorkflowTaskHandler; import io.temporal.internal.sync.POJOWorkflowImplementationFactory; import io.temporal.internal.sync.WorkflowThreadExecutor; @@ -54,6 +55,8 @@ public class SyncWorkflowWorker implements SuspendableWorker { private final POJOWorkflowImplementationFactory factory; private final DataConverter dataConverter; private final ActivityTaskHandlerImpl laTaskHandler; + private final CancelSource storageCancellation = + new CancelSource<>(() -> new CancellationException("Worker shutdown")); private boolean runningLocalActivityWorker; public SyncWorkflowWorker( @@ -71,6 +74,10 @@ public SyncWorkflowWorker( @Nonnull SlotSupplier slotSupplier, @Nonnull SlotSupplier laSlotSupplier, @Nonnull NamespaceCapabilities namespaceCapabilities) { + singleWorkerOptions = + SingleWorkerOptions.newBuilder(singleWorkerOptions) + .setStorageCancellation(storageCancellation.token()) + .build(); this.identity = singleWorkerOptions.getIdentity(); this.namespace = namespace; this.taskQueue = taskQueue; @@ -175,8 +182,12 @@ public boolean start() { @Override public CompletableFuture shutdown(ShutdownManager shutdownManager, boolean interruptTasks) { - return workflowWorker - .shutdown(shutdownManager, interruptTasks) + CompletableFuture workflowWorkerShutdown = + workflowWorker.shutdown(shutdownManager, interruptTasks); + if (interruptTasks) { + storageCancellation.cancel(); + } + return workflowWorkerShutdown .thenCompose(ignore -> laWorker.shutdown(shutdownManager, interruptTasks)) .exceptionally( e -> { diff --git a/temporal-sdk/src/main/java/io/temporal/internal/worker/WorkflowTaskHandler.java b/temporal-sdk/src/main/java/io/temporal/internal/worker/WorkflowTaskHandler.java index 129847fffc..3f4a95959c 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/worker/WorkflowTaskHandler.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/worker/WorkflowTaskHandler.java @@ -1,11 +1,13 @@ package io.temporal.internal.worker; +import io.temporal.api.common.v1.WorkflowExecution; import io.temporal.api.workflowservice.v1.PollWorkflowTaskQueueResponse; import io.temporal.api.workflowservice.v1.RespondQueryTaskCompletedRequest; import io.temporal.api.workflowservice.v1.RespondWorkflowTaskCompletedRequest; import io.temporal.api.workflowservice.v1.RespondWorkflowTaskFailedRequest; import io.temporal.serviceclient.RpcRetryOptions; import io.temporal.workflow.Functions; +import javax.annotation.Nullable; /** * Interface of workflow task handlers. @@ -23,6 +25,7 @@ final class Result { private final boolean completionCommand; private final Functions.Proc1 resetEventIdHandle; private final Runnable applyPostCompletionMetrics; + private final @Nullable WorkflowExecution completionParentExecution; public Result( String workflowType, @@ -33,6 +36,29 @@ public Result( boolean completionCommand, Functions.Proc1 resetEventIdHandle, Runnable applyPostCompletionMetrics) { + this( + workflowType, + taskCompleted, + taskFailed, + queryCompleted, + requestRetryOptions, + completionCommand, + resetEventIdHandle, + applyPostCompletionMetrics, + null); + } + + public Result( + String workflowType, + RespondWorkflowTaskCompletedRequest taskCompleted, + RespondWorkflowTaskFailedRequest taskFailed, + RespondQueryTaskCompletedRequest queryCompleted, + RpcRetryOptions requestRetryOptions, + boolean completionCommand, + Functions.Proc1 resetEventIdHandle, + Runnable applyPostCompletionMetrics, + @Nullable WorkflowExecution completionParentExecution) { + this.completionParentExecution = completionParentExecution; this.workflowType = workflowType; this.taskCompleted = taskCompleted; this.taskFailed = taskFailed; @@ -43,6 +69,15 @@ public Result( this.applyPostCompletionMetrics = applyPostCompletionMetrics; } + /** + * The workflow to attribute this workflow's own result to, or {@code null} to attribute it to + * the workflow itself. + */ + @Nullable + public WorkflowExecution getCompletionParentExecution() { + return completionParentExecution; + } + public RespondWorkflowTaskCompletedRequest getTaskCompleted() { return taskCompleted; } diff --git a/temporal-sdk/src/main/java/io/temporal/internal/worker/WorkflowWorker.java b/temporal-sdk/src/main/java/io/temporal/internal/worker/WorkflowWorker.java index 98660034d4..cdadd14414 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/worker/WorkflowWorker.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/worker/WorkflowWorker.java @@ -6,10 +6,12 @@ import com.google.common.base.Preconditions; import com.google.common.base.Strings; import com.google.protobuf.ByteString; +import com.google.protobuf.MessageOrBuilder; import com.uber.m3.tally.Scope; import com.uber.m3.tally.Stopwatch; import com.uber.m3.util.ImmutableMap; import io.grpc.StatusRuntimeException; +import io.temporal.api.command.v1.*; import io.temporal.api.common.v1.WorkflowExecution; import io.temporal.api.enums.v1.QueryResultType; import io.temporal.api.enums.v1.TaskQueueKind; @@ -18,15 +20,20 @@ import io.temporal.api.workflowservice.v1.*; import io.temporal.failure.ApplicationFailure; import io.temporal.internal.logging.LoggerTag; +import io.temporal.internal.payload.storage.ExternalStorageRunner; +import io.temporal.internal.payload.visitor.MessageVisitor; import io.temporal.internal.retryer.GrpcMessageTooLargeException; import io.temporal.internal.retryer.GrpcRetryer; import io.temporal.payload.context.WorkflowSerializationContext; +import io.temporal.payload.storage.StorageDriverTargetInfo; +import io.temporal.payload.storage.StorageDriverWorkflowInfo; import io.temporal.serviceclient.MetricsTag; import io.temporal.serviceclient.RpcRetryOptions; import io.temporal.serviceclient.WorkflowServiceStubs; import io.temporal.worker.*; import io.temporal.worker.tuning.*; import java.util.*; +import java.util.concurrent.CancellationException; import java.util.concurrent.CompletableFuture; import java.util.concurrent.RejectedExecutionException; import java.util.concurrent.TimeUnit; @@ -381,6 +388,118 @@ public String toString() { options.getIdentity(), namespace, taskQueue); } + private void storeOutboundPayloads( + com.google.protobuf.Message.Builder builder, @Nullable StorageDriverTargetInfo target) { + storeOutboundPayloads(builder, target, null); + } + + private void storeOutboundPayloads( + com.google.protobuf.Message.Builder builder, + @Nullable StorageDriverTargetInfo target, + @Nullable MessageVisitor targetVisitor) { + ExternalStorageRunner externalStorageRunner = options.getExternalStorageRunner(); + if (externalStorageRunner == null) { + return; + } + try { + externalStorageRunner.store(builder, target, targetVisitor, options.getStorageCancellation()); + } catch (CancellationException e) { + // if the worker is shutting down, extstore will throw a CancellationException and we need to + // rethrow it here so the handle() method can decide what to do. + throw e; + } catch (Exception e) { + throw new ExternalStorageTaskFailure("External storage store failed", e); + } + } + + private static final class ExternalStorageTaskFailure extends RuntimeException { + ExternalStorageTaskFailure(String message, Throwable cause) { + super(message, cause); + } + } + + @Nullable + private StorageDriverTargetInfo parentStorageTarget(@Nullable WorkflowExecution parent) { + if (parent == null || options.getExternalStorageRunner() == null) { + return null; + } + return new StorageDriverWorkflowInfo( + namespace, + Strings.emptyToNull(parent.getWorkflowId()), + Strings.emptyToNull(parent.getRunId()), + null); + } + + @Nullable + private StorageDriverTargetInfo workflowStorageTarget( + WorkflowExecution execution, String workflowType) { + if (options.getExternalStorageRunner() == null) { + return null; + } + return new StorageDriverWorkflowInfo( + namespace, execution.getWorkflowId(), execution.getRunId(), workflowType); + } + + static StorageDriverTargetInfo deriveStorageTarget( + String namespace, StorageDriverTargetInfo current, MessageOrBuilder message) { + return deriveStorageTarget(namespace, current, message, null); + } + + static StorageDriverTargetInfo deriveStorageTarget( + String namespace, + StorageDriverTargetInfo current, + MessageOrBuilder message, + @Nullable StorageDriverTargetInfo completionTarget) { + if (!(message instanceof CommandOrBuilder)) { + return current; + } + CommandOrBuilder command = (CommandOrBuilder) message; + // Keep this exhaustive so new command attributes require an explicit target decision. + switch (command.getAttributesCase()) { + case START_CHILD_WORKFLOW_EXECUTION_COMMAND_ATTRIBUTES: + StartChildWorkflowExecutionCommandAttributesOrBuilder child = + command.getStartChildWorkflowExecutionCommandAttributesOrBuilder(); + return new StorageDriverWorkflowInfo( + namespace, child.getWorkflowId(), null, child.getWorkflowType().getName()); + case SIGNAL_EXTERNAL_WORKFLOW_EXECUTION_COMMAND_ATTRIBUTES: + WorkflowExecution execution = + command.getSignalExternalWorkflowExecutionCommandAttributes().getExecution(); + return new StorageDriverWorkflowInfo( + namespace, execution.getWorkflowId(), execution.getRunId(), null); + case CONTINUE_AS_NEW_WORKFLOW_EXECUTION_COMMAND_ATTRIBUTES: + if (current instanceof StorageDriverWorkflowInfo) { + ContinueAsNewWorkflowExecutionCommandAttributesOrBuilder continueAsNew = + command.getContinueAsNewWorkflowExecutionCommandAttributesOrBuilder(); + StorageDriverWorkflowInfo currentWorkflow = (StorageDriverWorkflowInfo) current; + String workflowType = continueAsNew.getWorkflowType().getName(); + return new StorageDriverWorkflowInfo( + namespace, + currentWorkflow.getId(), + null, + Strings.isNullOrEmpty(workflowType) ? currentWorkflow.getType() : workflowType); + } + return current; + case COMPLETE_WORKFLOW_EXECUTION_COMMAND_ATTRIBUTES: + return completionTarget != null ? completionTarget : current; + case SCHEDULE_ACTIVITY_TASK_COMMAND_ATTRIBUTES: + case ATTRIBUTES_NOT_SET: + case START_TIMER_COMMAND_ATTRIBUTES: + case FAIL_WORKFLOW_EXECUTION_COMMAND_ATTRIBUTES: + case REQUEST_CANCEL_ACTIVITY_TASK_COMMAND_ATTRIBUTES: + case CANCEL_TIMER_COMMAND_ATTRIBUTES: + case CANCEL_WORKFLOW_EXECUTION_COMMAND_ATTRIBUTES: + case REQUEST_CANCEL_EXTERNAL_WORKFLOW_EXECUTION_COMMAND_ATTRIBUTES: + case RECORD_MARKER_COMMAND_ATTRIBUTES: + case UPSERT_WORKFLOW_SEARCH_ATTRIBUTES_COMMAND_ATTRIBUTES: + case PROTOCOL_MESSAGE_COMMAND_ATTRIBUTES: + case MODIFY_WORKFLOW_PROPERTIES_COMMAND_ATTRIBUTES: + case SCHEDULE_NEXUS_OPERATION_COMMAND_ATTRIBUTES: + case REQUEST_CANCEL_NEXUS_OPERATION_COMMAND_ATTRIBUTES: + return current; + } + throw new IllegalStateException("Unhandled command attributes: " + command.getAttributesCase()); + } + private class TaskHandlerImpl implements PollTaskExecutor.TaskHandler { final WorkflowTaskHandler handler; @@ -453,7 +572,26 @@ public void handle(WorkflowTask task) throws Exception { if (queryCompleted != null) { try { sendDirectQueryCompletedResponse( - currentTask.getTaskToken(), queryCompleted.toBuilder(), workflowTypeScope); + currentTask.getTaskToken(), + queryCompleted.toBuilder(), + workflowTypeScope, + workflowStorageTarget(workflowExecution, workflowType)); + } catch (ExternalStorageTaskFailure e) { + Failure failure = + storageFailure( + workflowExecution.getWorkflowId(), e, "Failed to send query response"); + RespondQueryTaskCompletedRequest.Builder queryFailedBuilder = + RespondQueryTaskCompletedRequest.newBuilder() + .setTaskToken(currentTask.getTaskToken()) + .setNamespace(namespace) + .setCompletedType(QueryResultType.QUERY_RESULT_TYPE_FAILED) + .setErrorMessage(failure.getMessage()) + .setFailure(failure); + sendDirectQueryCompletedResponse( + currentTask.getTaskToken(), + queryFailedBuilder, + workflowTypeScope, + workflowStorageTarget(workflowExecution, workflowType)); } catch (StatusRuntimeException e) { GrpcMessageTooLargeException tooLargeException = GrpcMessageTooLargeException.tryWrap(e); @@ -473,7 +611,10 @@ public void handle(WorkflowTask task) throws Exception { .setErrorMessage(failure.getMessage()) .setFailure(failure); sendDirectQueryCompletedResponse( - currentTask.getTaskToken(), queryFailedBuilder, workflowTypeScope); + currentTask.getTaskToken(), + queryFailedBuilder, + workflowTypeScope, + workflowStorageTarget(workflowExecution, workflowType)); } } else { try { @@ -489,7 +630,9 @@ public void handle(WorkflowTask task) throws Exception { currentTask.getTaskToken(), requestBuilder, result.getRequestRetryOptions(), - workflowTypeScope); + workflowTypeScope, + workflowStorageTarget(workflowExecution, workflowType), + parentStorageTarget(result.getCompletionParentExecution())); // If we were processing a speculative WFT the server may instruct us that the // task was dropped by resting out event ID. long resetEventId = response.getResetHistoryEventId(); @@ -509,7 +652,8 @@ public void handle(WorkflowTask task) throws Exception { currentTask.getTaskToken(), taskFailed.toBuilder(), result.getRequestRetryOptions(), - workflowTypeScope); + workflowTypeScope, + workflowStorageTarget(workflowExecution, workflowType)); } // Apply post-completion metrics only if runnable present and the above succeeded @@ -546,9 +690,41 @@ public void handle(WorkflowTask task) throws Exception { currentTask.getTaskToken(), taskFailedBuilder, result.getRequestRetryOptions(), - workflowTypeScope); + workflowTypeScope, + workflowStorageTarget(workflowExecution, workflowType)); + } catch (ExternalStorageTaskFailure e) { + releaseReason = SlotReleaseReason.error(e); + handleReportingFailure( + e, currentTask, result, workflowExecution, workflowTypeScope); + taskFailedCause = + WorkflowTaskFailedCause + .WORKFLOW_TASK_FAILED_CAUSE_WORKFLOW_WORKER_UNHANDLED_FAILURE; + + String messagePrefix = + String.format( + "Failed to send workflow task %s", + taskFailed == null ? "completion" : "failure"); + RespondWorkflowTaskFailedRequest.Builder storageFailedBuilder = + RespondWorkflowTaskFailedRequest.newBuilder() + .setFailure( + storageFailure(workflowExecution.getWorkflowId(), e, messagePrefix)) + .setCause( + WorkflowTaskFailedCause + .WORKFLOW_TASK_FAILED_CAUSE_WORKFLOW_WORKER_UNHANDLED_FAILURE); + sendTaskFailed( + currentTask.getTaskToken(), + storageFailedBuilder, + result.getRequestRetryOptions(), + workflowTypeScope, + workflowStorageTarget(workflowExecution, workflowType)); } } + } catch (CancellationException e) { + if (!options.getStorageCancellation().isCancellationRequested()) { + throw e; + } + log.trace("Abandoned a workflow task while the worker was shutting down", e); + return; } catch (Exception e) { iterationFailed = true; releaseReason = SlotReleaseReason.error(e); @@ -581,6 +757,11 @@ public void handle(WorkflowTask task) throws Exception { workflowTypeScope.counter(MetricsType.WORKFLOW_TASK_HEARTBEAT_COUNTER).inc(1); } } catch (Exception e) { + if (e instanceof CancellationException + && options.getStorageCancellation().isCancellationRequested()) { + log.trace("Abandoned a workflow task while the worker was shutting down", e); + return; + } iterationFailed = true; throw e; } finally { @@ -653,7 +834,9 @@ private RespondWorkflowTaskCompletedResponse sendTaskCompleted( ByteString taskToken, RespondWorkflowTaskCompletedRequest.Builder taskCompleted, RpcRetryOptions retryOptions, - Scope workflowTypeMetricsScope) { + Scope workflowTypeMetricsScope, + @Nullable StorageDriverTargetInfo storageTarget, + @Nullable StorageDriverTargetInfo completionTarget) { GrpcRetryer.GrpcRetryerOptions grpcRetryOptions = new GrpcRetryer.GrpcRetryerOptions( RpcRetryOptions.newBuilder().buildWithDefaultsFrom(retryOptions), null); @@ -676,12 +859,16 @@ private RespondWorkflowTaskCompletedResponse sendTaskCompleted( taskCompleted.setBinaryChecksum(options.getBuildId()); } + MessageVisitor storageTargetVisitor = + (current, message) -> deriveStorageTarget(namespace, current, message, completionTarget); + storeOutboundPayloads(taskCompleted, storageTarget, storageTargetVisitor); + RespondWorkflowTaskCompletedRequest request = taskCompleted.build(); return grpcRetryer.retryWithResult( () -> service .blockingStub() .withOption(METRICS_TAGS_CALL_OPTIONS_KEY, workflowTypeMetricsScope) - .respondWorkflowTaskCompleted(taskCompleted.build()), + .respondWorkflowTaskCompleted(request), grpcRetryOptions); } @@ -690,7 +877,8 @@ private void sendTaskFailed( ByteString taskToken, RespondWorkflowTaskFailedRequest.Builder taskFailed, RpcRetryOptions retryOptions, - Scope workflowTypeMetricsScope) { + Scope workflowTypeMetricsScope, + @Nullable StorageDriverTargetInfo storageTarget) { GrpcRetryer.GrpcRetryerOptions grpcRetryOptions = new GrpcRetryer.GrpcRetryerOptions( RpcRetryOptions.newBuilder().buildWithDefaultsFrom(retryOptions), null); @@ -704,25 +892,30 @@ private void sendTaskFailed( taskFailed.setWorkerVersion(options.workerVersionStamp()); } + storeOutboundPayloads(taskFailed, storageTarget); + RespondWorkflowTaskFailedRequest request = taskFailed.build(); grpcRetryer.retry( () -> service .blockingStub() .withOption(METRICS_TAGS_CALL_OPTIONS_KEY, workflowTypeMetricsScope) - .respondWorkflowTaskFailed(taskFailed.build()), + .respondWorkflowTaskFailed(request), grpcRetryOptions); } private void sendDirectQueryCompletedResponse( ByteString taskToken, RespondQueryTaskCompletedRequest.Builder queryCompleted, - Scope workflowTypeMetricsScope) { + Scope workflowTypeMetricsScope, + @Nullable StorageDriverTargetInfo storageTarget) { queryCompleted.setTaskToken(taskToken).setNamespace(namespace); + storeOutboundPayloads(queryCompleted, storageTarget); + RespondQueryTaskCompletedRequest request = queryCompleted.build(); // Do not retry query response service .blockingStub() .withOption(METRICS_TAGS_CALL_OPTIONS_KEY, workflowTypeMetricsScope) - .respondQueryTaskCompleted(queryCompleted.build()); + .respondQueryTaskCompleted(request); } private void logExceptionDuringResultReporting( @@ -760,6 +953,20 @@ private void handleReportingFailure( workflowExecution, workflowTypeScope, "Failed result reporting to the server", e); } + private Failure storageFailure( + String workflowId, ExternalStorageTaskFailure e, String messagePrefix) { + ApplicationFailure applicationFailure = + ApplicationFailure.newBuilder() + .setMessage(messagePrefix + ": " + (e.getCause() != null ? e.getCause() : e)) + .setType(ExternalStorageTaskFailure.class.getSimpleName()) + .build(); + applicationFailure.setStackTrace(new StackTraceElement[0]); + return options + .getDataConverter() + .withContext(new WorkflowSerializationContext(namespace, workflowId)) + .exceptionToFailure(applicationFailure); + } + private Failure grpcMessageTooLargeFailure( String workflowId, GrpcMessageTooLargeException e, String messagePrefix) { ApplicationFailure applicationFailure = diff --git a/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/ExternalStorageDataConverterTest.java b/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/ExternalStorageDataConverterTest.java index b4fd9cc787..116880aef5 100644 --- a/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/ExternalStorageDataConverterTest.java +++ b/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/ExternalStorageDataConverterTest.java @@ -17,19 +17,12 @@ import io.temporal.payload.codec.PayloadCodec; import io.temporal.payload.storage.ExternalStorage; import io.temporal.payload.storage.StorageDriver; -import io.temporal.payload.storage.StorageDriverClaim; -import io.temporal.payload.storage.StorageDriverRetrieveContext; -import io.temporal.payload.storage.StorageDriverStoreContext; -import io.temporal.payload.storage.StorageDriverTargetInfo; import io.temporal.payload.storage.StorageDriverWorkflowInfo; import java.lang.reflect.Type; import java.util.ArrayList; import java.util.Collections; -import java.util.HashMap; import java.util.List; -import java.util.Map; import java.util.Optional; -import java.util.concurrent.CompletableFuture; import java.util.concurrent.atomic.AtomicInteger; import org.junit.Test; @@ -39,7 +32,7 @@ public class ExternalStorageDataConverterTest { @Test public void payloadsRoundTripThroughStorage() { - RecordingDriver driver = new RecordingDriver(); + TestStorageDriver driver = TestStorageDriver.create(); DataConverter converter = resolving(driver, 0); Optional stored = converter.toPayloads("a", "b"); @@ -53,19 +46,19 @@ public void payloadsRoundTripThroughStorage() { @Test public void payloadsBelowThresholdStayInline() { - RecordingDriver driver = new RecordingDriver(); + TestStorageDriver driver = TestStorageDriver.create(); DataConverter converter = resolving(driver, 1024); Optional stored = converter.toPayloads("small"); assertFalse(ExternalStorageReferences.isReference(stored.get().getPayloads(0))); - assertTrue(driver.objects.isEmpty()); + assertTrue(driver.storedPayloads().isEmpty()); assertEquals("small", converter.fromPayloads(0, stored, String.class, String.class)); } @Test public void readingOneArgumentDoesNotFetchTheRest() { - RecordingDriver driver = new RecordingDriver(); + TestStorageDriver driver = TestStorageDriver.create(); DataConverter converter = resolving(driver, 0); Optional stored = converter.toPayloads("first", "second", "third"); @@ -78,7 +71,7 @@ public void readingOneArgumentDoesNotFetchTheRest() { @Test public void singlePayloadRoundTrips() { - DataConverter converter = resolving(new RecordingDriver(), 0); + DataConverter converter = resolving(TestStorageDriver.create(), 0); Optional stored = converter.toPayload("value"); @@ -88,7 +81,7 @@ public void singlePayloadRoundTrips() { @Test public void failureDetailsRoundTrip() { - DataConverter converter = resolving(new RecordingDriver(), 0); + DataConverter converter = resolving(TestStorageDriver.create(), 0); Failure failure = converter.exceptionToFailure( @@ -104,27 +97,27 @@ public void failureDetailsRoundTrip() { @Test public void storageTargetReachesTheDriver() { - RecordingDriver driver = new RecordingDriver(); + TestStorageDriver driver = TestStorageDriver.create(); StorageDriverWorkflowInfo target = new StorageDriverWorkflowInfo("ns", "wf-1", null, null); ExternalStorageDataConverter converter = new ExternalStorageDataConverter(plain, runner(driver, 0)).withStorageTarget(target); converter.toPayloads("x"); - assertEquals(target, driver.lastTarget); + assertEquals(target, driver.lastTarget()); } @Test public void withoutATargetTheDriverSeesNone() { - RecordingDriver driver = new RecordingDriver(); + TestStorageDriver driver = TestStorageDriver.create(); resolving(driver, 0).toPayloads("x"); - assertNull(driver.lastTarget); + assertNull(driver.lastTarget()); } @Test public void arrayFromPayloadsRoundTrips() { - RecordingDriver driver = new RecordingDriver(); + TestStorageDriver driver = TestStorageDriver.create(); DataConverter converter = resolving(driver, 0); Optional stored = converter.toPayloads("a", 42); @@ -141,7 +134,7 @@ public void arrayFromPayloadsRoundTrips() { @Test public void arrayFromPayloadsWithAbsentContentUsesDefaults() { - DataConverter converter = resolving(new RecordingDriver(), 0); + DataConverter converter = resolving(TestStorageDriver.create(), 0); Object[] values = converter.fromPayloads( @@ -152,7 +145,7 @@ public void arrayFromPayloadsWithAbsentContentUsesDefaults() { @Test public void arrayFromPayloadsDecodesThroughTheCodecInOneBatch() { - RecordingDriver driver = new RecordingDriver(); + TestStorageDriver driver = TestStorageDriver.create(); CountingCodec codec = new CountingCodec(); DataConverter converter = codecBacked(driver, codec); @@ -175,14 +168,14 @@ public void arrayFromPayloadsDecodesThroughTheCodecInOneBatch() { */ @Test public void driversOnlyEverSeeCodecEncodedPayloads() { - RecordingDriver driver = new RecordingDriver(); + TestStorageDriver driver = TestStorageDriver.create(); CountingCodec codec = new CountingCodec(); DataConverter converter = codecBacked(driver, codec); Optional stored = converter.toPayloads("a", "b", "c"); - assertEquals(3, driver.objects.size()); - for (Payload payload : driver.objects.values()) { + assertEquals(3, driver.storedCount()); + for (Payload payload : driver.storedPayloads()) { String data = payload.getData().toStringUtf8(); assertFalse(data.contains("\"a\"")); assertFalse(data.contains("\"b\"")); @@ -242,46 +235,4 @@ private static List apply(List payloads) { return out; } } - - private static final class RecordingDriver implements StorageDriver { - final Map objects = new HashMap<>(); - final List retrievedKeys = new ArrayList<>(); - volatile StorageDriverTargetInfo lastTarget; - private int counter = 0; - - @Override - public String getName() { - return "test"; - } - - @Override - public String getType() { - return "test.inmemory"; - } - - @Override - public synchronized CompletableFuture> store( - StorageDriverStoreContext context, List payloads) { - lastTarget = context.getTarget(); - List claims = new ArrayList<>(); - for (Payload payload : payloads) { - String key = "k-" + (counter++); - objects.put(key, payload); - claims.add(new StorageDriverClaim(Collections.singletonMap("key", key))); - } - return CompletableFuture.completedFuture(claims); - } - - @Override - public synchronized CompletableFuture> retrieve( - StorageDriverRetrieveContext context, List claims) { - List payloads = new ArrayList<>(); - for (StorageDriverClaim claim : claims) { - String key = claim.getClaimData().get("key"); - retrievedKeys.add(key); - payloads.add(objects.get(key)); - } - return CompletableFuture.completedFuture(payloads); - } - } } diff --git a/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/ExternalStoragePayloadTransformerTest.java b/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/ExternalStoragePayloadTransformerTest.java index 2dcb58f384..8a48625450 100644 --- a/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/ExternalStoragePayloadTransformerTest.java +++ b/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/ExternalStoragePayloadTransformerTest.java @@ -18,7 +18,6 @@ import io.temporal.payload.storage.StorageDriverRetrieveContext; import io.temporal.payload.storage.StorageDriverSelector; import io.temporal.payload.storage.StorageDriverStoreContext; -import java.util.ArrayList; import java.util.Arrays; import java.util.Collections; import java.util.HashMap; @@ -36,7 +35,7 @@ public class ExternalStoragePayloadTransformerTest { @Test public void storesAndRetrievesRoundTrip() throws Exception { - InMemoryDriver driver = new InMemoryDriver("d1"); + TestStorageDriver driver = TestStorageDriver.named("d1"); ExternalStoragePayloadTransformer transformer = transformer(driver, 0); List input = Arrays.asList(payload("a"), payload("b")); @@ -56,7 +55,7 @@ public void storesAndRetrievesRoundTrip() throws Exception { @Test public void payloadBelowThresholdStaysInline() throws Exception { - InMemoryDriver driver = new InMemoryDriver("d1"); + TestStorageDriver driver = TestStorageDriver.named("d1"); ExternalStoragePayloadTransformer transformer = transformer(driver, 100); Payload small = payload("x"); Payload large = payload(repeat("y", 200)); @@ -72,7 +71,7 @@ public void payloadBelowThresholdStaysInline() throws Exception { @Test public void selectorReturningNullKeepsInline() throws Exception { - InMemoryDriver driver = new InMemoryDriver("d1"); + TestStorageDriver driver = TestStorageDriver.named("d1"); ExternalStoragePayloadTransformer transformer = ExternalStoragePayloadTransformer.fromOptions( ExternalStorage.newBuilder() @@ -92,8 +91,8 @@ public void selectorReturningNullKeepsInline() throws Exception { @Test public void multipleDriversBatchPerDriverAndPreserveOrder() throws Exception { - InMemoryDriver d1 = new InMemoryDriver("d1"); - InMemoryDriver d2 = new InMemoryDriver("d2"); + TestStorageDriver d1 = TestStorageDriver.named("d1"); + TestStorageDriver d2 = TestStorageDriver.named("d2"); Map byPrefix = new HashMap<>(); byPrefix.put("1", d1); byPrefix.put("2", d2); @@ -179,7 +178,7 @@ public CompletableFuture> retrieve( @Test public void unknownDriverOnRetrieveFails() { - InMemoryDriver driver = new InMemoryDriver("d1"); + TestStorageDriver driver = TestStorageDriver.named("d1"); ExternalStoragePayloadTransformer transformer = transformer(driver, 0); Payload reference = ExternalStorageReferences.toReferencePayload( @@ -194,8 +193,8 @@ public void unknownDriverOnRetrieveFails() { @Test public void selectorReturningUnregisteredDriverFails() { - InMemoryDriver registered = new InMemoryDriver("d1"); - InMemoryDriver stranger = new InMemoryDriver("d2"); + TestStorageDriver registered = TestStorageDriver.named("d1"); + TestStorageDriver stranger = TestStorageDriver.named("d2"); ExternalStoragePayloadTransformer transformer = ExternalStoragePayloadTransformer.fromOptions( ExternalStorage.newBuilder() @@ -309,7 +308,7 @@ public CompletableFuture> retrieve( @Test public void selectorObservesCallerCancellationToken() { - InMemoryDriver driver = new InMemoryDriver("d1"); + TestStorageDriver driver = TestStorageDriver.named("d1"); CancelSource caller = new CancelSource<>(CancellationException::new); AtomicReference> observed = new AtomicReference<>(); ExternalStoragePayloadTransformer transformer = @@ -400,39 +399,4 @@ public CompletableFuture> retrieve( throw new UnsupportedOperationException(); } } - - private static class InMemoryDriver extends FakeDriver { - final Map objects = new HashMap<>(); - final List storeBatchSizes = new ArrayList<>(); - final List retrieveBatchSizes = new ArrayList<>(); - private int counter = 0; - - InMemoryDriver(String name) { - super(name); - } - - @Override - public CompletableFuture> store( - StorageDriverStoreContext context, List payloads) { - storeBatchSizes.add(payloads.size()); - List claims = new ArrayList<>(); - for (Payload payload : payloads) { - String key = getName() + "-" + (counter++); - objects.put(key, payload); - claims.add(new StorageDriverClaim(Collections.singletonMap("key", key))); - } - return CompletableFuture.completedFuture(claims); - } - - @Override - public CompletableFuture> retrieve( - StorageDriverRetrieveContext context, List claims) { - retrieveBatchSizes.add(claims.size()); - List payloads = new ArrayList<>(); - for (StorageDriverClaim claim : claims) { - payloads.add(objects.get(claim.getClaimData().get("key"))); - } - return CompletableFuture.completedFuture(payloads); - } - } } diff --git a/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/ExternalStorageRunnerTest.java b/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/ExternalStorageRunnerTest.java index 7cef800eb5..73d5010256 100644 --- a/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/ExternalStorageRunnerTest.java +++ b/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/ExternalStorageRunnerTest.java @@ -8,6 +8,7 @@ import com.google.protobuf.ByteString; import io.temporal.api.command.v1.Command; +import io.temporal.api.command.v1.CommandOrBuilder; import io.temporal.api.command.v1.CompleteWorkflowExecutionCommandAttributes; import io.temporal.api.command.v1.ScheduleActivityTaskCommandAttributes; import io.temporal.api.command.v1.ScheduleActivityTaskCommandAttributesOrBuilder; @@ -16,6 +17,7 @@ import io.temporal.api.common.v1.Payload; import io.temporal.api.common.v1.Payloads; import io.temporal.api.common.v1.SearchAttributes; +import io.temporal.api.sdk.v1.UserMetadata; import io.temporal.api.workflowservice.v1.RespondWorkflowTaskCompletedRequest; import io.temporal.common.CancellationToken; import io.temporal.internal.concurrent.structured.CancelSource; @@ -23,18 +25,9 @@ import io.temporal.payload.storage.ExternalStorage; import io.temporal.payload.storage.StorageDriver; import io.temporal.payload.storage.StorageDriverActivityInfo; -import io.temporal.payload.storage.StorageDriverClaim; -import io.temporal.payload.storage.StorageDriverRetrieveContext; -import io.temporal.payload.storage.StorageDriverStoreContext; import io.temporal.payload.storage.StorageDriverTargetInfo; import io.temporal.payload.storage.StorageDriverWorkflowInfo; -import java.util.ArrayList; -import java.util.Collections; -import java.util.HashMap; -import java.util.List; -import java.util.Map; import java.util.concurrent.CancellationException; -import java.util.concurrent.CompletableFuture; import java.util.concurrent.atomic.AtomicInteger; import org.junit.Test; @@ -43,7 +36,7 @@ public class ExternalStorageRunnerTest { @Test public void storeAndRetrieveRoundTripsOverAMessage() throws Exception { - InMemoryDriver driver = new InMemoryDriver("d1"); + TestStorageDriver driver = TestStorageDriver.named("d1"); ExternalStorageRunner transformer = transformer(driver, 0); Payloads message = Payloads.newBuilder().addPayloads(payload("a")).addPayloads(payload("b")).build(); @@ -61,7 +54,7 @@ public void storeAndRetrieveRoundTripsOverAMessage() throws Exception { @Test public void walksNestedPayloads() throws Exception { - InMemoryDriver driver = new InMemoryDriver("d1"); + TestStorageDriver driver = TestStorageDriver.named("d1"); ExternalStorageRunner transformer = transformer(driver, 0); Command command = Command.newBuilder() @@ -81,7 +74,7 @@ public void walksNestedPayloads() throws Exception { @Test public void payloadBelowThresholdLeavesMessageUnchanged() throws Exception { - InMemoryDriver driver = new InMemoryDriver("d1"); + TestStorageDriver driver = TestStorageDriver.named("d1"); ExternalStorageRunner transformer = transformer(driver, 1024); Payloads message = Payloads.newBuilder().addPayloads(payload("small")).build(); @@ -96,7 +89,7 @@ public void payloadBelowThresholdLeavesMessageUnchanged() throws Exception { @Test public void searchAttributesAreNotOffloaded() throws Exception { - InMemoryDriver driver = new InMemoryDriver("d1"); + TestStorageDriver driver = TestStorageDriver.named("d1"); ExternalStorageRunner transformer = transformer(driver, 0); Command command = Command.newBuilder() @@ -122,7 +115,7 @@ public void searchAttributesAreNotOffloaded() throws Exception { @Test public void throwIfContainsReferenceThrowsOnANestedReference() throws Exception { - ExternalStorageRunner transformer = transformer(new InMemoryDriver("d1"), 0); + ExternalStorageRunner transformer = transformer(TestStorageDriver.named("d1"), 0); RespondWorkflowTaskCompletedRequest.Builder request = RespondWorkflowTaskCompletedRequest.newBuilder() .addCommands( @@ -142,7 +135,7 @@ public void throwIfContainsReferenceThrowsOnANestedReference() throws Exception @Test public void throwIfContainsReferenceThrowsOnReference() throws Exception { - InMemoryDriver driver = new InMemoryDriver("d1"); + TestStorageDriver driver = TestStorageDriver.named("d1"); ExternalStorageRunner transformer = transformer(driver, 0); Payloads.Builder builder = Payloads.newBuilder().addPayloads(payload("a")); transformer.store(builder, null, null, CancellationToken.none()); @@ -160,14 +153,16 @@ public void throwIfContainsReferenceAllowsInlinePayloads() { } @Test - public void storeAppliesPerCommandTargetFromMessageVisitor() { - TargetCapturingDriver driver = new TargetCapturingDriver("d1"); + public void storeScopesCommandTargetOverAttributesAndMetadata() { + TestStorageDriver driver = TestStorageDriver.named("d1"); ExternalStorageRunner storage = transformer(driver, 0); RespondWorkflowTaskCompletedRequest.Builder request = RespondWorkflowTaskCompletedRequest.newBuilder() .addCommands( Command.newBuilder() + .setUserMetadata( + UserMetadata.newBuilder().setSummary(payload("activity-summary"))) .setScheduleActivityTaskCommandAttributes( ScheduleActivityTaskCommandAttributes.newBuilder() .setActivityId("act-1") @@ -184,11 +179,15 @@ public void storeAppliesPerCommandTargetFromMessageVisitor() { new StorageDriverWorkflowInfo("ns", "wf-1", "run-1", "MyWorkflow"); MessageVisitor visitor = (current, message) -> { - if (message instanceof ScheduleActivityTaskCommandAttributesOrBuilder) { - ScheduleActivityTaskCommandAttributesOrBuilder attrs = - (ScheduleActivityTaskCommandAttributesOrBuilder) message; - return new StorageDriverActivityInfo( - "ns", attrs.getActivityId(), null, attrs.getActivityType().getName()); + if (message instanceof CommandOrBuilder) { + CommandOrBuilder command = (CommandOrBuilder) message; + if (command.getAttributesCase() + == Command.AttributesCase.SCHEDULE_ACTIVITY_TASK_COMMAND_ATTRIBUTES) { + ScheduleActivityTaskCommandAttributesOrBuilder attrs = + command.getScheduleActivityTaskCommandAttributesOrBuilder(); + return new StorageDriverActivityInfo( + "ns", attrs.getActivityId(), null, attrs.getActivityType().getName()); + } } return current; }; @@ -198,12 +197,15 @@ public void storeAppliesPerCommandTargetFromMessageVisitor() { assertEquals( new StorageDriverActivityInfo("ns", "act-1", null, "MyActivity"), driver.targetFor("activity-input")); + assertEquals( + new StorageDriverActivityInfo("ns", "act-1", null, "MyActivity"), + driver.targetFor("activity-summary")); assertEquals(workflowTarget, driver.targetFor("wf-result")); } @Test public void callerCancellationAbortsStore() { - ExternalStorageRunner storage = transformer(new HangingDriver("d1"), 0); + ExternalStorageRunner storage = transformer(TestStorageDriver.named("d1").neverAnswers(), 0); CancelSource caller = new CancelSource<>(CancellationException::new); caller.cancel(); Payloads message = Payloads.newBuilder().addPayloads(payload("big")).build(); @@ -216,7 +218,7 @@ public void callerCancellationAbortsStore() { @Test public void completedOperationsReleaseTheirCancellationRegistrations() { RegistrationCountingToken token = new RegistrationCountingToken(); - ExternalStorageRunner storage = transformer(new InMemoryDriver("d1"), 0); + ExternalStorageRunner storage = transformer(TestStorageDriver.named("d1"), 0); for (int i = 0; i < 5; i++) { Payloads.Builder builder = Payloads.newBuilder().addPayloads(payload("a")); @@ -263,121 +265,4 @@ public Registration onCancel(Runnable callback) { return open::decrementAndGet; } } - - private static final class InMemoryDriver implements StorageDriver { - private final String name; - private final Map objects = new HashMap<>(); - final List storeBatchSizes = new ArrayList<>(); - private int counter = 0; - - InMemoryDriver(String name) { - this.name = name; - } - - @Override - public String getName() { - return name; - } - - @Override - public String getType() { - return "test.inmemory"; - } - - @Override - public synchronized CompletableFuture> store( - StorageDriverStoreContext context, List payloads) { - storeBatchSizes.add(payloads.size()); - List claims = new ArrayList<>(); - for (Payload payload : payloads) { - String key = name + "-" + (counter++); - objects.put(key, payload); - claims.add(new StorageDriverClaim(Collections.singletonMap("key", key))); - } - return CompletableFuture.completedFuture(claims); - } - - @Override - public synchronized CompletableFuture> retrieve( - StorageDriverRetrieveContext context, List claims) { - List payloads = new ArrayList<>(); - for (StorageDriverClaim claim : claims) { - payloads.add(objects.get(claim.getClaimData().get("key"))); - } - return CompletableFuture.completedFuture(payloads); - } - } - - private static final class TargetCapturingDriver implements StorageDriver { - private final String name; - private final Map targetByData = new HashMap<>(); - private int counter = 0; - - TargetCapturingDriver(String name) { - this.name = name; - } - - @Override - public String getName() { - return name; - } - - @Override - public String getType() { - return "test.capture"; - } - - @Override - public synchronized CompletableFuture> store( - StorageDriverStoreContext context, List payloads) { - List claims = new ArrayList<>(); - for (Payload payload : payloads) { - targetByData.put(payload.getData().toStringUtf8(), context.getTarget()); - claims.add( - new StorageDriverClaim(Collections.singletonMap("key", name + "-" + (counter++)))); - } - return CompletableFuture.completedFuture(claims); - } - - synchronized StorageDriverTargetInfo targetFor(String data) { - return targetByData.get(data); - } - - @Override - public CompletableFuture> retrieve( - StorageDriverRetrieveContext context, List claims) { - throw new UnsupportedOperationException(); - } - } - - /** Driver whose operations never settle, so only cancellation can end a blocking call. */ - private static final class HangingDriver implements StorageDriver { - private final String name; - - HangingDriver(String name) { - this.name = name; - } - - @Override - public String getName() { - return name; - } - - @Override - public String getType() { - return "test.hanging"; - } - - @Override - public CompletableFuture> store( - StorageDriverStoreContext context, List payloads) { - return new CompletableFuture<>(); - } - - @Override - public CompletableFuture> retrieve( - StorageDriverRetrieveContext context, List claims) { - return new CompletableFuture<>(); - } - } } diff --git a/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/TestStorageDriver.java b/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/TestStorageDriver.java new file mode 100644 index 0000000000..b7e21ce63b --- /dev/null +++ b/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/TestStorageDriver.java @@ -0,0 +1,247 @@ +package io.temporal.internal.payload.storage; + +import io.temporal.api.common.v1.Payload; +import io.temporal.payload.storage.StorageDriver; +import io.temporal.payload.storage.StorageDriverClaim; +import io.temporal.payload.storage.StorageDriverRetrieveContext; +import io.temporal.payload.storage.StorageDriverStoreContext; +import io.temporal.payload.storage.StorageDriverTargetInfo; +import java.util.ArrayList; +import java.util.Collection; +import java.util.Collections; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.concurrent.CancellationException; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CopyOnWriteArrayList; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; +import javax.annotation.Nullable; + +/** + * In-memory storage driver for tests. Records what it was asked to do, and can be told to fail, + * block or never answer so that one driver covers the cases the tests need. + */ +public final class TestStorageDriver implements StorageDriver { + + private final String name; + private final Map objects = new HashMap<>(); + private final Map targetByData = new HashMap<>(); + private int counter; + + public final List storeBatchSizes = new CopyOnWriteArrayList<>(); + public final List retrieveBatchSizes = new CopyOnWriteArrayList<>(); + public final List retrievedKeys = new CopyOnWriteArrayList<>(); + public final List targets = new CopyOnWriteArrayList<>(); + public final List storedData = new CopyOnWriteArrayList<>(); + public final AtomicInteger stores = new AtomicInteger(); + public final AtomicInteger retrieves = new AtomicInteger(); + public final AtomicInteger injectedFailures = new AtomicInteger(); + + private volatile boolean neverAnswers; + private volatile boolean cancelInsteadOfFailing; + private final AtomicInteger storeFailures = new AtomicInteger(); + private final AtomicInteger retrieveFailures = new AtomicInteger(); + private volatile @Nullable String failStoresContaining; + private volatile @Nullable CountDownLatch storeEntered; + private volatile @Nullable CountDownLatch releaseStore; + + private TestStorageDriver(String name) { + this.name = name; + } + + public static TestStorageDriver create() { + return new TestStorageDriver("test"); + } + + public static TestStorageDriver named(String name) { + return new TestStorageDriver(name); + } + + /** Neither storing nor retrieving ever finishes, so only cancellation can end the call. */ + public TestStorageDriver neverAnswers() { + this.neverAnswers = true; + return this; + } + + public TestStorageDriver failStores(int times) { + this.storeFailures.set(times); + return this; + } + + /** Fails the next {@code times} stores the way an abandoned call does. */ + public TestStorageDriver cancelStores(int times) { + this.storeFailures.set(times); + this.cancelInsteadOfFailing = true; + return this; + } + + /** Fails the next {@code times} stores that carry a payload containing {@code marker}. */ + public TestStorageDriver failStoresContaining(String marker, int times) { + this.failStoresContaining = marker; + this.storeFailures.set(times); + return this; + } + + public TestStorageDriver failRetrieves(int times) { + this.retrieveFailures.set(times); + return this; + } + + /** Holds each store until {@code release}, counting down {@code entered} on the way in. */ + public TestStorageDriver blockStores(CountDownLatch entered, CountDownLatch release) { + this.storeEntered = entered; + this.releaseStore = release; + return this; + } + + @Override + public String getName() { + return name; + } + + @Override + public String getType() { + return "test.in-memory"; + } + + @Override + public synchronized CompletableFuture> store( + StorageDriverStoreContext context, List payloads) { + stores.incrementAndGet(); + storeBatchSizes.add(payloads.size()); + targets.add(context.getTarget()); + + CountDownLatch entered = storeEntered; + if (entered != null) { + entered.countDown(); + } + CountDownLatch release = releaseStore; + if (release != null) { + try { + release.await(10, TimeUnit.SECONDS); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + } + if (neverAnswers) { + return new CompletableFuture<>(); + } + if (shouldFailStore(payloads)) { + injectedFailures.incrementAndGet(); + return cancelInsteadOfFailing + ? cancelled("external storage stopped") + : failed("storage unavailable"); + } + + List claims = new ArrayList<>(); + for (Payload payload : payloads) { + String data = payload.getData().toStringUtf8(); + storedData.add(data); + targetByData.put(data, context.getTarget()); + String key = name + "-" + (counter++); + objects.put(key, payload); + claims.add(new StorageDriverClaim(Collections.singletonMap("key", key))); + } + return CompletableFuture.completedFuture(claims); + } + + @Override + public synchronized CompletableFuture> retrieve( + StorageDriverRetrieveContext context, List claims) { + retrieves.incrementAndGet(); + retrieveBatchSizes.add(claims.size()); + for (StorageDriverClaim claim : claims) { + retrievedKeys.add(claim.getClaimData().get("key")); + } + if (neverAnswers) { + return new CompletableFuture<>(); + } + if (retrieveFailures.get() > 0) { + retrieveFailures.decrementAndGet(); + injectedFailures.incrementAndGet(); + return failed("storage unavailable"); + } + List payloads = new ArrayList<>(); + for (StorageDriverClaim claim : claims) { + payloads.add(objects.get(claim.getClaimData().get("key"))); + } + return CompletableFuture.completedFuture(payloads); + } + + /** Forgets everything stored and recorded, and clears any injected behaviour. */ + public synchronized void reset() { + objects.clear(); + targetByData.clear(); + counter = 0; + storeBatchSizes.clear(); + retrieveBatchSizes.clear(); + retrievedKeys.clear(); + targets.clear(); + storedData.clear(); + stores.set(0); + retrieves.set(0); + injectedFailures.set(0); + storeFailures.set(0); + retrieveFailures.set(0); + failStoresContaining = null; + storeEntered = null; + releaseStore = null; + neverAnswers = false; + cancelInsteadOfFailing = false; + } + + /** The target supplied when the payload with this data was stored. */ + public synchronized StorageDriverTargetInfo targetFor(String data) { + return targetByData.get(data); + } + + public synchronized int storedCount() { + return objects.size(); + } + + public synchronized Collection storedPayloads() { + return new ArrayList<>(objects.values()); + } + + /** The target supplied with the most recent store, or {@code null} if nothing was stored. */ + public StorageDriverTargetInfo lastTarget() { + return targets.isEmpty() ? null : targets.get(targets.size() - 1); + } + + public boolean stored(String substring) { + return storedData.stream().anyMatch(data -> data.contains(substring)); + } + + private boolean shouldFailStore(List payloads) { + if (storeFailures.get() <= 0) { + return false; + } + String marker = failStoresContaining; + if (marker == null) { + storeFailures.decrementAndGet(); + return true; + } + for (Payload payload : payloads) { + if (payload.getData().toStringUtf8().contains(marker)) { + storeFailures.decrementAndGet(); + return true; + } + } + return false; + } + + private static CompletableFuture cancelled(String message) { + CompletableFuture future = new CompletableFuture<>(); + future.completeExceptionally(new CancellationException(message)); + return future; + } + + private static CompletableFuture failed(String message) { + CompletableFuture future = new CompletableFuture<>(); + future.completeExceptionally(new IllegalStateException(message)); + return future; + } +} diff --git a/temporal-sdk/src/test/java/io/temporal/internal/replay/ReplayWorkflowRunTaskHandlerTaskHandlerTests.java b/temporal-sdk/src/test/java/io/temporal/internal/replay/ReplayWorkflowRunTaskHandlerTaskHandlerTests.java index ed6446678a..046255218d 100644 --- a/temporal-sdk/src/test/java/io/temporal/internal/replay/ReplayWorkflowRunTaskHandlerTaskHandlerTests.java +++ b/temporal-sdk/src/test/java/io/temporal/internal/replay/ReplayWorkflowRunTaskHandlerTaskHandlerTests.java @@ -3,26 +3,36 @@ import static junit.framework.TestCase.assertEquals; import static junit.framework.TestCase.assertNotNull; import static org.junit.Assert.assertFalse; +import static org.junit.Assert.assertNotEquals; +import static org.junit.Assert.assertThrows; import static org.junit.Assert.assertTrue; import static org.junit.Assume.assumeFalse; import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; import com.google.protobuf.ByteString; import com.google.protobuf.util.Durations; import com.uber.m3.tally.NoopScope; +import io.temporal.api.common.v1.Payload; +import io.temporal.api.common.v1.Payloads; import io.temporal.api.enums.v1.EventType; import io.temporal.api.history.v1.History; import io.temporal.api.history.v1.HistoryEvent; import io.temporal.api.taskqueue.v1.StickyExecutionAttributes; import io.temporal.api.workflowservice.v1.*; +import io.temporal.common.CancellationToken; import io.temporal.internal.common.InternalUtils; +import io.temporal.internal.concurrent.structured.CancelSource; +import io.temporal.internal.payload.storage.ExternalStorageRunner; +import io.temporal.internal.payload.storage.TestStorageDriver; import io.temporal.internal.statemachines.ExecuteLocalActivityParameters; import io.temporal.internal.worker.SingleWorkerOptions; import io.temporal.internal.worker.WorkflowExecutorCache; import io.temporal.internal.worker.WorkflowRunLockManager; import io.temporal.internal.worker.WorkflowTaskHandler; +import io.temporal.payload.storage.ExternalStorage; import io.temporal.serviceclient.Version; import io.temporal.serviceclient.WorkflowServiceStubs; import io.temporal.testUtils.HistoryUtils; @@ -31,8 +41,10 @@ import java.util.HashMap; import java.util.List; import java.util.Optional; +import java.util.concurrent.CancellationException; import org.junit.Rule; import org.junit.Test; +import org.mockito.ArgumentCaptor; public class ReplayWorkflowRunTaskHandlerTaskHandlerTests { @@ -121,6 +133,231 @@ public void workflowTaskFailOnIncompleteHistory() throws Throwable { result.getTaskFailed().getFailure().getMessage()); } + @Test + public void resolvesExternalStorageReferencesInTheWorkflowTaskItself() throws Throwable { + TestStorageDriver driver = TestStorageDriver.create(); + ExternalStorageRunner externalStorage = + ExternalStorageRunner.create( + ExternalStorage.newBuilder().setDriver(driver).setPayloadSizeThreshold(0).build()); + PollWorkflowTaskQueueResponse fullTask = HistoryUtils.generateWorkflowTaskWithInitialHistory(); + HistoryEvent startedEvent = fullTask.getHistory().getEvents(0); + Payload input = Payload.newBuilder().setData(ByteString.copyFromUtf8("input")).build(); + History.Builder storedHistory = + fullTask.getHistory().toBuilder() + .setEvents( + 0, + startedEvent.toBuilder() + .setWorkflowExecutionStartedEventAttributes( + startedEvent.getWorkflowExecutionStartedEventAttributes().toBuilder() + .setInput(Payloads.newBuilder().addPayloads(input)))); + externalStorage.store(storedHistory, null, null, CancellationToken.none()); + assertNotEquals( + "the payload must be replaced by a reference, otherwise this test proves nothing", + input, + storedInput(storedHistory)); + + WorkflowServiceStubs client = mock(WorkflowServiceStubs.class); + when(client.getServerCapabilities()) + .thenReturn(() -> GetSystemInfoResponse.Capabilities.newBuilder().build()); + + ReplayWorkflow workflow = mock(ReplayWorkflow.class); + when(workflow.eventLoop()).thenReturn(true); + when(workflow.getOutput()).thenReturn(Optional.empty()); + WorkflowContext workflowContext = mock(WorkflowContext.class); + when(workflowContext.getRunningUpdateHandlers()).thenReturn(new HashMap<>()); + when(workflow.getWorkflowContext()).thenReturn(workflowContext); + ReplayWorkflowFactory workflowFactory = mock(ReplayWorkflowFactory.class); + when(workflowFactory.getWorkflow(any(), any())).thenReturn(workflow); + + WorkflowTaskHandler taskHandler = + new ReplayWorkflowTaskHandler( + "namespace", + workflowFactory, + new WorkflowExecutorCache(10, new WorkflowRunLockManager(), new NoopScope()), + SingleWorkerOptions.newBuilder().setExternalStorageRunner(externalStorage).build(), + null, + Duration.ofSeconds(5), + client, + null); + + taskHandler.handleWorkflowTask(fullTask.toBuilder().setHistory(storedHistory).build()); + + ArgumentCaptor event = ArgumentCaptor.forClass(HistoryEvent.class); + verify(workflow).start(event.capture(), any()); + assertEquals( + input, + event.getValue().getWorkflowExecutionStartedEventAttributes().getInput().getPayloads(0)); + } + + private static Payload storedInput(History.Builder history) { + return history + .getEvents(0) + .getWorkflowExecutionStartedEventAttributes() + .getInput() + .getPayloads(0); + } + + @Test + public void aCancelledDownloadIsNotReportedAsAWorkflowTaskFailure() throws Throwable { + TestStorageDriver driver = TestStorageDriver.create(); + ExternalStorageRunner externalStorage = + ExternalStorageRunner.create( + ExternalStorage.newBuilder().setDriver(driver).setPayloadSizeThreshold(0).build()); + PollWorkflowTaskQueueResponse fullTask = HistoryUtils.generateWorkflowTaskWithInitialHistory(); + HistoryEvent startedEvent = fullTask.getHistory().getEvents(0); + Payload input = Payload.newBuilder().setData(ByteString.copyFromUtf8("input")).build(); + History.Builder storedHistory = + fullTask.getHistory().toBuilder() + .setEvents( + 0, + startedEvent.toBuilder() + .setWorkflowExecutionStartedEventAttributes( + startedEvent.getWorkflowExecutionStartedEventAttributes().toBuilder() + .setInput(Payloads.newBuilder().addPayloads(input)))); + externalStorage.store(storedHistory, null, null, CancellationToken.none()); + driver.neverAnswers(); + + CancelSource stopping = + new CancelSource<>(() -> new CancellationException("Worker shutdown")); + stopping.cancel(); + + WorkflowServiceStubs client = mock(WorkflowServiceStubs.class); + when(client.getServerCapabilities()) + .thenReturn(() -> GetSystemInfoResponse.Capabilities.newBuilder().build()); + + WorkflowTaskHandler taskHandler = + new ReplayWorkflowTaskHandler( + "namespace", + setUpMockWorkflowFactory(), + new WorkflowExecutorCache(10, new WorkflowRunLockManager(), new NoopScope()), + SingleWorkerOptions.newBuilder() + .setExternalStorageRunner(externalStorage) + .setStorageCancellation(stopping.token()) + .build(), + null, + Duration.ofSeconds(5), + client, + null); + + assertThrows( + "stopping storage must not be turned into a workflow task failure", + CancellationException.class, + () -> + taskHandler.handleWorkflowTask(fullTask.toBuilder().setHistory(storedHistory).build())); + } + + @Test + public void aFailedDownloadIsReportedAsAWorkflowTaskFailure() throws Throwable { + TestStorageDriver driver = TestStorageDriver.create(); + ExternalStorageRunner externalStorage = + ExternalStorageRunner.create( + ExternalStorage.newBuilder().setDriver(driver).setPayloadSizeThreshold(0).build()); + PollWorkflowTaskQueueResponse fullTask = HistoryUtils.generateWorkflowTaskWithInitialHistory(); + HistoryEvent startedEvent = fullTask.getHistory().getEvents(0); + Payload input = Payload.newBuilder().setData(ByteString.copyFromUtf8("input")).build(); + History.Builder storedHistory = + fullTask.getHistory().toBuilder() + .setEvents( + 0, + startedEvent.toBuilder() + .setWorkflowExecutionStartedEventAttributes( + startedEvent.getWorkflowExecutionStartedEventAttributes().toBuilder() + .setInput(Payloads.newBuilder().addPayloads(input)))); + externalStorage.store(storedHistory, null, null, CancellationToken.none()); + driver.failRetrieves(1); + + WorkflowServiceStubs client = mock(WorkflowServiceStubs.class); + when(client.getServerCapabilities()) + .thenReturn(() -> GetSystemInfoResponse.Capabilities.newBuilder().build()); + + WorkflowTaskHandler taskHandler = + new ReplayWorkflowTaskHandler( + "namespace", + setUpMockWorkflowFactory(), + new WorkflowExecutorCache(10, new WorkflowRunLockManager(), new NoopScope()), + SingleWorkerOptions.newBuilder().setExternalStorageRunner(externalStorage).build(), + null, + Duration.ofSeconds(5), + client, + null); + + WorkflowTaskHandler.Result result = + taskHandler.handleWorkflowTask(fullTask.toBuilder().setHistory(storedHistory).build()); + + assertNotNull( + "a failed download must be reported rather than ending the task", result.getTaskFailed()); + assertTrue(result.getTaskFailed().hasFailure()); + assertTrue( + "the reported failure must say what went wrong, got: " + + result.getTaskFailed().getFailure().getMessage(), + result.getTaskFailed().getFailure().getMessage().contains("storage unavailable")); + } + + @Test + public void resolvesExternalStorageReferencesInFetchedFullHistory() throws Throwable { + ExternalStorageRunner externalStorage = + ExternalStorageRunner.create( + ExternalStorage.newBuilder() + .setDriver(TestStorageDriver.create()) + .setPayloadSizeThreshold(0) + .build()); + PollWorkflowTaskQueueResponse fullTask = HistoryUtils.generateWorkflowTaskWithInitialHistory(); + HistoryEvent startedEvent = fullTask.getHistory().getEvents(0); + Payload input = Payload.newBuilder().setData(ByteString.copyFromUtf8("input")).build(); + History.Builder storedHistory = + fullTask.getHistory().toBuilder() + .setEvents( + 0, + startedEvent.toBuilder() + .setWorkflowExecutionStartedEventAttributes( + startedEvent.getWorkflowExecutionStartedEventAttributes().toBuilder() + .setInput(Payloads.newBuilder().addPayloads(input)))); + externalStorage.store(storedHistory, null, null, CancellationToken.none()); + assertNotEquals( + "the payload must be replaced by a reference, otherwise this test proves nothing", + input, + storedInput(storedHistory)); + + WorkflowServiceStubs client = mock(WorkflowServiceStubs.class); + when(client.getServerCapabilities()) + .thenReturn(() -> GetSystemInfoResponse.Capabilities.newBuilder().build()); + WorkflowServiceGrpc.WorkflowServiceBlockingStub blockingStub = + mock(WorkflowServiceGrpc.WorkflowServiceBlockingStub.class); + when(client.blockingStub()).thenReturn(blockingStub); + when(blockingStub.withOption(any(), any())).thenReturn(blockingStub); + when(blockingStub.getWorkflowExecutionHistory(any())) + .thenReturn( + GetWorkflowExecutionHistoryResponse.newBuilder().setHistory(storedHistory).build()); + + ReplayWorkflow workflow = mock(ReplayWorkflow.class); + when(workflow.eventLoop()).thenReturn(true); + when(workflow.getOutput()).thenReturn(Optional.empty()); + WorkflowContext workflowContext = mock(WorkflowContext.class); + when(workflowContext.getRunningUpdateHandlers()).thenReturn(new HashMap<>()); + when(workflow.getWorkflowContext()).thenReturn(workflowContext); + ReplayWorkflowFactory workflowFactory = mock(ReplayWorkflowFactory.class); + when(workflowFactory.getWorkflow(any(), any())).thenReturn(workflow); + WorkflowTaskHandler taskHandler = + new ReplayWorkflowTaskHandler( + "namespace", + workflowFactory, + new WorkflowExecutorCache(10, new WorkflowRunLockManager(), new NoopScope()), + SingleWorkerOptions.newBuilder().setExternalStorageRunner(externalStorage).build(), + null, + Duration.ofSeconds(5), + client, + null); + + taskHandler.handleWorkflowTask( + fullTask.toBuilder().setHistory(History.getDefaultInstance()).build()); + + ArgumentCaptor event = ArgumentCaptor.forClass(HistoryEvent.class); + verify(workflow).start(event.capture(), any()); + assertEquals( + input, + event.getValue().getWorkflowExecutionStartedEventAttributes().getInput().getPayloads(0)); + } + @Test public void localActivityMeteringHelper() { ReplayWorkflowRunTaskHandler.LocalActivityMeteringHelper laMeteringHelper = diff --git a/temporal-sdk/src/test/java/io/temporal/internal/replay/ServiceWorkflowHistoryIteratorTest.java b/temporal-sdk/src/test/java/io/temporal/internal/replay/ServiceWorkflowHistoryIteratorTest.java index ad0c665800..672b0a7515 100644 --- a/temporal-sdk/src/test/java/io/temporal/internal/replay/ServiceWorkflowHistoryIteratorTest.java +++ b/temporal-sdk/src/test/java/io/temporal/internal/replay/ServiceWorkflowHistoryIteratorTest.java @@ -1,12 +1,23 @@ package io.temporal.internal.replay; import com.google.protobuf.ByteString; +import io.temporal.api.common.v1.Payload; +import io.temporal.api.common.v1.Payloads; import io.temporal.api.history.v1.History; +import io.temporal.api.history.v1.HistoryEvent; +import io.temporal.api.history.v1.WorkflowExecutionStartedEventAttributes; import io.temporal.api.workflowservice.v1.GetWorkflowExecutionHistoryResponse; import io.temporal.api.workflowservice.v1.PollWorkflowTaskQueueResponse; +import io.temporal.common.CancellationToken; +import io.temporal.internal.concurrent.structured.CancelSource; +import io.temporal.internal.payload.storage.ExternalStorageNotConfiguredException; +import io.temporal.internal.payload.storage.ExternalStorageRunner; +import io.temporal.internal.payload.storage.TestStorageDriver; +import io.temporal.payload.storage.ExternalStorage; import io.temporal.testUtils.HistoryUtils; import java.nio.charset.Charset; import java.util.NoSuchElementException; +import java.util.concurrent.CancellationException; import java.util.concurrent.atomic.AtomicInteger; import org.junit.Assert; import org.junit.Test; @@ -84,4 +95,91 @@ GetWorkflowExecutionHistoryResponse queryWorkflowExecutionHistory() { Assert.assertThrows(NoSuchElementException.class, iterator::next); Assert.assertEquals(4, timesCalledServer.get()); } + + @Test + public void resolvesExternalStorageReferencesInFetchedPages() { + ExternalStorageRunner storage = inMemoryStorage(); + History inline = historyWithInput(payload("big-input")); + History.Builder builder = inline.toBuilder(); + storage.store(builder, null, null, CancellationToken.none()); + History stored = builder.build(); + Assert.assertNotEquals( + "stored history should hold a reference, not the inline payload", inline, stored); + + ServiceWorkflowHistoryIterator iterator = fetchingIterator(stored, storage); + + HistoryEvent event = iterator.next(); + Assert.assertEquals( + payload("big-input"), + event.getWorkflowExecutionStartedEventAttributes().getInput().getPayloads(0)); + } + + @Test + public void failsLoudWhenAFetchedPageHasAReferenceAndStorageIsNotConfigured() { + History.Builder builder = historyWithInput(payload("big-input")).toBuilder(); + inMemoryStorage().store(builder, null, null, CancellationToken.none()); + History stored = builder.build(); + + ServiceWorkflowHistoryIterator iterator = fetchingIterator(stored, null); + + Assert.assertThrows(ExternalStorageNotConfiguredException.class, iterator::hasNext); + } + + @Test + public void aCancelledTokenAbortsRetrievalOfAFetchedPage() { + ExternalStorageRunner storage = inMemoryStorage(); + History.Builder builder = historyWithInput(payload("big-input")).toBuilder(); + storage.store(builder, null, null, CancellationToken.none()); + History stored = builder.build(); + + CancelSource source = + new CancelSource<>(() -> new CancellationException("Worker shutdown")); + source.cancel(); + + ServiceWorkflowHistoryIterator iterator = fetchingIterator(stored, storage, source.token()); + + Assert.assertThrows(CancellationException.class, iterator::hasNext); + } + + private static ServiceWorkflowHistoryIterator fetchingIterator( + History page, ExternalStorageRunner storage) { + return fetchingIterator(page, storage, CancellationToken.none()); + } + + private static ServiceWorkflowHistoryIterator fetchingIterator( + History page, + ExternalStorageRunner storage, + CancellationToken storageCancellation) { + PollWorkflowTaskQueueResponse workflowTask = + PollWorkflowTaskQueueResponse.newBuilder().setNextPageToken(NEXT_PAGE_TOKEN).build(); + return new ServiceWorkflowHistoryIterator( + null, "default", workflowTask, null, storage, storageCancellation) { + @Override + GetWorkflowExecutionHistoryResponse queryWorkflowExecutionHistory() { + return GetWorkflowExecutionHistoryResponse.newBuilder().setHistory(page).build(); + } + }; + } + + private static ExternalStorageRunner inMemoryStorage() { + return ExternalStorageRunner.create( + ExternalStorage.newBuilder() + .setDriver(TestStorageDriver.create()) + .setPayloadSizeThreshold(0) + .build()); + } + + private static History historyWithInput(Payload payload) { + return History.newBuilder() + .addEvents( + HistoryEvent.newBuilder() + .setWorkflowExecutionStartedEventAttributes( + WorkflowExecutionStartedEventAttributes.newBuilder() + .setInput(Payloads.newBuilder().addPayloads(payload)))) + .build(); + } + + private static Payload payload(String data) { + return Payload.newBuilder().setData(ByteString.copyFromUtf8(data)).build(); + } } diff --git a/temporal-sdk/src/test/java/io/temporal/internal/worker/WorkflowWorkerExternalStorageFailureTest.java b/temporal-sdk/src/test/java/io/temporal/internal/worker/WorkflowWorkerExternalStorageFailureTest.java new file mode 100644 index 0000000000..a1713e16e4 --- /dev/null +++ b/temporal-sdk/src/test/java/io/temporal/internal/worker/WorkflowWorkerExternalStorageFailureTest.java @@ -0,0 +1,120 @@ +package io.temporal.internal.worker; + +import io.temporal.api.enums.v1.EventType; +import io.temporal.client.WorkflowClient; +import io.temporal.client.WorkflowClientOptions; +import io.temporal.client.WorkflowOptions; +import io.temporal.internal.payload.storage.TestStorageDriver; +import io.temporal.payload.storage.ExternalStorage; +import io.temporal.testing.internal.SDKTestWorkflowRule; +import io.temporal.workflow.shared.TestWorkflows; +import java.time.Duration; +import java.util.UUID; +import org.junit.Assert; +import org.junit.Before; +import org.junit.Rule; +import org.junit.Test; + +public class WorkflowWorkerExternalStorageFailureTest { + + private static final TestStorageDriver driver = TestStorageDriver.named("wf-flaky"); + + private static final ExternalStorage storage = + ExternalStorage.newBuilder().setDriver(driver).setPayloadSizeThreshold(0).build(); + + @Rule + public SDKTestWorkflowRule testWorkflowRule = + SDKTestWorkflowRule.newBuilder() + .setWorkflowTypes(EchoWorkflowImpl.class) + .setWorkflowClientOptions( + WorkflowClientOptions.newBuilder().setExternalStorage(storage).build()) + .build(); + + @Before + public void resetDriver() { + driver.reset(); + } + + @Test + public void aFailedOutboundStoreFailsTheWorkflowTaskInsteadOfTimingOut() throws Exception { + String workflowId = "extstore-wft-" + UUID.randomUUID(); + String input = "wft-store-" + UUID.randomUUID(); + driver.failStoresContaining("echo: " + input, 1); + + TestWorkflows.TestWorkflow1 workflow = + testWorkflowRule + .getWorkflowClient() + .newWorkflowStub( + TestWorkflows.TestWorkflow1.class, + WorkflowOptions.newBuilder() + .setTaskQueue(testWorkflowRule.getTaskQueue()) + .setWorkflowId(workflowId) + .build()); + WorkflowClient.start(workflow::execute, input); + + awaitEvent(workflowId, EventType.EVENT_TYPE_WORKFLOW_EXECUTION_COMPLETED); + + Assert.assertEquals( + "expected exactly one injected store failure", 1, driver.injectedFailures.get()); + testWorkflowRule.assertHistoryEvent(workflowId, EventType.EVENT_TYPE_WORKFLOW_TASK_FAILED); + String reported = + testWorkflowRule + .getHistoryEvent(workflowId, EventType.EVENT_TYPE_WORKFLOW_TASK_FAILED) + .getWorkflowTaskFailedEventAttributes() + .getFailure() + .getMessage(); + Assert.assertTrue( + "the reported failure must say what went wrong, got: " + reported, + reported.contains("storage unavailable")); + Assert.assertTrue( + "a reported failure must not leave a workflow task timeout in history", + testWorkflowRule + .getHistoryEvents(workflowId, EventType.EVENT_TYPE_WORKFLOW_TASK_TIMED_OUT) + .isEmpty()); + } + + @Test + public void aStorageFailureIsReportedOnEveryAttemptNotJustTheFirst() throws Exception { + String workflowId = "extstore-wft-retry-" + UUID.randomUUID(); + String input = "wft-retry-" + UUID.randomUUID(); + driver.failStoresContaining("echo: " + input, 2); + + TestWorkflows.TestWorkflow1 workflow = + testWorkflowRule + .getWorkflowClient() + .newWorkflowStub( + TestWorkflows.TestWorkflow1.class, + WorkflowOptions.newBuilder() + .setTaskQueue(testWorkflowRule.getTaskQueue()) + .setWorkflowId(workflowId) + .build()); + WorkflowClient.start(workflow::execute, input); + + awaitEvent(workflowId, EventType.EVENT_TYPE_WORKFLOW_EXECUTION_COMPLETED); + + Assert.assertEquals("expected two injected store failures", 2, driver.injectedFailures.get()); + Assert.assertTrue( + "a reported failure must not leave a workflow task timeout in history", + testWorkflowRule + .getHistoryEvents(workflowId, EventType.EVENT_TYPE_WORKFLOW_TASK_TIMED_OUT) + .isEmpty()); + } + + private void awaitEvent(String workflowId, EventType eventType) throws InterruptedException { + long deadline = System.nanoTime() + Duration.ofSeconds(8).toNanos(); + while (System.nanoTime() < deadline) { + if (!testWorkflowRule.getHistoryEvents(workflowId, eventType).isEmpty()) { + return; + } + Thread.sleep(100); + } + Assert.fail("timed out waiting for " + eventType + " on " + workflowId); + } + + public static class EchoWorkflowImpl implements TestWorkflows.TestWorkflow1 { + @Override + public String execute(String input) { + return "echo: " + input; + } + } +} diff --git a/temporal-sdk/src/test/java/io/temporal/internal/worker/WorkflowWorkerTest.java b/temporal-sdk/src/test/java/io/temporal/internal/worker/WorkflowWorkerTest.java index 5cd1fc8d3e..e7ae93456e 100644 --- a/temporal-sdk/src/test/java/io/temporal/internal/worker/WorkflowWorkerTest.java +++ b/temporal-sdk/src/test/java/io/temporal/internal/worker/WorkflowWorkerTest.java @@ -3,23 +3,51 @@ import static java.nio.charset.StandardCharsets.UTF_8; import static junit.framework.TestCase.assertEquals; import static org.junit.Assert.*; +import static org.junit.Assert.assertFalse; +import static org.junit.Assert.assertNotEquals; import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.*; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import ch.qos.logback.classic.Level; +import ch.qos.logback.classic.LoggerContext; +import ch.qos.logback.classic.spi.ILoggingEvent; +import ch.qos.logback.core.read.ListAppender; import com.google.common.util.concurrent.Futures; import com.google.protobuf.ByteString; import com.uber.m3.tally.NoopScope; import com.uber.m3.tally.RootScopeBuilder; import com.uber.m3.tally.Scope; import com.uber.m3.util.ImmutableMap; +import io.temporal.api.command.v1.Command; +import io.temporal.api.command.v1.CompleteWorkflowExecutionCommandAttributes; +import io.temporal.api.command.v1.ContinueAsNewWorkflowExecutionCommandAttributes; +import io.temporal.api.command.v1.ScheduleActivityTaskCommandAttributes; +import io.temporal.api.command.v1.SignalExternalWorkflowExecutionCommandAttributes; +import io.temporal.api.command.v1.StartChildWorkflowExecutionCommandAttributes; +import io.temporal.api.common.v1.ActivityType; +import io.temporal.api.common.v1.Payload; +import io.temporal.api.common.v1.Payloads; import io.temporal.api.common.v1.WorkflowExecution; import io.temporal.api.common.v1.WorkflowType; +import io.temporal.api.failure.v1.ApplicationFailureInfo; +import io.temporal.api.failure.v1.Failure; import io.temporal.api.workflowservice.v1.*; +import io.temporal.api.workflowservice.v1.RespondQueryTaskCompletedRequest; +import io.temporal.api.workflowservice.v1.RespondWorkflowTaskFailedRequest; +import io.temporal.common.CancellationToken; import io.temporal.common.reporter.TestStatsReporter; import io.temporal.internal.common.InternalUtils; +import io.temporal.internal.concurrent.structured.CancelSource; +import io.temporal.internal.payload.storage.ExternalStorageRunner; +import io.temporal.internal.payload.storage.TestStorageDriver; import io.temporal.internal.replay.ReplayWorkflow; import io.temporal.internal.replay.ReplayWorkflowFactory; import io.temporal.internal.replay.ReplayWorkflowTaskHandler; +import io.temporal.payload.storage.ExternalStorage; +import io.temporal.payload.storage.StorageDriverTargetInfo; +import io.temporal.payload.storage.StorageDriverWorkflowInfo; import io.temporal.serviceclient.WorkflowServiceStubs; import io.temporal.testUtils.Eventually; import io.temporal.testUtils.HistoryUtils; @@ -31,7 +59,11 @@ import java.time.Duration; import java.util.UUID; import java.util.concurrent.*; +import java.util.concurrent.CancellationException; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.TimeUnit; import org.junit.Test; +import org.mockito.ArgumentCaptor; import org.mockito.stubbing.Answer; import org.slf4j.Logger; import org.slf4j.LoggerFactory; @@ -448,4 +480,565 @@ private ReplayWorkflowFactory setUpMockWorkflowFactory() throws Throwable { when(mockWorkflow.eventLoop()).thenReturn(false); return mockFactory; } + + @Test + public void aTaskAbandonedWhileShuttingDownIsNotReported() throws Exception { + WorkflowServiceStubs client = mock(WorkflowServiceStubs.class); + when(client.getServerCapabilities()) + .thenReturn(() -> GetSystemInfoResponse.Capabilities.newBuilder().build()); + WorkflowRunLockManager runLockManager = new WorkflowRunLockManager(); + Scope metricsScope = + new RootScopeBuilder() + .reporter(reporter) + .reportEvery(com.uber.m3.util.Duration.ofMillis(1)); + WorkflowExecutorCache cache = new WorkflowExecutorCache(10, runLockManager, metricsScope); + WorkflowTaskHandler taskHandler = mock(WorkflowTaskHandler.class); + when(taskHandler.isAnyTypeSupported()).thenReturn(true); + + CountDownLatch handlerEntered = new CountDownLatch(1); + CountDownLatch releaseHandler = new CountDownLatch(1); + CountDownLatch escaped = new CountDownLatch(1); + CancelSource storageCancellation = + new CancelSource<>(() -> new CancellationException("Worker shutdown")); + + WorkflowWorker worker = + new WorkflowWorker( + client, + "default", + "task_queue", + "sticky_task_queue", + SingleWorkerOptions.newBuilder() + .setIdentity("test_identity") + .setBuildId(UUID.randomUUID().toString()) + .setWorkerInstanceKey(UUID.randomUUID().toString()) + .setPollerOptions( + PollerOptions.newBuilder() + .setPollerBehavior(new PollerBehaviorSimpleMaximum(1)) + .setUncaughtExceptionHandler((thread, error) -> escaped.countDown()) + .build()) + .setMetricsScope(metricsScope) + .setStorageCancellation(storageCancellation.token()) + .build(), + runLockManager, + cache, + taskHandler, + mock(EagerActivityDispatcher.class), + 3, + new FixedSizeSlotSupplier<>(10), + new NamespaceCapabilities()); + + WorkflowServiceGrpc.WorkflowServiceFutureStub futureStub = + mock(WorkflowServiceGrpc.WorkflowServiceFutureStub.class); + when(futureStub.shutdownWorker(any(ShutdownWorkerRequest.class))) + .thenReturn(Futures.immediateFuture(ShutdownWorkerResponse.newBuilder().build())); + WorkflowServiceGrpc.WorkflowServiceBlockingStub blockingStub = + mock(WorkflowServiceGrpc.WorkflowServiceBlockingStub.class); + when(client.blockingStub()).thenReturn(blockingStub); + when(client.futureStub()).thenReturn(futureStub); + when(blockingStub.withOption(any(), any())).thenReturn(blockingStub); + + PollWorkflowTaskQueueResponse pollResponse = + PollWorkflowTaskQueueResponse.newBuilder() + .setTaskToken(ByteString.copyFrom("token", UTF_8)) + .setWorkflowExecution( + WorkflowExecution.newBuilder().setWorkflowId(WORKFLOW_ID).setRunId(RUN_ID).build()) + .setWorkflowType(WorkflowType.newBuilder().setName(WORKFLOW_TYPE).build()) + .build(); + CountDownLatch blockPolls = new CountDownLatch(1); + when(blockingStub.pollWorkflowTaskQueue(any(PollWorkflowTaskQueueRequest.class))) + .thenReturn(pollResponse) + .thenAnswer( + (Answer) + invocation -> { + blockPolls.await(); + return null; + }); + + // The task is abandoned part way through, which is what stopping storage looks like. + when(taskHandler.handleWorkflowTask(any(PollWorkflowTaskQueueResponse.class))) + .thenAnswer( + (Answer) + invocation -> { + handlerEntered.countDown(); + try { + releaseHandler.await(10, TimeUnit.SECONDS); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + throw new CancellationException("Worker shutdown"); + }); + + assertTrue(worker.start()); + assertTrue(handlerEntered.await(10, TimeUnit.SECONDS)); + storageCancellation.cancel(); + CompletableFuture shutdown = worker.shutdown(new ShutdownManager(), true); + releaseHandler.countDown(); + + assertFalse( + "abandoning a task while shutting down must not surface as an error", + escaped.await(2, TimeUnit.SECONDS)); + verify(blockingStub, never()) + .respondWorkflowTaskFailed(any(RespondWorkflowTaskFailedRequest.class)); + assertEquals( + "a task abandoned while shutting down must not count as a failed task", + 0, + worker.getTaskCounter().getTotalFailed()); + shutdown.get(); + } + + @Test + public void storageBreakingDuringAForcedShutdownIsStillReported() throws Exception { + // Cancelling storage means we abandoned the work. Storage genuinely breaking at the same + // moment is a different thing and must not disappear with it. + TestStorageDriver driver = TestStorageDriver.create().failStores(1); + Payload result = Payload.newBuilder().setData(ByteString.copyFrom("result", UTF_8)).build(); + RespondWorkflowTaskCompletedRequest taskCompleted = + RespondWorkflowTaskCompletedRequest.newBuilder() + .addCommands( + Command.newBuilder() + .setCompleteWorkflowExecutionCommandAttributes( + CompleteWorkflowExecutionCommandAttributes.newBuilder() + .setResult(Payloads.newBuilder().addPayloads(result)))) + .build(); + CancelSource storageCancellation = + new CancelSource<>(() -> new CancellationException("Worker shutdown")); + storageCancellation.cancel(); + + runOneTask( + driver, + new WorkflowTaskHandler.Result( + WORKFLOW_TYPE, taskCompleted, null, null, null, false, null, null), + storageCancellation.token(), + blockingStub -> + verify(blockingStub) + .respondWorkflowTaskFailed(any(RespondWorkflowTaskFailedRequest.class))); + + assertEquals("expected one injected store failure", 1, driver.injectedFailures.get()); + } + + @Test + public void payloadsInAFailedWorkflowTaskAreOffloaded() throws Exception { + TestStorageDriver driver = TestStorageDriver.create(); + Payload details = Payload.newBuilder().setData(ByteString.copyFrom("details", UTF_8)).build(); + RespondWorkflowTaskFailedRequest taskFailed = + RespondWorkflowTaskFailedRequest.newBuilder() + .setFailure( + Failure.newBuilder() + .setMessage("boom") + .setApplicationFailureInfo( + ApplicationFailureInfo.newBuilder() + .setDetails(Payloads.newBuilder().addPayloads(details)))) + .build(); + + ArgumentCaptor sent = + ArgumentCaptor.forClass(RespondWorkflowTaskFailedRequest.class); + runOneTask( + driver, + new WorkflowTaskHandler.Result( + WORKFLOW_TYPE, null, taskFailed, null, null, false, null, null), + blockingStub -> verify(blockingStub).respondWorkflowTaskFailed(sent.capture())); + + assertEquals("the failure details must be offloaded", 1, driver.storedCount()); + assertNotEquals( + "the failure details must be replaced by a reference", + details, + sent.getValue().getFailure().getApplicationFailureInfo().getDetails().getPayloads(0)); + } + + @Test + public void payloadsInADirectQueryResponseAreOffloaded() throws Exception { + TestStorageDriver driver = TestStorageDriver.create(); + Payload answer = Payload.newBuilder().setData(ByteString.copyFrom("answer", UTF_8)).build(); + RespondQueryTaskCompletedRequest queryCompleted = + RespondQueryTaskCompletedRequest.newBuilder() + .setQueryResult(Payloads.newBuilder().addPayloads(answer)) + .build(); + + ArgumentCaptor sent = + ArgumentCaptor.forClass(RespondQueryTaskCompletedRequest.class); + runOneTask( + driver, + new WorkflowTaskHandler.Result( + WORKFLOW_TYPE, null, null, queryCompleted, null, false, null, null), + blockingStub -> verify(blockingStub).respondQueryTaskCompleted(sent.capture())); + + assertEquals("the query answer must be offloaded", 1, driver.storedCount()); + assertNotEquals( + "the query answer must be replaced by a reference", + answer, + sent.getValue().getQueryResult().getPayloads(0)); + } + + /** Runs a single workflow task through a worker wired to {@code driver}, then verifies. */ + private void runOneTask( + TestStorageDriver driver, + WorkflowTaskHandler.Result handlerResult, + java.util.function.Consumer verification) + throws Exception { + runOneTask(driver, handlerResult, CancellationToken.none(), verification); + } + + private void runOneTask( + TestStorageDriver driver, + WorkflowTaskHandler.Result handlerResult, + CancellationToken storageCancellation, + java.util.function.Consumer verification) + throws Exception { + WorkflowServiceStubs client = mock(WorkflowServiceStubs.class); + when(client.getServerCapabilities()) + .thenReturn(() -> GetSystemInfoResponse.Capabilities.newBuilder().build()); + WorkflowRunLockManager runLockManager = new WorkflowRunLockManager(); + Scope metricsScope = + new RootScopeBuilder() + .reporter(reporter) + .reportEvery(com.uber.m3.util.Duration.ofMillis(1)); + WorkflowExecutorCache cache = new WorkflowExecutorCache(10, runLockManager, metricsScope); + WorkflowTaskHandler taskHandler = mock(WorkflowTaskHandler.class); + when(taskHandler.isAnyTypeSupported()).thenReturn(true); + + WorkflowWorker worker = + new WorkflowWorker( + client, + "default", + "task_queue", + "sticky_task_queue", + SingleWorkerOptions.newBuilder() + .setIdentity("test_identity") + .setBuildId(UUID.randomUUID().toString()) + .setWorkerInstanceKey(UUID.randomUUID().toString()) + .setPollerOptions( + PollerOptions.newBuilder() + .setPollerBehavior(new PollerBehaviorSimpleMaximum(1)) + .build()) + .setMetricsScope(metricsScope) + .setExternalStorageRunner( + ExternalStorageRunner.create( + ExternalStorage.newBuilder() + .setDriver(driver) + .setPayloadSizeThreshold(0) + .build())) + .setStorageCancellation(storageCancellation) + .build(), + runLockManager, + cache, + taskHandler, + mock(EagerActivityDispatcher.class), + 3, + new FixedSizeSlotSupplier<>(10), + new NamespaceCapabilities()); + + WorkflowServiceGrpc.WorkflowServiceFutureStub futureStub = + mock(WorkflowServiceGrpc.WorkflowServiceFutureStub.class); + when(futureStub.shutdownWorker(any(ShutdownWorkerRequest.class))) + .thenReturn(Futures.immediateFuture(ShutdownWorkerResponse.newBuilder().build())); + WorkflowServiceGrpc.WorkflowServiceBlockingStub blockingStub = + mock(WorkflowServiceGrpc.WorkflowServiceBlockingStub.class); + when(client.blockingStub()).thenReturn(blockingStub); + when(client.futureStub()).thenReturn(futureStub); + when(blockingStub.withOption(any(), any())).thenReturn(blockingStub); + + PollWorkflowTaskQueueResponse pollResponse = + PollWorkflowTaskQueueResponse.newBuilder() + .setTaskToken(ByteString.copyFrom("token", UTF_8)) + .setWorkflowExecution( + WorkflowExecution.newBuilder().setWorkflowId(WORKFLOW_ID).setRunId(RUN_ID).build()) + .setWorkflowType(WorkflowType.newBuilder().setName(WORKFLOW_TYPE).build()) + .build(); + CountDownLatch blockPolls = new CountDownLatch(1); + when(blockingStub.pollWorkflowTaskQueue(any(PollWorkflowTaskQueueRequest.class))) + .thenReturn(pollResponse) + .thenAnswer( + (Answer) + invocation -> { + blockPolls.await(); + return null; + }); + + CountDownLatch handled = new CountDownLatch(1); + when(taskHandler.handleWorkflowTask(any(PollWorkflowTaskQueueResponse.class))) + .thenAnswer( + (Answer) + invocation -> { + handled.countDown(); + return handlerResult; + }); + + assertTrue(worker.start()); + assertTrue(handled.await(10, TimeUnit.SECONDS)); + worker.shutdown(new ShutdownManager(), false).get(); + verification.accept(blockingStub); + } + + @Test + public void aStoreThatFailsWhileShuttingDownIsNotTreatedAsAProblem() throws Exception { + LoggerContext loggerContext = (LoggerContext) LoggerFactory.getILoggerFactory(); + ListAppender logs = new ListAppender<>(); + logs.setContext(loggerContext); + logs.start(); + ch.qos.logback.classic.Logger workerLog = + loggerContext.getLogger(WorkflowWorker.class.getName()); + workerLog.addAppender(logs); + try { + WorkflowServiceStubs client = mock(WorkflowServiceStubs.class); + when(client.getServerCapabilities()) + .thenReturn(() -> GetSystemInfoResponse.Capabilities.newBuilder().build()); + + WorkflowRunLockManager runLockManager = new WorkflowRunLockManager(); + Scope metricsScope = + new RootScopeBuilder() + .reporter(reporter) + .reportEvery(com.uber.m3.util.Duration.ofMillis(1)); + WorkflowExecutorCache cache = new WorkflowExecutorCache(10, runLockManager, metricsScope); + SlotSupplier slotSupplier = new FixedSizeSlotSupplier<>(10); + + WorkflowTaskHandler taskHandler = mock(WorkflowTaskHandler.class); + when(taskHandler.isAnyTypeSupported()).thenReturn(true); + + CountDownLatch storeEntered = new CountDownLatch(1); + CountDownLatch releaseStore = new CountDownLatch(1); + TestStorageDriver driver = + TestStorageDriver.create().blockStores(storeEntered, releaseStore).cancelStores(1); + CountDownLatch escaped = new CountDownLatch(1); + CancelSource storageCancellation = + new CancelSource<>(() -> new CancellationException("Worker shutdown")); + + WorkflowWorker worker = + new WorkflowWorker( + client, + "default", + "task_queue", + "sticky_task_queue", + SingleWorkerOptions.newBuilder() + .setIdentity("test_identity") + .setBuildId(UUID.randomUUID().toString()) + .setWorkerInstanceKey(UUID.randomUUID().toString()) + .setPollerOptions( + PollerOptions.newBuilder() + .setPollerBehavior(new PollerBehaviorSimpleMaximum(1)) + .setUncaughtExceptionHandler((thread, error) -> escaped.countDown()) + .build()) + .setMetricsScope(metricsScope) + .setExternalStorageRunner( + ExternalStorageRunner.create( + ExternalStorage.newBuilder() + .setDriver(driver) + .setPayloadSizeThreshold(0) + .build())) + .setStorageCancellation(storageCancellation.token()) + .build(), + runLockManager, + cache, + taskHandler, + mock(EagerActivityDispatcher.class), + 3, + slotSupplier, + new NamespaceCapabilities()); + + WorkflowServiceGrpc.WorkflowServiceFutureStub futureStub = + mock(WorkflowServiceGrpc.WorkflowServiceFutureStub.class); + when(futureStub.shutdownWorker(any(ShutdownWorkerRequest.class))) + .thenReturn(Futures.immediateFuture(ShutdownWorkerResponse.newBuilder().build())); + WorkflowServiceGrpc.WorkflowServiceBlockingStub blockingStub = + mock(WorkflowServiceGrpc.WorkflowServiceBlockingStub.class); + when(client.blockingStub()).thenReturn(blockingStub); + when(client.futureStub()).thenReturn(futureStub); + when(blockingStub.withOption(any(), any())).thenReturn(blockingStub); + + PollWorkflowTaskQueueResponse pollResponse = + PollWorkflowTaskQueueResponse.newBuilder() + .setTaskToken(ByteString.copyFrom("token", UTF_8)) + .setWorkflowExecution( + WorkflowExecution.newBuilder() + .setWorkflowId(WORKFLOW_ID) + .setRunId(RUN_ID) + .build()) + .setWorkflowType(WorkflowType.newBuilder().setName(WORKFLOW_TYPE).build()) + .build(); + CountDownLatch pollTaskQueueLatch = new CountDownLatch(1); + CountDownLatch blockPollTaskQueueLatch = new CountDownLatch(1); + when(blockingStub.pollWorkflowTaskQueue(any(PollWorkflowTaskQueueRequest.class))) + .thenReturn(pollResponse) + .thenAnswer( + (Answer) + invocation -> { + pollTaskQueueLatch.countDown(); + blockPollTaskQueueLatch.await(); + return null; + }); + + when(taskHandler.handleWorkflowTask(any(PollWorkflowTaskQueueResponse.class))) + .thenAnswer( + (Answer) + invocation -> + new WorkflowTaskHandler.Result( + WORKFLOW_TYPE, + RespondWorkflowTaskCompletedRequest.newBuilder() + .addCommands( + Command.newBuilder() + .setCompleteWorkflowExecutionCommandAttributes( + CompleteWorkflowExecutionCommandAttributes.newBuilder() + .setResult( + Payloads.newBuilder() + .addPayloads( + Payload.newBuilder() + .setData( + ByteString.copyFrom( + "result", UTF_8)))))) + .build(), + null, + null, + null, + false, + null, + null)); + + assertTrue(worker.start()); + assertTrue(storeEntered.await(10, TimeUnit.SECONDS)); + + CompletableFuture shutdown = worker.shutdown(new ShutdownManager(), true); + storageCancellation.cancel(); + releaseStore.countDown(); + + assertFalse( + "a store that fails while shutting down must not surface as an error on the task", + escaped.await(2, TimeUnit.SECONDS)); + verify(blockingStub, never()) + .respondWorkflowTaskFailed(any(RespondWorkflowTaskFailedRequest.class)); + shutdown.get(); + + assertFalse( + "shutting down must not be logged as a failure to report progress", + logs.list.stream() + .anyMatch( + event -> + event.getLevel() == Level.WARN + && event.getMessage().contains("Failure while reporting"))); + } finally { + workerLog.detachAppender(logs); + logs.stop(); + } + } + + /** One driver for these tests: stores in memory, and can block or fail on demand. */ + @Test + public void deriveStorageTargetPointsACompletionAtItsParent() { + StorageDriverTargetInfo child = new StorageDriverWorkflowInfo("ns", "child", "run-1", "Child"); + StorageDriverTargetInfo parent = + new StorageDriverWorkflowInfo("ns", "parent", "parent-run", null); + Command command = + Command.newBuilder() + .setCompleteWorkflowExecutionCommandAttributes( + CompleteWorkflowExecutionCommandAttributes.newBuilder()) + .build(); + + assertEquals(parent, WorkflowWorker.deriveStorageTarget("ns", child, command, parent)); + } + + @Test + public void deriveStorageTargetKeepsACompletionOnItselfWithoutAParent() { + StorageDriverTargetInfo self = + new StorageDriverWorkflowInfo("ns", "wf-1", "run-1", "MyWorkflow"); + Command command = + Command.newBuilder() + .setCompleteWorkflowExecutionCommandAttributes( + CompleteWorkflowExecutionCommandAttributes.newBuilder()) + .build(); + + assertEquals(self, WorkflowWorker.deriveStorageTarget("ns", self, command, null)); + } + + @Test + public void deriveStorageTargetKeepsActivityCommandsOnTheWorkflow() { + StorageDriverTargetInfo workflowDefault = + new StorageDriverWorkflowInfo("ns", "wf-1", "run-1", "MyWorkflow"); + Command command = + Command.newBuilder() + .setScheduleActivityTaskCommandAttributes( + ScheduleActivityTaskCommandAttributes.newBuilder() + .setActivityId("act-1") + .setActivityType(ActivityType.newBuilder().setName("MyActivity"))) + .build(); + + assertEquals( + workflowDefault, WorkflowWorker.deriveStorageTarget("ns", workflowDefault, command)); + } + + @Test + public void deriveStorageTargetPointsChildWorkflowCommandsAtTheChild() { + StorageDriverTargetInfo parent = + new StorageDriverWorkflowInfo("ns", "parent", "parent-run", "Parent"); + Command command = + Command.newBuilder() + .setStartChildWorkflowExecutionCommandAttributes( + StartChildWorkflowExecutionCommandAttributes.newBuilder() + .setWorkflowId("child-1") + .setWorkflowType(WorkflowType.newBuilder().setName("Child"))) + .build(); + + assertEquals( + new StorageDriverWorkflowInfo("ns", "child-1", null, "Child"), + WorkflowWorker.deriveStorageTarget("ns", parent, command)); + } + + @Test + public void deriveStorageTargetPointsSignalCommandsAtTheTargetWorkflow() { + StorageDriverTargetInfo self = new StorageDriverWorkflowInfo("ns", "self", "self-run", "Self"); + Command command = + Command.newBuilder() + .setSignalExternalWorkflowExecutionCommandAttributes( + SignalExternalWorkflowExecutionCommandAttributes.newBuilder() + .setExecution( + WorkflowExecution.newBuilder() + .setWorkflowId("other") + .setRunId("other-run"))) + .build(); + + assertEquals( + new StorageDriverWorkflowInfo("ns", "other", "other-run", null), + WorkflowWorker.deriveStorageTarget("ns", self, command)); + } + + @Test + public void deriveStorageTargetPointsContinueAsNewAtTheNewRun() { + StorageDriverTargetInfo current = + new StorageDriverWorkflowInfo("ns", "wf-1", "run-1", "CurrentWorkflow"); + Command command = + Command.newBuilder() + .setContinueAsNewWorkflowExecutionCommandAttributes( + ContinueAsNewWorkflowExecutionCommandAttributes.newBuilder() + .setWorkflowType(WorkflowType.newBuilder().setName("NextWorkflow"))) + .build(); + + assertEquals( + new StorageDriverWorkflowInfo("ns", "wf-1", null, "NextWorkflow"), + WorkflowWorker.deriveStorageTarget("ns", current, command)); + } + + @Test + public void deriveStorageTargetKeepsWorkflowTypeForContinueAsNewWithoutOverride() { + StorageDriverTargetInfo current = + new StorageDriverWorkflowInfo("ns", "wf-1", "run-1", "CurrentWorkflow"); + Command command = + Command.newBuilder() + .setContinueAsNewWorkflowExecutionCommandAttributes( + ContinueAsNewWorkflowExecutionCommandAttributes.newBuilder()) + .build(); + + assertEquals( + new StorageDriverWorkflowInfo("ns", "wf-1", null, "CurrentWorkflow"), + WorkflowWorker.deriveStorageTarget("ns", current, command)); + } + + @Test + public void deriveStorageTargetKeepsTheCurrentTargetForOtherCommands() { + StorageDriverTargetInfo current = + new StorageDriverWorkflowInfo("ns", "wf-1", "run-1", "MyWorkflow"); + Command command = + Command.newBuilder() + .setCompleteWorkflowExecutionCommandAttributes( + CompleteWorkflowExecutionCommandAttributes.newBuilder()) + .build(); + + assertSame(current, WorkflowWorker.deriveStorageTarget("ns", current, command)); + } }