From 7c6ab0de5c52a2de6c8a5a15631d2b2bc413a7b3 Mon Sep 17 00:00:00 2001 From: wangbill Date: Fri, 9 Oct 2026 09:57:21 -0400 Subject: [PATCH 1/2] Honor explicit abandonment in the in-process test host Assign a fresh completion token to each activity and orchestration delivery. Coordinate completion, partial responses, explicit abandonment, and cleanup using that delivery's ownership so stale requests cannot settle a replacement. Cancel explicitly abandoned executions so the existing dispatcher requeues them, while preserving the stream-disconnect policy. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: 4a00a0f1-0b30-44f9-83a2-27dad922625f --- src/InProcessTestHost/README.md | 17 + .../Sidecar/Grpc/TaskHubGrpcServer.cs | 289 ++++----- .../WorkItemAbandonmentIntegrationTests.cs | 292 +++++++++ .../WorkItemAbandonmentTests.cs | 606 ++++++++++++++++++ .../WorkerHistorySnapshotTests.cs | 51 +- 5 files changed, 1086 insertions(+), 169 deletions(-) create mode 100644 test/InProcessTestHost.Tests/WorkItemAbandonmentIntegrationTests.cs create mode 100644 test/InProcessTestHost.Tests/WorkItemAbandonmentTests.cs diff --git a/src/InProcessTestHost/README.md b/src/InProcessTestHost/README.md index e583dea5..de3e78d6 100644 --- a/src/InProcessTestHost/README.md +++ b/src/InProcessTestHost/README.md @@ -70,6 +70,23 @@ PurgeResult purgeResult = await testHost.Client.PurgeAllInstancesAsync( creation-time range. You can provide `CreatedTo` without `CreatedFrom`. Explicit bounds are inclusive and are evaluated in UTC. +## Explicit Work-Item Abandonment + +The in-process sidecar honors explicit worker abandonment of activity and orchestration +work items. Abandonment cancels that delivery's pending execution, allowing the existing +dispatcher to release and requeue the work item for another attempt. + +Each delivery has a fresh completion token in the existing gRPC `completionToken` field. +Completion and abandonment requests must echo that token. Missing tokens return +`InvalidArgument`; unknown or already settled tokens return `NotFound`. Duplicate +abandonment and late completion cannot affect a replacement delivery. Abandoning an +orchestration also discards its partial response actions and releases its temporary +worker history snapshot, without deleting committed history. + +Closing the `GetWorkItems` stream does **not** implicitly abandon work that was already +delivered. An activity can still finish and send its completion through an independent +RPC. This behavior does not add an activity timeout, heartbeat, or lease policy. + ## Dependency Injection When your activities depend on services, there are two approaches: diff --git a/src/InProcessTestHost/Sidecar/Grpc/TaskHubGrpcServer.cs b/src/InProcessTestHost/Sidecar/Grpc/TaskHubGrpcServer.cs index 8331789b..c2946e4c 100644 --- a/src/InProcessTestHost/Sidecar/Grpc/TaskHubGrpcServer.cs +++ b/src/InProcessTestHost/Sidecar/Grpc/TaskHubGrpcServer.cs @@ -3,7 +3,6 @@ using System.Collections.Concurrent; using System.Diagnostics; -using System.Globalization; using System.Linq; using DurableTask.Core; using DurableTask.Core.Command; @@ -33,30 +32,29 @@ public class TaskHubGrpcServer : P.TaskHubSidecarService.TaskHubSidecarServiceBa { static readonly Task EmptyCompleteTaskResponse = Task.FromResult(new P.CompleteTaskResponse()); - readonly ConcurrentDictionary> pendingOrchestratorTasks = new(StringComparer.OrdinalIgnoreCase); - readonly ConcurrentDictionary> pendingActivityTasks = new(StringComparer.OrdinalIgnoreCase); - readonly ConcurrentDictionary partialOrchestratorChunks = new(StringComparer.OrdinalIgnoreCase); + // Token lookup, partial responses, and settlement share the same delivery-ownership boundary. + readonly object pendingTasksLock = new(); + readonly Dictionary pendingOrchestratorTasks = new(StringComparer.Ordinal); + readonly Dictionary pendingActivityTasks = new(StringComparer.Ordinal); - /// - /// Helper class to accumulate partial orchestrator chunks. - /// - sealed class PartialOrchestratorChunk + sealed class PendingOrchestratorTask(string instanceId) { - readonly object lockObject = new(); + public string InstanceId { get; } = instanceId; - public TaskCompletionSource TaskCompletionSource { get; set; } = null!; - public List AccumulatedActions { get; } = new(); + public TaskCompletionSource CompletionSource { get; } = + new(TaskCreationOptions.RunContinuationsAsynchronously); - /// - /// Thread-safely adds actions to the accumulated actions list. - /// - public void AddActions(IEnumerable actions) - { - lock (this.lockObject) - { - this.AccumulatedActions.AddRange(actions); - } - } + public List Actions { get; } = new(); + } + + sealed class PendingActivityTask(string instanceId, int taskId) + { + public string InstanceId { get; } = instanceId; + + public int TaskId { get; } = taskId; + + public TaskCompletionSource CompletionSource { get; } = + new(TaskCreationOptions.RunContinuationsAsynchronously); } readonly ILogger log; @@ -582,73 +580,41 @@ static P.GetInstanceResponse CreateGetInstanceResponse(OrchestrationState state, /// Returns an empty ack back to the remote SDK that we've received the completion. public override Task CompleteOrchestratorTask(P.OrchestratorResponse request, ServerCallContext context) { -#pragma warning disable CS0612 // isPartial is deprecated but still required for chunked response wire compatibility. - if (request.IsPartial) -#pragma warning restore CS0612 + ValidateCompletionToken(request.CompletionToken); + lock (this.pendingTasksLock) { - // This is a partial chunk - accumulate actions but don't complete yet - PartialOrchestratorChunk partialChunk = this.partialOrchestratorChunks.GetOrAdd( - request.InstanceId, - _ => - { - // First chunk - get the TCS and initialize the partial chunk - if (!this.pendingOrchestratorTasks.TryGetValue(request.InstanceId, out TaskCompletionSource? tcs)) - { - throw new RpcException(new Status(StatusCode.NotFound, $"Orchestration with instance ID '{request.InstanceId}' not found")); - } - - return new PartialOrchestratorChunk - { - TaskCompletionSource = tcs, - }; - }); - - // Accumulate actions from this chunk (thread-safe) - partialChunk.AddActions(request.Actions.Select(ProtobufUtils.ToOrchestratorAction)); + if (!this.pendingOrchestratorTasks.TryGetValue(request.CompletionToken, out PendingOrchestratorTask? pending)) + { + throw new RpcException(new Status(StatusCode.NotFound, "Orchestrator delivery not found.")); + } - return EmptyCompleteTaskResponse; - } + if (!string.Equals(pending.InstanceId, request.InstanceId, StringComparison.OrdinalIgnoreCase)) + { + throw new RpcException(new Status(StatusCode.InvalidArgument, "Completion token does not match the orchestration instance.")); + } - // This is the final chunk (or a single non-chunked response) - if (this.partialOrchestratorChunks.TryRemove(request.InstanceId, out PartialOrchestratorChunk? existingPartialChunk)) - { - // We've been accumulating chunks - combine with final chunk (thread-safe) - existingPartialChunk.AddActions(request.Actions.Select(ProtobufUtils.ToOrchestratorAction)); + List actions = request.Actions.Select(ProtobufUtils.ToOrchestratorAction).ToList(); +#pragma warning disable CS0612 // isPartial is deprecated but still required for chunked response wire compatibility. + if (request.IsPartial) +#pragma warning restore CS0612 + { + pending.Actions.AddRange(actions); + return EmptyCompleteTaskResponse; + } - GrpcOrchestratorExecutionResult res = new() + GrpcOrchestratorExecutionResult result = new() { - Actions = existingPartialChunk.AccumulatedActions, + Actions = pending.Actions, CustomStatus = request.CustomStatus, + OrchestrationActivitySpanId = request.OrchestrationTraceContext?.SpanID, + OrchestrationActivityStartTime = request.OrchestrationTraceContext?.SpanStartTime?.ToDateTimeOffset(), }; - // Remove the TCS from pending tasks and complete it - this.pendingOrchestratorTasks.TryRemove(request.InstanceId, out _); - existingPartialChunk.TaskCompletionSource.TrySetResult(res); - + pending.Actions.AddRange(actions); + this.pendingOrchestratorTasks.Remove(request.CompletionToken); + pending.CompletionSource.SetResult(result); return EmptyCompleteTaskResponse; } - - // Single non-chunked response (no partial chunks) - if (!this.pendingOrchestratorTasks.TryRemove( - request.InstanceId, - out TaskCompletionSource? tcs)) - { - // TODO: Log? - // CA2201: Use specific exception types - throw new RpcException(new Status(StatusCode.NotFound, $"Orchestration not found")); - } - - GrpcOrchestratorExecutionResult result = new() - { - Actions = request.Actions.Select(ProtobufUtils.ToOrchestratorAction), - CustomStatus = request.CustomStatus, - OrchestrationActivitySpanId = request.OrchestrationTraceContext?.SpanID, - OrchestrationActivityStartTime = request.OrchestrationTraceContext?.SpanStartTime?.ToDateTimeOffset(), - }; - - tcs.TrySetResult(result); - - return EmptyCompleteTaskResponse; } /// @@ -659,31 +625,39 @@ static P.GetInstanceResponse CreateGetInstanceResponse(OrchestrationState state, /// Returns an empty ack back to the remote SDK that we've received the completion. public override Task CompleteActivityTask(P.ActivityResponse request, ServerCallContext context) { - string taskIdKey = GetTaskIdKey(request.InstanceId, request.TaskId); - if (!this.pendingActivityTasks.TryRemove(taskIdKey, out TaskCompletionSource? tcs)) + ValidateCompletionToken(request.CompletionToken); + lock (this.pendingTasksLock) { - // TODO: Log? - // CA2201: Use specific exception types - throw new RpcException(new Status(StatusCode.NotFound, $"Activity not found")); - } + if (!this.pendingActivityTasks.TryGetValue(request.CompletionToken, out PendingActivityTask? pending)) + { + throw new RpcException(new Status(StatusCode.NotFound, "Activity delivery not found.")); + } - HistoryEvent resultEvent; - if (request.FailureDetails == null) - { - resultEvent = new TaskCompletedEvent(-1, request.TaskId, request.Result); - } - else - { - resultEvent = new TaskFailedEvent( - eventId: -1, - taskScheduledId: request.TaskId, - reason: null, - details: null, - failureDetails: ProtobufUtils.GetFailureDetails(request.FailureDetails)); - } + if (!string.Equals(pending.InstanceId, request.InstanceId, StringComparison.OrdinalIgnoreCase) || + pending.TaskId != request.TaskId) + { + throw new RpcException(new Status(StatusCode.InvalidArgument, "Completion token does not match the activity instance and task ID.")); + } + + HistoryEvent resultEvent; + if (request.FailureDetails == null) + { + resultEvent = new TaskCompletedEvent(-1, request.TaskId, request.Result); + } + else + { + resultEvent = new TaskFailedEvent( + eventId: -1, + taskScheduledId: request.TaskId, + reason: null, + details: null, + failureDetails: ProtobufUtils.GetFailureDetails(request.FailureDetails)); + } - tcs.TrySetResult(new ActivityExecutionResult { ResponseEvent = resultEvent }); - return EmptyCompleteTaskResponse; + this.pendingActivityTasks.Remove(request.CompletionToken); + pending.CompletionSource.SetResult(new ActivityExecutionResult { ResponseEvent = resultEvent }); + return EmptyCompleteTaskResponse; + } } /// @@ -818,10 +792,13 @@ async Task ITaskExecutor.ExecuteOrchestrator( } : null; - // Create a task completion source that represents the async completion of the orchestrator execution. - // This must be done before we start the orchestrator execution. - TaskCompletionSource tcs = - this.CreateTaskCompletionSourceForOrchestrator(instance.InstanceId); + string completionToken = Guid.NewGuid().ToString("N"); + PendingOrchestratorTask pending = new(instance.InstanceId); + lock (this.pendingTasksLock) + { + this.pendingOrchestratorTasks.Add(completionToken, pending); + } + List? streamedPastEvents = null; try @@ -859,21 +836,21 @@ async Task ITaskExecutor.ExecuteOrchestrator( await this.SendWorkItemToClientAsync(new P.WorkItem { OrchestratorRequest = orkRequest, + CompletionToken = completionToken, }); // The TCS will be completed on the message stream handler when it gets a response back from the remote process // TODO: How should we handle timeouts if the remote process never sends a response? // Probably need to have a static timeout (e.g. 5 minutes). - return await tcs.Task; - } - catch - { - // Remove the TaskCompletionSource that we just created - this.RemoveOrchestratorTaskCompletionSource(instance.InstanceId); - throw; + return await pending.CompletionSource.Task; } finally { + lock (this.pendingTasksLock) + { + this.pendingOrchestratorTasks.Remove(completionToken); + } + if (streamedPastEvents is not null) { this.streamingPastEvents.TryRemove( @@ -884,16 +861,18 @@ await this.SendWorkItemToClientAsync(new P.WorkItem async Task ITaskExecutor.ExecuteActivity(OrchestrationInstance instance, TaskScheduledEvent activityEvent) { - // Create a task completion source that represents the async completion of the activity. - // This must be done before we start the activity execution. - TaskCompletionSource tcs = this.CreateTaskCompletionSourceForActivity( - instance.InstanceId, - activityEvent.EventId); + string completionToken = Guid.NewGuid().ToString("N"); + PendingActivityTask pending = new(instance.InstanceId, activityEvent.EventId); + lock (this.pendingTasksLock) + { + this.pendingActivityTasks.Add(completionToken, pending); + } try { P.WorkItem workItem = new() { + CompletionToken = completionToken, ActivityRequest = new P.ActivityRequest { Name = activityEvent.Name, @@ -921,19 +900,17 @@ async Task ITaskExecutor.ExecuteActivity(OrchestrationI } await this.SendWorkItemToClientAsync(workItem); + + // Explicit abandonment cancels this delivery; disconnecting the stream does not. + return await pending.CompletionSource.Task; } - catch + finally { - // Remove the TaskCompletionSource that we just created - this.RemoveActivityTaskCompletionSource(instance.InstanceId, activityEvent.EventId); - throw; + lock (this.pendingTasksLock) + { + this.pendingActivityTasks.Remove(completionToken); + } } - - // The TCS will be completed on the message stream handler when it gets a response back from the remote process. - // TODO: How should we handle timeouts if the remote process never sends a response? - // Probably need a timeout feature for activities and/or a heartbeat API that activities - // can use to signal that they're still running. - return await tcs.Task; } async Task SendWorkItemToClientAsync(P.WorkItem workItem) @@ -985,36 +962,12 @@ async Task SendWorkItemToClientAsync(P.WorkItem workItem) } } - TaskCompletionSource CreateTaskCompletionSourceForOrchestrator(string instanceId) + static void ValidateCompletionToken(string completionToken) { - TaskCompletionSource tcs = new(TaskCreationOptions.RunContinuationsAsynchronously); - this.pendingOrchestratorTasks.TryAdd(instanceId, tcs); - return tcs; - } - - void RemoveOrchestratorTaskCompletionSource(string instanceId) - { - this.pendingOrchestratorTasks.TryRemove(instanceId, out _); - this.partialOrchestratorChunks.TryRemove(instanceId, out _); - } - - TaskCompletionSource CreateTaskCompletionSourceForActivity(string instanceId, int taskId) - { - string taskIdKey = GetTaskIdKey(instanceId, taskId); - TaskCompletionSource tcs = new(TaskCreationOptions.RunContinuationsAsynchronously); - this.pendingActivityTasks.TryAdd(taskIdKey, tcs); - return tcs; - } - - void RemoveActivityTaskCompletionSource(string instanceId, int taskId) - { - string taskIdKey = GetTaskIdKey(instanceId, taskId); - this.pendingActivityTasks.TryRemove(taskIdKey, out _); - } - - static string GetTaskIdKey(string instanceId, int taskId) - { - return string.Concat(instanceId, "_", taskId.ToString(CultureInfo.InvariantCulture)); + if (string.IsNullOrEmpty(completionToken)) + { + throw new RpcException(new Status(StatusCode.InvalidArgument, "A completion token is required.")); + } } /// @@ -1025,7 +978,18 @@ static string GetTaskIdKey(string instanceId, int taskId) /// An abandon activity task response. public override Task AbandonTaskActivityWorkItem(P.AbandonActivityTaskRequest request, ServerCallContext context) { - return Task.FromResult(new()); + ValidateCompletionToken(request.CompletionToken); + lock (this.pendingTasksLock) + { + if (!this.pendingActivityTasks.Remove(request.CompletionToken, out PendingActivityTask? pending)) + { + throw new RpcException(new Status(StatusCode.NotFound, "Activity delivery not found.")); + } + + pending.CompletionSource.SetCanceled(); + } + + return Task.FromResult(new P.AbandonActivityTaskResponse()); } /// @@ -1036,7 +1000,18 @@ static string GetTaskIdKey(string instanceId, int taskId) /// An abandon orchestration task response. public override Task AbandonTaskOrchestratorWorkItem(P.AbandonOrchestrationTaskRequest request, ServerCallContext context) { - return Task.FromResult(new()); + ValidateCompletionToken(request.CompletionToken); + lock (this.pendingTasksLock) + { + if (!this.pendingOrchestratorTasks.Remove(request.CompletionToken, out PendingOrchestratorTask? pending)) + { + throw new RpcException(new Status(StatusCode.NotFound, "Orchestrator delivery not found.")); + } + + pending.CompletionSource.SetCanceled(); + } + + return Task.FromResult(new P.AbandonOrchestrationTaskResponse()); } /// diff --git a/test/InProcessTestHost.Tests/WorkItemAbandonmentIntegrationTests.cs b/test/InProcessTestHost.Tests/WorkItemAbandonmentIntegrationTests.cs new file mode 100644 index 00000000..9dd7766c --- /dev/null +++ b/test/InProcessTestHost.Tests/WorkItemAbandonmentIntegrationTests.cs @@ -0,0 +1,292 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Collections.Concurrent; +using System.Net; +using DurableTask.Core; +using Grpc.Core; +using Grpc.Core.Interceptors; +using Grpc.Net.Client; +using Microsoft.AspNetCore.Builder; +using Microsoft.AspNetCore.Hosting; +using Microsoft.AspNetCore.Hosting.Server; +using Microsoft.AspNetCore.Hosting.Server.Features; +using Microsoft.AspNetCore.Server.Kestrel.Core; +using Microsoft.DurableTask; +using Microsoft.DurableTask.Testing.Sidecar; +using Microsoft.DurableTask.Testing.Sidecar.Grpc; +using Microsoft.DurableTask.Worker; +using Microsoft.Extensions.DependencyInjection; +using Microsoft.Extensions.Hosting; +using Microsoft.Extensions.Logging; +using Xunit; +using Xunit.Abstractions; +using P = Microsoft.DurableTask.Protobuf; + +namespace InProcessTestHost.Tests; + +/// +/// Tests explicit abandonment and redelivery through localhost gRPC and the in-memory dispatcher. +/// +public class WorkItemAbandonmentIntegrationTests(ITestOutputHelper output) +{ + static readonly TimeSpan Timeout = TimeSpan.FromSeconds(15); + + /// + /// A rejecting worker releases one delivery, then a compatible SDK worker completes it and unrelated work. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task ExplicitAbandonment_CompatibleWorkerCompletesRedeliveryAsync(bool activity) + { + // Arrange + InMemoryOrchestrationService service = new(); + WorkItemObserver observer = new(); + using IHost sidecar = Host.CreateDefaultBuilder() + .ConfigureLogging(logging => logging.ClearProviders()) + .ConfigureWebHostDefaults(web => web + .UseKestrel(options => options.Listen(IPAddress.Loopback, 0, endpoint => endpoint.Protocols = HttpProtocols.Http2)) + .ConfigureServices(services => + { + services.AddSingleton(observer); + services.AddGrpc(options => options.Interceptors.Add()); + services.AddSingleton(service); + services.AddSingleton(service); + services.AddSingleton(); + }) + .Configure(app => + { + app.UseRouting(); + app.UseEndpoints(endpoints => endpoints.MapGrpcService()); + })) + .Build(); + await sidecar.StartAsync().WaitAsync(Timeout); + string address = Assert.Single(sidecar.Services.GetRequiredService() + .Features.Get()!.Addresses); + using GrpcChannel channel = GrpcChannel.ForAddress(address); + P.TaskHubSidecarService.TaskHubSidecarServiceClient client = new(channel); + using CancellationTokenSource connectionCancellation = new(); + using AsyncServerStreamingCall rejectingWorker = client.GetWorkItems( + new(), cancellationToken: connectionCancellation.Token); + IHost? compatibleWorker = null; + P.WorkItem? rejected = null; + + try + { + await client.StartInstanceAsync(new() + { + InstanceId = "rejected", + Name = "Workflow", + Version = "2", + Input = "\"input\"", + }).ResponseAsync.WaitAsync(Timeout); + Assert.True(await rejectingWorker.ResponseStream.MoveNext(default).WaitAsync(Timeout)); + rejected = rejectingWorker.ResponseStream.Current; + if (activity) + { + await client.CompleteOrchestratorTaskAsync(new() + { + InstanceId = "rejected", + CompletionToken = rejected.CompletionToken, + Actions = { new P.OrchestratorAction + { + Id = 0, + ScheduleTask = new() { Name = "Echo", Version = "2", Input = "\"input\"" }, + } }, + }).ResponseAsync.WaitAsync(Timeout); + Assert.True(await rejectingWorker.ResponseStream.MoveNext(default).WaitAsync(Timeout)); + rejected = rejectingWorker.ResponseStream.Current; + Assert.NotNull(rejected.ActivityRequest); + Assert.Equal("2", rejected.ActivityRequest.Version); + } + else + { + Assert.Equal("2", Assert.Single(rejected.OrchestratorRequest.NewEvents, + historyEvent => historyEvent.ExecutionStarted is not null).ExecutionStarted.Version); + } + + // Stop fetching before rejecting so the replacement is delivered only to the compatible worker. + connectionCancellation.Cancel(); + await observer.FirstStreamClosed.Task.WaitAsync(Timeout); + + // Act + if (activity) + { + await client.AbandonTaskActivityWorkItemAsync( + new() { CompletionToken = rejected.CompletionToken }).ResponseAsync.WaitAsync(Timeout); + } + else + { + await client.AbandonTaskOrchestratorWorkItemAsync( + new() { CompletionToken = rejected.CompletionToken }).ResponseAsync.WaitAsync(Timeout); + } + + await client.StartInstanceAsync(new() + { + InstanceId = "other", + Name = "Workflow", + Version = "2", + Input = "\"other\"", + }).ResponseAsync.WaitAsync(Timeout); + compatibleWorker = Host.CreateDefaultBuilder() + .ConfigureLogging(logging => logging.ClearProviders()) + .ConfigureServices(services => services.AddDurableTaskWorker(builder => + { + builder.UseGrpc(channel); + builder.UseVersioning(new() + { + Version = "2", + MatchStrategy = DurableTaskWorkerOptions.VersionMatchStrategy.Strict, + }); + builder.AddTasks(tasks => + { + tasks.AddOrchestratorFunc("Workflow", new TaskVersion("2"), + (context, input) => context.CallActivityAsync("Echo", input)); + tasks.AddActivityFunc("Echo", new TaskVersion("2"), + (_, input) => Task.FromResult($"completed {input}")); + }); + })) + .Build(); + await compatibleWorker.StartAsync().WaitAsync(Timeout); + using CancellationTokenSource waitCancellation = new(); + Task result = client.WaitForInstanceCompletionAsync( + new() { InstanceId = "rejected", GetInputsAndOutputs = true }, cancellationToken: waitCancellation.Token).ResponseAsync; + Task otherResult = client.WaitForInstanceCompletionAsync( + new() { InstanceId = "other", GetInputsAndOutputs = true }, cancellationToken: waitCancellation.Token).ResponseAsync; + P.GetInstanceResponse[] completed; + try + { + completed = await Task.WhenAll(result, otherResult).WaitAsync(Timeout); + } + finally + { + waitCancellation.Cancel(); + } + + // Assert + Assert.Equal(P.OrchestrationStatus.Completed, completed[0].OrchestrationState.OrchestrationStatus); + Assert.Equal("\"completed input\"", completed[0].OrchestrationState.Output); + Assert.Equal(P.OrchestrationStatus.Completed, completed[1].OrchestrationState.OrchestrationStatus); + Assert.Equal("\"completed other\"", completed[1].OrchestrationState.Output); + Assert.Equal(1, observer.AbandonmentCalls); + P.WorkItem[] deliveries = observer.Deliveries.Where(item => activity + ? item.ActivityRequest?.OrchestrationInstance.InstanceId == "rejected" + : item.OrchestratorRequest?.InstanceId == "rejected" && + item.OrchestratorRequest.NewEvents.Any(historyEvent => historyEvent.ExecutionStarted is not null)).ToArray(); + output.WriteLine("Delivered work: " + string.Join(", ", observer.Deliveries.Select(item => + item.ActivityRequest is { } request + ? $"activity:{request.OrchestrationInstance.InstanceId}:{request.TaskId}" + : $"orchestrator:{item.OrchestratorRequest.InstanceId}"))); + Assert.Equal(2, deliveries.Length); + Assert.All(deliveries, delivery => Assert.NotEmpty(delivery.CompletionToken)); + Assert.NotEqual(deliveries[0].CompletionToken, deliveries[1].CompletionToken); + } + finally + { + connectionCancellation.Cancel(); + if (rejected is not null) + { + try + { + if (activity && rejected.ActivityRequest is { } request) + { + await client.CompleteActivityTaskAsync(new() + { + InstanceId = request.OrchestrationInstance.InstanceId, + TaskId = request.TaskId, + CompletionToken = rejected.CompletionToken, + Result = "\"cleanup\"", + }).ResponseAsync.WaitAsync(Timeout); + } + else + { + await client.CompleteOrchestratorTaskAsync(new() + { + InstanceId = "rejected", + CompletionToken = rejected.CompletionToken, + Actions = { new P.OrchestratorAction + { + Id = 0, + CompleteOrchestration = new() { OrchestrationStatus = P.OrchestrationStatus.Completed, Result = "\"cleanup\"" }, + } }, + }).ResponseAsync.WaitAsync(Timeout); + } + } + catch (RpcException exception) when (exception.StatusCode == StatusCode.NotFound) + { + // The rejected delivery has already been settled. + } + } + + if (compatibleWorker is not null) + { + await compatibleWorker.StopAsync().WaitAsync(Timeout); + compatibleWorker.Dispose(); + } + + await sidecar.StopAsync().WaitAsync(Timeout); + } + } + + sealed class WorkItemObserver : Interceptor + { + readonly ConcurrentQueue deliveries = new(); + int abandonmentCalls; + + internal TaskCompletionSource FirstStreamClosed { get; } = new(TaskCreationOptions.RunContinuationsAsynchronously); + + internal IEnumerable Deliveries => this.deliveries; + + internal int AbandonmentCalls => Volatile.Read(ref this.abandonmentCalls); + + public override Task UnaryServerHandler( + TRequest request, ServerCallContext context, UnaryServerMethod continuation) + { + if (request is P.AbandonActivityTaskRequest or P.AbandonOrchestrationTaskRequest) + { + Interlocked.Increment(ref this.abandonmentCalls); + } + + return continuation(request, context); + } + + public override async Task ServerStreamingServerHandler( + TRequest request, IServerStreamWriter responseStream, ServerCallContext context, + ServerStreamingServerMethod continuation) + { + try + { + await continuation(request, new RecordingWriter(responseStream, item => + { + if (item is P.WorkItem delivery) + { + this.deliveries.Enqueue(delivery); + } + }), context); + } + finally + { + if (request is P.GetWorkItemsRequest) + { + this.FirstStreamClosed.TrySetResult(); + } + } + } + } + + sealed class RecordingWriter(IServerStreamWriter inner, Action record) : IServerStreamWriter + { + public WriteOptions? WriteOptions + { + get => inner.WriteOptions; + set => inner.WriteOptions = value; + } + + public async Task WriteAsync(T message) + { + await inner.WriteAsync(message); + record(message); + } + } +} diff --git a/test/InProcessTestHost.Tests/WorkItemAbandonmentTests.cs b/test/InProcessTestHost.Tests/WorkItemAbandonmentTests.cs new file mode 100644 index 00000000..e37aca18 --- /dev/null +++ b/test/InProcessTestHost.Tests/WorkItemAbandonmentTests.cs @@ -0,0 +1,606 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Threading.Channels; +using DurableTask.Core; +using DurableTask.Core.Command; +using DurableTask.Core.History; +using Grpc.Core; +using Microsoft.DurableTask.Testing.Sidecar; +using Microsoft.DurableTask.Testing.Sidecar.Dispatcher; +using Microsoft.DurableTask.Testing.Sidecar.Grpc; +using Microsoft.Extensions.Hosting; +using Microsoft.Extensions.Logging.Abstractions; +using Microsoft.Extensions.Options; +using Moq; +using Moq.Protected; +using Xunit; +using P = Microsoft.DurableTask.Protobuf; + +namespace InProcessTestHost.Tests; + +/// +/// Tests explicit abandonment and the ownership of individual work-item deliveries. +/// +public class WorkItemAbandonmentTests +{ + static readonly TimeSpan Timeout = TimeSpan.FromSeconds(5); + + /// + /// Explicit abandonment settles the execution awaited by the dispatcher. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task AbandonWorkItem_CancelsPendingExecutionAsync(bool activity) + { + // Arrange + await using ServerSession session = new(); + Task execution = session.StartExecution(activity); + P.WorkItem delivery = await session.ReadAsync(); + + // Act + await AbandonAsync(session.Server, delivery); + + // Assert + await Assert.ThrowsAnyAsync(() => execution.WaitAsync(Timeout)); + Assert.True(execution.IsCanceled); + } + + /// + /// Repeated deliveries of the same logical work item have distinct nonempty tokens. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task Dispatch_AssignsUniqueDeliveryTokensAsync(bool activity) + { + // Arrange + await using ServerSession session = new(); + Task firstExecution = session.StartExecution(activity); + P.WorkItem first = await session.ReadAsync(); + await CompleteAsync(session.Server, first); + await firstExecution.WaitAsync(Timeout); + + // Act + Task nextExecution = session.StartExecution(activity); + P.WorkItem next = await session.ReadAsync(); + + // Assert + Assert.NotEmpty(first.CompletionToken); + Assert.NotEmpty(next.CompletionToken); + Assert.NotEqual(first.CompletionToken, next.CompletionToken); + RpcException lateAbandon = await Assert.ThrowsAsync(() => AbandonAsync(session.Server, first)); + Assert.Equal(StatusCode.NotFound, lateAbandon.StatusCode); + Assert.False(nextExecution.IsCompleted); + await CompleteAsync(session.Server, next); + await nextExecution.WaitAsync(Timeout); + } + + /// + /// Stale abandonment and completion cannot affect a replacement or an unrelated delivery. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task AbandonWorkItem_StaleTokenDoesNotSettleRedeliveryAsync(bool activity) + { + // Arrange + await using ServerSession session = new(); + Task firstExecution = session.StartExecution(activity); + P.WorkItem first = await session.ReadAsync(); + await AbandonAsync(session.Server, first); + await Assert.ThrowsAnyAsync(() => firstExecution.WaitAsync(Timeout)); + Task replacementExecution = session.StartExecution(activity); + P.WorkItem replacement = await session.ReadAsync(); + Task otherExecution = session.StartExecution(activity, "other"); + P.WorkItem other = await session.ReadAsync(); + + // Act + RpcException duplicate = await Assert.ThrowsAsync(() => AbandonAsync(session.Server, first)); + RpcException lateCompletion = await Assert.ThrowsAsync(() => CompleteAsync(session.Server, first)); + + // Assert + Assert.Equal(StatusCode.NotFound, duplicate.StatusCode); + Assert.Equal(StatusCode.NotFound, lateCompletion.StatusCode); + Assert.NotEqual(first.CompletionToken, replacement.CompletionToken); + Assert.False(replacementExecution.IsCompleted); + Assert.False(otherExecution.IsCompleted); + await CompleteAsync(session.Server, replacement); + await CompleteAsync(session.Server, other); + await Task.WhenAll(replacementExecution, otherExecution).WaitAsync(Timeout); + } + + /// + /// Missing, unknown, and wrong-kind abandonment tokens leave active deliveries untouched. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task AbandonWorkItem_InvalidTokensAreRejectedAsync(bool activity) + { + // Arrange + await using ServerSession session = new(); + Task execution = session.StartExecution(activity); + P.WorkItem delivery = await session.ReadAsync(); + Task otherExecution = session.StartExecution(!activity, "other"); + P.WorkItem other = await session.ReadAsync(); + + // Act + RpcException empty = await Assert.ThrowsAsync(() => AbandonTokenAsync(session.Server, activity, string.Empty)); + RpcException unknown = await Assert.ThrowsAsync(() => AbandonTokenAsync(session.Server, activity, "unknown")); + RpcException wrongKind = await Assert.ThrowsAsync(() => AbandonTokenAsync(session.Server, activity, other.CompletionToken)); + + // Assert + Assert.Equal(StatusCode.InvalidArgument, empty.StatusCode); + Assert.Equal(StatusCode.NotFound, unknown.StatusCode); + Assert.Equal(StatusCode.NotFound, wrongKind.StatusCode); + Assert.False(execution.IsCompleted); + Assert.False(otherExecution.IsCompleted); + await CompleteAsync(session.Server, delivery); + await CompleteAsync(session.Server, other); + await Task.WhenAll(execution, otherExecution).WaitAsync(Timeout); + } + + /// + /// Completion checks both the delivery token and the logical work-item identity. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task CompleteWorkItem_InvalidIdentityDoesNotRemoveDeliveryAsync(bool activity) + { + // Arrange + await using ServerSession session = new(); + Task execution = session.StartExecution(activity); + P.WorkItem delivery = await session.ReadAsync(); + P.WorkItem missingToken = delivery.Clone(); + missingToken.CompletionToken = string.Empty; + P.WorkItem wrongInstance = delivery.Clone(); + if (activity) + { + wrongInstance.ActivityRequest.OrchestrationInstance.InstanceId = "other"; + } + else + { + wrongInstance.OrchestratorRequest.InstanceId = "other"; + } + + // Act + RpcException missing = await Assert.ThrowsAsync(() => CompleteAsync(session.Server, missingToken)); + RpcException mismatch = await Assert.ThrowsAsync(() => CompleteAsync(session.Server, wrongInstance)); + + // Assert + Assert.Equal(StatusCode.InvalidArgument, missing.StatusCode); + Assert.Equal(StatusCode.InvalidArgument, mismatch.StatusCode); + Assert.False(execution.IsCompleted); + if (activity) + { + P.WorkItem wrongTask = delivery.Clone(); + wrongTask.ActivityRequest.TaskId++; + RpcException wrongTaskId = await Assert.ThrowsAsync(() => CompleteAsync(session.Server, wrongTask)); + Assert.Equal(StatusCode.InvalidArgument, wrongTaskId.StatusCode); + Assert.False(execution.IsCompleted); + } + + await CompleteAsync(session.Server, delivery); + await execution.WaitAsync(Timeout); + } + + /// + /// Concurrent completion and abandonment have exactly one winner. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task CompleteAndAbandon_OnlyOneClaimSucceedsAsync(bool activity) + { + // Arrange + await using ServerSession session = new(); + Task execution = session.StartExecution(activity); + P.WorkItem delivery = await session.ReadAsync(); + TaskCompletionSource start = new(TaskCreationOptions.RunContinuationsAsynchronously); + Task completion = Task.Run(() => SettleAfterSignalAsync( + start.Task, () => CompleteAsync(session.Server, delivery))); + Task abandonment = Task.Run(() => SettleAfterSignalAsync( + start.Task, () => AbandonAsync(session.Server, delivery))); + + // Act + start.SetResult(); + StatusCode[] outcomes = await Task.WhenAll(completion, abandonment).WaitAsync(Timeout); + + // Assert + Assert.Single(outcomes, status => status == StatusCode.OK); + Assert.Single(outcomes, status => status == StatusCode.NotFound); + if (outcomes[0] == StatusCode.OK) + { + await execution.WaitAsync(Timeout); + } + else + { + await Assert.ThrowsAnyAsync(() => execution.WaitAsync(Timeout)); + } + } + + /// + /// Activity responses still carry successful results and application failures. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task CompleteActivityTask_PreservesResultAsync(bool failed) + { + // Arrange + await using ServerSession session = new(); + Task execution = session.Executor.ExecuteActivity( + new() { InstanceId = "instance", ExecutionId = "current" }, new TaskScheduledEvent(1, "Activity", string.Empty, null)); + P.WorkItem delivery = await session.ReadAsync(); + + // Act + await session.Server.CompleteActivityTask(new() + { + InstanceId = "instance", + TaskId = 1, + CompletionToken = delivery.CompletionToken, + Result = "\"result\"", + FailureDetails = failed ? new() { ErrorType = "ExpectedFailure", ErrorMessage = "failed" } : null, + }, CreateContext()); + ActivityExecutionResult result = await execution.WaitAsync(Timeout); + + // Assert + if (failed) + { + TaskFailedEvent failure = Assert.IsType(result.ResponseEvent); + Assert.Equal(1, failure.TaskScheduledId); + Assert.Equal("ExpectedFailure", failure.FailureDetails?.ErrorType); + } + else + { + TaskCompletedEvent completed = Assert.IsType(result.ResponseEvent); + Assert.Equal(1, completed.TaskScheduledId); + Assert.Equal("\"result\"", completed.Result); + } + + RpcException lateAbandon = await Assert.ThrowsAsync(() => AbandonAsync(session.Server, delivery)); + Assert.Equal(StatusCode.NotFound, lateAbandon.StatusCode); + } + + /// + /// Abandonment discards partial actions and the snapshot without contaminating the next episode. + /// + [Fact] + public async Task AbandonOrchestrator_ReleasesPartialResponseAndSnapshotAsync() + { + // Arrange + await using ServerSession session = new(); + Task firstExecution = session.StartExecution(activity: false, historyPayload: 'x'); + P.WorkItem first = await session.ReadAsync(); + Assert.True(first.OrchestratorRequest.RequiresHistoryStreaming); + Assert.NotEmpty(await session.ReadHistoryAsync()); + await AddPartialResponseAsync(session.Server, first, "Discarded"); + Assert.False(firstExecution.IsCompleted); + + // Act + await AbandonAsync(session.Server, first); + await Assert.ThrowsAnyAsync(() => firstExecution.WaitAsync(Timeout)); + + // Assert + Assert.Empty(WorkerHistorySnapshotTestHelpers.GetSnapshots(session.Server)); + Assert.Empty(await session.ReadHistoryAsync()); + Task nextExecution = session.StartExecution(activity: false, historyPayload: 'y'); + P.WorkItem next = await session.ReadAsync(); + List nextHistory = await session.ReadHistoryAsync(); + RpcException latePartial = await Assert.ThrowsAsync(() => AddPartialResponseAsync(session.Server, first, "Late")); + Assert.Equal(StatusCode.NotFound, latePartial.StatusCode); + Assert.Equal(nextHistory, await session.ReadHistoryAsync()); + await AddPartialResponseAsync(session.Server, next, "Current"); + await session.Server.CompleteOrchestratorTask(new() + { + InstanceId = next.OrchestratorRequest.InstanceId, + CompletionToken = next.CompletionToken, + CustomStatus = "\"current\"", + }, CreateContext()); + GrpcOrchestratorExecutionResult result = await ((Task)nextExecution).WaitAsync(Timeout); + Assert.Equal("Current", Assert.IsAssignableFrom(Assert.Single(result.Actions)).Name); + Assert.Equal("\"current\"", result.CustomStatus); + Assert.Empty(WorkerHistorySnapshotTestHelpers.GetSnapshots(session.Server)); + } + + /// + /// An already captured history reader survives abandonment without reading the replacement's snapshot. + /// + [Fact] + public async Task AbandonOrchestrator_PreservesCapturedReaderAndReplacementSnapshotAsync() + { + // Arrange + await using ServerSession session = new(); + Task firstExecution = session.StartExecution(activity: false, historyPayload: 'x'); + P.WorkItem first = await session.ReadAsync(); + List firstHistory = await session.ReadHistoryAsync(); + TaskCompletionSource firstChunk = new(TaskCreationOptions.RunContinuationsAsynchronously); + TaskCompletionSource releaseReader = new(TaskCreationOptions.RunContinuationsAsynchronously); + List chunks = new(); + Mock> writer = new(); + writer.Setup(stream => stream.WriteAsync(It.IsAny())).Returns(async chunk => + { + chunks.Add(chunk.Clone()); + if (chunks.Count == 1) + { + firstChunk.SetResult(); + await releaseReader.Task.WaitAsync(Timeout); + } + }); + Task reader = session.Server.StreamInstanceHistory( + new() { InstanceId = "instance", ForWorkItemProcessing = true }, writer.Object, CreateContext()); + + try + { + await firstChunk.Task.WaitAsync(Timeout); + + // Act + await AbandonAsync(session.Server, first); + await Assert.ThrowsAnyAsync(() => firstExecution.WaitAsync(Timeout)); + Task nextExecution = session.StartExecution(activity: false, historyPayload: 'y'); + P.WorkItem next = await session.ReadAsync(); + List nextHistory = await session.ReadHistoryAsync(); + releaseReader.SetResult(); + await reader.WaitAsync(Timeout); + + // Assert + Assert.Equal(2, chunks.Count); + Assert.Equal(firstHistory, chunks.SelectMany(chunk => chunk.Events)); + Assert.Equal(nextHistory, await session.ReadHistoryAsync()); + Assert.False(nextExecution.IsCompleted); + await CompleteAsync(session.Server, next); + await nextExecution.WaitAsync(Timeout); + Assert.Empty(WorkerHistorySnapshotTestHelpers.GetSnapshots(session.Server)); + } + finally + { + releaseReader.TrySetResult(); + await reader.WaitAsync(Timeout); + } + } + + /// + /// Send-failure cleanup from an abandoned delivery cannot erase the replacement's ownership. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task Dispatch_FailedOldSendDoesNotRemoveReplacementAsync(bool activity) + { + // Arrange + TaskCompletionSource releaseWrite = new(TaskCreationOptions.RunContinuationsAsynchronously); + int writes = 0; + await using ServerSession session = new(async _ => + { + if (Interlocked.Increment(ref writes) == 1) + { + await releaseWrite.Task.WaitAsync(Timeout); + throw new InvalidOperationException("Expected send failure"); + } + }); + Task firstExecution = session.StartExecution(activity, historyPayload: 'x'); + P.WorkItem first = await session.ReadAsync(); + + try + { + await AbandonAsync(session.Server, first); + Task nextExecution = session.StartExecution(activity, historyPayload: 'y'); + + // Act + releaseWrite.TrySetResult(); + await Assert.ThrowsAsync(() => firstExecution.WaitAsync(Timeout)); + P.WorkItem next = await session.ReadAsync(); + + // Assert + Assert.NotEqual(first.CompletionToken, next.CompletionToken); + Assert.False(nextExecution.IsCompleted); + if (!activity) + { + Assert.NotEmpty(await session.ReadHistoryAsync()); + Assert.Equal('y', WorkerHistorySnapshotTestHelpers.GetSnapshots(session.Server)["instance"][0].ExecutionStarted.Input![0]); + } + + await CompleteAsync(session.Server, next); + await nextExecution.WaitAsync(Timeout); + } + finally + { + releaseWrite.TrySetResult(); + } + } + + /// + /// A disconnected work-item stream does not implicitly abandon already delivered work. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task Disconnect_DoesNotAbandonDeliveredWorkAsync(bool activity) + { + // Arrange + await using ServerSession session = new(); + Task execution = session.StartExecution(activity); + P.WorkItem delivery = await session.ReadAsync(); + + // Act + await session.DisconnectAsync(); + + // Assert + Assert.False(execution.IsCompleted); + await CompleteAsync(session.Server, delivery); + await execution.WaitAsync(Timeout); + } + + static Task AbandonAsync(TaskHubGrpcServer server, P.WorkItem delivery) => + AbandonTokenAsync(server, delivery.ActivityRequest is not null, delivery.CompletionToken); + + static Task AbandonTokenAsync(TaskHubGrpcServer server, bool activity, string completionToken) => + activity + ? server.AbandonTaskActivityWorkItem(new() { CompletionToken = completionToken }, CreateContext()) + : server.AbandonTaskOrchestratorWorkItem(new() { CompletionToken = completionToken }, CreateContext()); + + static Task CompleteAsync(TaskHubGrpcServer server, P.WorkItem delivery) => + delivery.ActivityRequest is { } activity + ? server.CompleteActivityTask(new() + { + InstanceId = activity.OrchestrationInstance.InstanceId, + TaskId = activity.TaskId, + CompletionToken = delivery.CompletionToken, + Result = "\"result\"", + }, CreateContext()) + : server.CompleteOrchestratorTask(new() + { + InstanceId = delivery.OrchestratorRequest.InstanceId, + CompletionToken = delivery.CompletionToken, + }, CreateContext()); + + static Task AddPartialResponseAsync(TaskHubGrpcServer server, P.WorkItem delivery, string activityName) + { +#pragma warning disable CS0612 // Exercise legacy chunked responses with per-delivery ownership. + return server.CompleteOrchestratorTask(new() + { + InstanceId = delivery.OrchestratorRequest.InstanceId, + CompletionToken = delivery.CompletionToken, + IsPartial = true, + Actions = { new P.OrchestratorAction { Id = 0, ScheduleTask = new() { Name = activityName } } }, + }, CreateContext()); +#pragma warning restore CS0612 + } + + static async Task SettleAfterSignalAsync(Task signal, Func settle) + { + await signal; + try + { + await settle(); + return StatusCode.OK; + } + catch (RpcException exception) + { + return exception.StatusCode; + } + } + + static ServerCallContext CreateContext(CancellationToken cancellation = default) + { + Mock context = new(); + context.Protected().SetupGet("CancellationTokenCore").Returns(cancellation); + return context.Object; + } + + sealed class ServerSession : IAsyncDisposable + { + readonly CancellationTokenSource stopping = new(); + readonly CancellationTokenSource connectionCancellation = new(); + readonly Channel workItems = Channel.CreateUnbounded(); + readonly List deliveries = new(); + readonly List executions = new(); + readonly Task connection; + + internal ServerSession(Func? afterWrite = null) + { + InMemoryOrchestrationService service = new(); + this.Server = new( + Mock.Of(lifetime => lifetime.ApplicationStopping == this.stopping.Token), + NullLoggerFactory.Instance, service, service, Options.Create(new TaskHubGrpcServerOptions())); + Mock> writer = new(); + writer.Setup(w => w.WriteAsync(It.IsAny())).Returns(async item => + { + this.deliveries.Add(item); + await this.workItems.Writer.WriteAsync(item); + if (afterWrite is not null) + { + await afterWrite(item); + } + }); + this.connection = this.Server.GetWorkItems( + new() { Capabilities = { P.WorkerCapability.HistoryStreaming } }, + writer.Object, CreateContext(this.connectionCancellation.Token)); + } + + internal TaskHubGrpcServer Server { get; } + + internal ITaskExecutor Executor => this.Server; + + internal Task StartExecution(bool activity, string instanceId = "instance", char? historyPayload = null) + { + OrchestrationInstance instance = new() { InstanceId = instanceId, ExecutionId = "current" }; + HistoryEvent[] history = historyPayload is { } payload + ? [ + new ExecutionStartedEvent(-1, new string(payload, 600 * 1024)) { Name = "Orchestrator", Version = string.Empty, OrchestrationInstance = instance }, + new TaskScheduledEvent(0, "Activity", string.Empty, new string(payload, 600 * 1024)), + ] + : []; + Task execution = activity + ? this.Executor.ExecuteActivity(instance, new TaskScheduledEvent(1, "Activity", string.Empty, null)) + : this.Executor.ExecuteOrchestrator(instance, history, [ + new ExecutionStartedEvent(-1, null) { Name = "Orchestrator", Version = string.Empty, OrchestrationInstance = instance }, + ]); + this.executions.Add(execution); + return execution; + } + + internal Task ReadAsync() => this.workItems.Reader.ReadAsync().AsTask().WaitAsync(Timeout); + + internal async Task> ReadHistoryAsync() + { + List history = new(); + Mock> writer = new(); + writer.Setup(w => w.WriteAsync(It.IsAny())).Returns(chunk => + { + history.AddRange(chunk.Events); + return Task.CompletedTask; + }); + await this.Server.StreamInstanceHistory( + new() { InstanceId = "instance", ForWorkItemProcessing = true }, writer.Object, CreateContext()); + return history; + } + + internal async Task DisconnectAsync() + { + this.connectionCancellation.Cancel(); + await this.connection.WaitAsync(Timeout); + } + + public async ValueTask DisposeAsync() + { + try + { + foreach (P.WorkItem delivery in this.deliveries) + { + try + { + await CompleteAsync(this.Server, delivery); + } + catch (RpcException exception) when (exception.StatusCode == StatusCode.NotFound) + { + // Already completed or abandoned by the test. + } + } + + foreach (Task execution in this.executions.Where(task => !task.IsCompleted)) + { + try + { + await execution.WaitAsync(Timeout); + } + catch (OperationCanceledException) + { + // Explicitly abandoned executions are expected to be canceled. + } + } + } + finally + { + this.stopping.Cancel(); + await this.DisconnectAsync(); + this.Server.Dispose(); + this.connectionCancellation.Dispose(); + this.stopping.Dispose(); + } + } + } +} diff --git a/test/InProcessTestHost.Tests/WorkerHistorySnapshotTests.cs b/test/InProcessTestHost.Tests/WorkerHistorySnapshotTests.cs index 0950c8cc..e597189e 100644 --- a/test/InProcessTestHost.Tests/WorkerHistorySnapshotTests.cs +++ b/test/InProcessTestHost.Tests/WorkerHistorySnapshotTests.cs @@ -49,7 +49,8 @@ public async Task ExecuteOrchestrator_FinalResponse_ReleasesSnapshot(bool partia try { - P.OrchestratorRequest request = (await workItems.Reader.ReadAsync(timeout.Token)).OrchestratorRequest; + P.WorkItem delivery = await workItems.Reader.ReadAsync(timeout.Token); + P.OrchestratorRequest request = delivery.OrchestratorRequest; Assert.True(request.RequiresHistoryStreaming); Assert.Empty(request.PastEvents); Assert.Equal(history.Select(ProtobufUtils.ToHistoryEventProto), @@ -61,6 +62,7 @@ public async Task ExecuteOrchestrator_FinalResponse_ReleasesSnapshot(bool partia await server.CompleteOrchestratorTask(new() { InstanceId = instance.InstanceId, + CompletionToken = delivery.CompletionToken, IsPartial = true, Actions = { new P.OrchestratorAction { Id = 0, ScheduleTask = new() { Name = "First" } } }, }, CreateContext()); @@ -76,6 +78,7 @@ await server.CompleteOrchestratorTask(new() await server.CompleteOrchestratorTask(new() { InstanceId = instance.InstanceId, + CompletionToken = delivery.CompletionToken, CustomStatus = "\"waiting for activity\"", Actions = { new P.OrchestratorAction { Id = 1, ScheduleTask = new() { Name = "Next" } } }, }, CreateContext()); @@ -107,11 +110,18 @@ public async Task ExecuteOrchestrator_SendFailure_ReleasesOnlyFailedSnapshot(boo using TaskHubGrpcServer server = CreateServer(timeout.Token); Channel workItems = Channel.CreateUnbounded(); Mock> writer = new(); + string failedCompletionToken = string.Empty; writer.Setup(w => w.WriteAsync(It.IsAny())).Returns(workItem => - workItem.OrchestratorRequest.InstanceId == "failed" - ? Task.FromException(new InvalidOperationException( - streamClosed ? "The request is complete." : "Expected dispatch failure")) - : workItems.Writer.WriteAsync(workItem).AsTask()); + { + if (workItem.OrchestratorRequest.InstanceId == "failed") + { + failedCompletionToken = workItem.CompletionToken; + return Task.FromException(new InvalidOperationException( + streamClosed ? "The request is complete." : "Expected dispatch failure")); + } + + return workItems.Writer.WriteAsync(workItem).AsTask(); + }); Task connection = server.GetWorkItems( new() { Capabilities = { P.WorkerCapability.HistoryStreaming } }, writer.Object, CreateContext(timeout.Token)); @@ -119,7 +129,7 @@ public async Task ExecuteOrchestrator_SendFailure_ReleasesOnlyFailedSnapshot(boo OrchestrationInstance other = new() { InstanceId = "other", ExecutionId = "current" }; HistoryEvent[] otherHistory = CreateHistory(other); Task otherEpisode = executor.ExecuteOrchestrator(other, otherHistory, []); - await workItems.Reader.ReadAsync(timeout.Token); + P.WorkItem otherDelivery = await workItems.Reader.ReadAsync(timeout.Token); OrchestrationInstance failed = new() { InstanceId = "failed", ExecutionId = "current" }; try @@ -147,11 +157,19 @@ public async Task ExecuteOrchestrator_SendFailure_ReleasesOnlyFailedSnapshot(boo Assert.Equal(otherHistory.Select(ProtobufUtils.ToHistoryEventProto), await ReadWorkerHistoryAsync(server, other.InstanceId)); RpcException missing = await Assert.ThrowsAsync(() => - server.CompleteOrchestratorTask(new() { InstanceId = failed.InstanceId }, CreateContext())); + server.CompleteOrchestratorTask(new() + { + InstanceId = failed.InstanceId, + CompletionToken = failedCompletionToken, + }, CreateContext())); Assert.Equal(StatusCode.NotFound, missing.StatusCode); Assert.False(otherEpisode.IsCompleted); - await server.CompleteOrchestratorTask(new() { InstanceId = other.InstanceId }, CreateContext()); + await server.CompleteOrchestratorTask(new() + { + InstanceId = other.InstanceId, + CompletionToken = otherDelivery.CompletionToken, + }, CreateContext()); await otherEpisode.WaitAsync(timeout.Token); Assert.Equal(0, snapshots.Count); } @@ -180,7 +198,7 @@ public async Task StreamHistoryAsync_EpisodeCompletion_PreservesReaderAndNextSna OrchestrationInstance instance = new() { InstanceId = "instance", ExecutionId = "previous" }; HistoryEvent[] history = CreateHistory(instance); Task episode = executor.ExecuteOrchestrator(instance, history, []); - await workItems.Reader.ReadAsync(timeout.Token); + P.WorkItem delivery = await workItems.Reader.ReadAsync(timeout.Token); TaskCompletionSource firstChunkWritten = new(TaskCreationOptions.RunContinuationsAsynchronously); TaskCompletionSource releaseReader = new(TaskCreationOptions.RunContinuationsAsynchronously); List chunks = new(); @@ -203,7 +221,11 @@ public async Task StreamHistoryAsync_EpisodeCompletion_PreservesReaderAndNextSna await firstChunkWritten.Task.WaitAsync(timeout.Token); // Act - await server.CompleteOrchestratorTask(new() { InstanceId = instance.InstanceId }, CreateContext()); + await server.CompleteOrchestratorTask(new() + { + InstanceId = instance.InstanceId, + CompletionToken = delivery.CompletionToken, + }, CreateContext()); await episode.WaitAsync(timeout.Token); // Assert @@ -211,7 +233,8 @@ public async Task StreamHistoryAsync_EpisodeCompletion_PreservesReaderAndNextSna OrchestrationInstance next = new() { InstanceId = instance.InstanceId, ExecutionId = "current" }; HistoryEvent[] nextHistory = CreateHistory(next, 'y'); Task nextEpisode = executor.ExecuteOrchestrator(next, nextHistory, []); - P.OrchestratorRequest nextRequest = (await workItems.Reader.ReadAsync(timeout.Token)).OrchestratorRequest; + P.WorkItem nextDelivery = await workItems.Reader.ReadAsync(timeout.Token); + P.OrchestratorRequest nextRequest = nextDelivery.OrchestratorRequest; Assert.Equal(next.ExecutionId, nextRequest.ExecutionId); Assert.True(nextRequest.RequiresHistoryStreaming); Assert.Equal(nextHistory.Select(ProtobufUtils.ToHistoryEventProto), @@ -226,7 +249,11 @@ public async Task StreamHistoryAsync_EpisodeCompletion_PreservesReaderAndNextSna WorkerHistorySnapshotTestHelpers.GetSnapshots(server)[next.InstanceId]); Assert.False(nextEpisode.IsCompleted); - await server.CompleteOrchestratorTask(new() { InstanceId = next.InstanceId }, CreateContext()); + await server.CompleteOrchestratorTask(new() + { + InstanceId = next.InstanceId, + CompletionToken = nextDelivery.CompletionToken, + }, CreateContext()); await nextEpisode.WaitAsync(timeout.Token); Assert.Equal(0, WorkerHistorySnapshotTestHelpers.GetSnapshots(server).Count); } From 8e2f5fb5c1e341fb5596c0023f9d92ccfd85a26e Mon Sep 17 00:00:00 2001 From: wangbill Date: Fri, 9 Oct 2026 17:10:14 -0400 Subject: [PATCH 2/2] Preserve accepted completion after late test-host send failures Let send-failure cleanup claim delivery ownership under the same lock as completion and abandonment. When a successful completion already won, return its accepted result for both activity and orchestrator execution instead of requeueing it. Add deterministic completion-before-send-failure coverage for both dispatch paths and document the settlement policy. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: 5777cc4b-513e-47b7-a159-75ad5e8a7298 --- src/InProcessTestHost/README.md | 3 + .../Sidecar/Grpc/TaskHubGrpcServer.cs | 28 +++++++ .../WorkItemAbandonmentTests.cs | 74 +++++++++++++++++++ 3 files changed, 105 insertions(+) diff --git a/src/InProcessTestHost/README.md b/src/InProcessTestHost/README.md index de3e78d6..2c9bff33 100644 --- a/src/InProcessTestHost/README.md +++ b/src/InProcessTestHost/README.md @@ -83,6 +83,9 @@ abandonment and late completion cannot affect a replacement delivery. Abandoning orchestration also discards its partial response actions and releases its temporary worker history snapshot, without deleting committed history. +An accepted completion takes precedence over a later failure of that delivery's pending +stream write. The send failure cannot replace the accepted result or requeue completed work. + Closing the `GetWorkItems` stream does **not** implicitly abandon work that was already delivered. An activity can still finish and send its completion through an independent RPC. This behavior does not add an activity timeout, heartbeat, or lease policy. diff --git a/src/InProcessTestHost/Sidecar/Grpc/TaskHubGrpcServer.cs b/src/InProcessTestHost/Sidecar/Grpc/TaskHubGrpcServer.cs index c2946e4c..95363db0 100644 --- a/src/InProcessTestHost/Sidecar/Grpc/TaskHubGrpcServer.cs +++ b/src/InProcessTestHost/Sidecar/Grpc/TaskHubGrpcServer.cs @@ -844,6 +844,20 @@ await this.SendWorkItemToClientAsync(new P.WorkItem // Probably need to have a static timeout (e.g. 5 minutes). return await pending.CompletionSource.Task; } + catch + { + lock (this.pendingTasksLock) + { + if (this.pendingOrchestratorTasks.Remove(completionToken) || + !pending.CompletionSource.Task.IsCompletedSuccessfully) + { + throw; + } + } + + // An accepted completion takes precedence over a late send failure. + return await pending.CompletionSource.Task; + } finally { lock (this.pendingTasksLock) @@ -904,6 +918,20 @@ async Task ITaskExecutor.ExecuteActivity(OrchestrationI // Explicit abandonment cancels this delivery; disconnecting the stream does not. return await pending.CompletionSource.Task; } + catch + { + lock (this.pendingTasksLock) + { + if (this.pendingActivityTasks.Remove(completionToken) || + !pending.CompletionSource.Task.IsCompletedSuccessfully) + { + throw; + } + } + + // An accepted completion takes precedence over a late send failure. + return await pending.CompletionSource.Task; + } finally { lock (this.pendingTasksLock) diff --git a/test/InProcessTestHost.Tests/WorkItemAbandonmentTests.cs b/test/InProcessTestHost.Tests/WorkItemAbandonmentTests.cs index e37aca18..d202b1c7 100644 --- a/test/InProcessTestHost.Tests/WorkItemAbandonmentTests.cs +++ b/test/InProcessTestHost.Tests/WorkItemAbandonmentTests.cs @@ -412,6 +412,80 @@ public async Task Dispatch_FailedOldSendDoesNotRemoveReplacementAsync(bool activ } } + /// + /// A send failure cannot override a completion accepted while that send was still pending. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task Dispatch_CompletedDeliveryPreservesResultAfterSendFailureAsync(bool activity) + { + // Arrange + TaskCompletionSource releaseWrite = new(TaskCreationOptions.RunContinuationsAsynchronously); + await using ServerSession session = new(async _ => + { + await releaseWrite.Task.WaitAsync(Timeout); + throw new InvalidOperationException("Expected send failure"); + }); + Task execution = session.StartExecution(activity, historyPayload: 'x'); + P.WorkItem delivery = await session.ReadAsync(); + + try + { + // Act + if (activity) + { + await session.Server.CompleteActivityTask(new() + { + InstanceId = delivery.ActivityRequest.OrchestrationInstance.InstanceId, + TaskId = delivery.ActivityRequest.TaskId, + CompletionToken = delivery.CompletionToken, + Result = "\"accepted result\"", + }, CreateContext()); + } + else + { + await session.Server.CompleteOrchestratorTask(new() + { + InstanceId = delivery.OrchestratorRequest.InstanceId, + CompletionToken = delivery.CompletionToken, + CustomStatus = "\"accepted status\"", + Actions = { new P.OrchestratorAction + { + Id = 7, + ScheduleTask = new() { Name = "Accepted", Input = "\"accepted input\"" }, + } }, + }, CreateContext()); + } + + Assert.False(execution.IsCompleted); + releaseWrite.SetResult(); + + // Assert + if (activity) + { + ActivityExecutionResult result = await ((Task)execution).WaitAsync(Timeout); + TaskCompletedEvent completed = Assert.IsType(result.ResponseEvent); + Assert.Equal(delivery.ActivityRequest.TaskId, completed.TaskScheduledId); + Assert.Equal("\"accepted result\"", completed.Result); + } + else + { + GrpcOrchestratorExecutionResult result = await ((Task)execution).WaitAsync(Timeout); + ScheduleTaskOrchestratorAction action = Assert.IsAssignableFrom(Assert.Single(result.Actions)); + Assert.Equal(7, action.Id); + Assert.Equal("Accepted", action.Name); + Assert.Equal("\"accepted input\"", action.Input); + Assert.Equal("\"accepted status\"", result.CustomStatus); + Assert.Empty(WorkerHistorySnapshotTestHelpers.GetSnapshots(session.Server)); + } + } + finally + { + releaseWrite.TrySetResult(); + } + } + /// /// A disconnected work-item stream does not implicitly abandon already delivered work. ///