diff --git a/src/InProcessTestHost/README.md b/src/InProcessTestHost/README.md index e583dea5..012af89a 100644 --- a/src/InProcessTestHost/README.md +++ b/src/InProcessTestHost/README.md @@ -70,6 +70,40 @@ PurgeResult purgeResult = await testHost.Client.PurgeAllInstancesAsync( creation-time range. You can provide `CreatedTo` without `CreatedFrom`. Explicit bounds are inclusive and are evaluated in UTC. +## Work-Item Completion Ownership + +The in-process sidecar assigns a fresh token to every activity and orchestration delivery +using the existing gRPC `WorkItem.completionToken` field. Completion is matched to that +delivery, not just its logical instance or activity task ID. Duplicate and stale completion +requests cannot settle a replacement delivery. + +**Migration for direct gRPC callers:** copy the received `WorkItem.completionToken` into +`ActivityResponse.completionToken` or `OrchestratorResponse.completionToken` on every +whole completion request. Keep the original instance ID and, for activities, the task ID +from that delivery. Missing tokens or mismatched identities return `InvalidArgument`; unknown, wrong-kind, or already +settled tokens return `NotFound`. A rejected request does not consume a valid delivery. +The SDK worker already echoes these tokens. SDK and protobuf implementations are unchanged. + +The test host accepts a single full orchestration completion response only. Deprecated partial +or chunked completion is not supported: `isPartial = true` or any present `chunkIndex` +(including zero on a final fragment) returns `InvalidArgument` without consuming the delivery +or retaining actions. Direct callers must send all actions in one response, with `isPartial` +false and `chunkIndex` omitted. The SDK's existing oversized-response chunking path is unchanged +and is not supported by this test host. + +Replay history streaming is separate and remains supported. A full response releases the +delivery's temporary replay snapshot, while an already captured history reader can still finish +independently. + +An accepted completion takes precedence over a later failure of that delivery's pending +stream write. A send failure without an accepted completion still propagates. Closing the +`GetWorkItems` stream does not implicitly settle already delivered work; completion can +arrive through an independent RPC. No timeout, heartbeat, or lease behavior is added. + +Explicit abandonment is unchanged: both Abandon RPCs only acknowledge requests and do not +validate tokens, cancel execution, or requeue work. Honoring abandonment is left to the +dependent follow-up [#814](https://github.com/microsoft/durabletask-dotnet/pull/814). + ## 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..ee0789ee 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,27 @@ 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 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); - } - } + 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 +578,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) + ValidateCompletionToken(request.CompletionToken); +#pragma warning disable CS0612 // Reject deprecated chunk fields instead of treating a fragment as a full response. + if (request.IsPartial || request.ChunkIndex is not null) #pragma warning restore CS0612 { - // 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)); - - return EmptyCompleteTaskResponse; + throw new RpcException(new Status( + StatusCode.InvalidArgument, + "Chunked orchestrator completion is not supported by the in-process test host. Send a single full response without chunk fields.")); } - // This is the final chunk (or a single non-chunked response) - if (this.partialOrchestratorChunks.TryRemove(request.InstanceId, out PartialOrchestratorChunk? existingPartialChunk)) + lock (this.pendingTasksLock) { - // We've been accumulating chunks - combine with final chunk (thread-safe) - existingPartialChunk.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.")); + } + + if (!string.Equals(pending.InstanceId, request.InstanceId, StringComparison.OrdinalIgnoreCase)) + { + throw new RpcException(new Status(StatusCode.InvalidArgument, "Completion token does not match the orchestration instance.")); + } - GrpcOrchestratorExecutionResult res = new() + List actions = request.Actions.Select(ProtobufUtils.ToOrchestratorAction).ToList(); + GrpcOrchestratorExecutionResult result = new() { - Actions = existingPartialChunk.AccumulatedActions, + Actions = 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); - + 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 +623,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.")); + } - tcs.TrySetResult(new ActivityExecutionResult { ResponseEvent = resultEvent }); - return EmptyCompleteTaskResponse; + 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)); + } + + this.pendingActivityTasks.Remove(request.CompletionToken); + pending.CompletionSource.SetResult(new ActivityExecutionResult { ResponseEvent = resultEvent }); + return EmptyCompleteTaskResponse; + } } /// @@ -818,10 +790,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 +834,35 @@ 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; + return await pending.CompletionSource.Task; } catch { - // Remove the TaskCompletionSource that we just created - this.RemoveOrchestratorTaskCompletionSource(instance.InstanceId); - throw; + 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) + { + this.pendingOrchestratorTasks.Remove(completionToken); + } + if (streamedPastEvents is not null) { this.streamingPastEvents.TryRemove( @@ -884,16 +873,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 +912,31 @@ async Task ITaskExecutor.ExecuteActivity(OrchestrationI } await this.SendWorkItemToClientAsync(workItem); + + // Closing the work-item stream does not settle an already delivered activity. + return await pending.CompletionSource.Task; } catch { - // Remove the TaskCompletionSource that we just created - this.RemoveActivityTaskCompletionSource(instance.InstanceId, activityEvent.EventId); - throw; - } + lock (this.pendingTasksLock) + { + if (this.pendingActivityTasks.Remove(completionToken) || + !pending.CompletionSource.Task.IsCompletedSuccessfully) + { + throw; + } + } - // 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; + // An accepted completion takes precedence over a late send failure. + return await pending.CompletionSource.Task; + } + finally + { + lock (this.pendingTasksLock) + { + this.pendingActivityTasks.Remove(completionToken); + } + } } async Task SendWorkItemToClientAsync(P.WorkItem workItem) @@ -985,36 +988,12 @@ async Task SendWorkItemToClientAsync(P.WorkItem workItem) } } - TaskCompletionSource CreateTaskCompletionSourceForOrchestrator(string instanceId) - { - TaskCompletionSource tcs = new(TaskCreationOptions.RunContinuationsAsynchronously); - this.pendingOrchestratorTasks.TryAdd(instanceId, tcs); - return tcs; - } - - void RemoveOrchestratorTaskCompletionSource(string instanceId) + static void ValidateCompletionToken(string completionToken) { - 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.")); + } } /// diff --git a/test/Grpc.IntegrationTests/AutochunkTests.cs b/test/Grpc.IntegrationTests/AutochunkTests.cs index 0bab2254..bb59e0f3 100644 --- a/test/Grpc.IntegrationTests/AutochunkTests.cs +++ b/test/Grpc.IntegrationTests/AutochunkTests.cs @@ -10,27 +10,26 @@ namespace Microsoft.DurableTask.Grpc.Tests; /// -/// Integration tests for validating autochunk functionality when orchestration completion responses -/// exceed the maximum chunk size and are automatically split into multiple chunks. +/// Integration tests for whole orchestration completion responses and single-action size validation. /// public class AutochunkTests(ITestOutputHelper output, GrpcSidecarFixture sidecarFixture) : IntegrationTestBase(output, sidecarFixture) { /// - /// Validates that orchestrations complete successfully when the completion response - /// exceeds the chunk size and must be split into multiple chunks. + /// Validates that multiple activity actions complete in one response below the configured size limit. /// [Fact] - public async Task Autochunk_MultipleChunks_CompletesSuccessfully() + public async Task FullResponse_MultipleActions_CompletesSuccessfully() { - const int ActivityCount = 36; + // Arrange + const int ActivityCount = 16; const int PayloadSizePerActivity = 30 * 1024; const int ChunkSize = GrpcDurableTaskWorkerOptions.MinCompleteOrchestrationWorkItemChunkSizeInBytes; // 1 MB (minimum allowed) - TaskName orchestratorName = nameof(Autochunk_MultipleChunks_CompletesSuccessfully); + TaskName orchestratorName = nameof(FullResponse_MultipleActions_CompletesSuccessfully); TaskName activityName = "Echo"; await using HostTestLifetime server = await this.StartWorkerAsync(b => { - // Set a small chunk size to force chunking + // Keep the response below the minimum configured chunk size. b.UseGrpc(opt => opt.CompleteOrchestrationWorkItemChunkSizeInBytes = ChunkSize); b.AddTasks(tasks => tasks .AddOrchestratorFunc(orchestratorName, async ctx => @@ -48,11 +47,13 @@ public async Task Autochunk_MultipleChunks_CompletesSuccessfully() .AddActivityFunc(activityName, (ctx, input) => Task.FromResult(input))); }); + // Act string instanceId = await server.Client.ScheduleNewOrchestrationInstanceAsync(orchestratorName); using CancellationTokenSource cts = new CancellationTokenSource(TimeSpan.FromSeconds(30)); OrchestrationMetadata metadata = await server.Client.WaitForInstanceCompletionAsync( instanceId, getInputsAndOutputs: true, cts.Token); + // Assert Assert.NotNull(metadata); Assert.Equal(instanceId, metadata.InstanceId); Assert.Equal(OrchestrationRuntimeStatus.Completed, metadata.RuntimeStatus); @@ -60,24 +61,24 @@ public async Task Autochunk_MultipleChunks_CompletesSuccessfully() } /// - /// Validates autochunking with mixed action types (activities, timers, sub-orchestrations). + /// Validates a full response with mixed action types (activities, timers, sub-orchestrations). /// [Fact] - public async Task Autochunk_MixedActions_CompletesSuccessfully() + public async Task FullResponse_MixedActions_CompletesSuccessfully() { - // Use minimum allowed chunk size (1 MB) and ensure total payload exceeds it to trigger chunking + // Arrange const int ActivityCount = 30; const int TimerCount = 100; const int SubOrchCount = 50; const int PayloadSizePerActivity = 20 * 1024; const int ChunkSize = GrpcDurableTaskWorkerOptions.MinCompleteOrchestrationWorkItemChunkSizeInBytes; // 1 MB (minimum allowed) - TaskName orchestratorName = nameof(Autochunk_MixedActions_CompletesSuccessfully); + TaskName orchestratorName = nameof(FullResponse_MixedActions_CompletesSuccessfully); TaskName activityName = "Echo"; TaskName subOrchName = "SubOrch"; await using HostTestLifetime server = await this.StartWorkerAsync(b => { - // Set a small chunk size to force chunking + // The mixed actions fit in a single response below the minimum configured chunk size. b.UseGrpc(opt => opt.CompleteOrchestrationWorkItemChunkSizeInBytes = ChunkSize); b.AddTasks(tasks => tasks .AddOrchestratorFunc(orchestratorName, async ctx => @@ -111,11 +112,13 @@ public async Task Autochunk_MixedActions_CompletesSuccessfully() .AddActivityFunc(activityName, (ctx, input) => Task.FromResult(input))); }); + // Act string instanceId = await server.Client.ScheduleNewOrchestrationInstanceAsync(orchestratorName); using CancellationTokenSource cts = new CancellationTokenSource(TimeSpan.FromSeconds(30)); OrchestrationMetadata metadata = await server.Client.WaitForInstanceCompletionAsync( instanceId, getInputsAndOutputs: true, cts.Token); + // Assert Assert.NotNull(metadata); Assert.Equal(instanceId, metadata.InstanceId); Assert.Equal(OrchestrationRuntimeStatus.Completed, metadata.RuntimeStatus); @@ -160,4 +163,3 @@ public async Task Autochunk_SingleActionExceedsChunkSize_CompletesWithFailedStat Assert.Equal("System.InvalidOperationException: A single orchestrator action of type ScheduleTask with id 0 exceeds the 1.00MB limit: 1.10MB. Enable large-payload externalization to Azure Blob Storage to support oversized actions.", metadata.FailureDetails.ToString()); } } - diff --git a/test/InProcessTestHost.Tests/WorkItemCompletionTests.cs b/test/InProcessTestHost.Tests/WorkItemCompletionTests.cs new file mode 100644 index 00000000..2de88d1b --- /dev/null +++ b/test/InProcessTestHost.Tests/WorkItemCompletionTests.cs @@ -0,0 +1,607 @@ +// 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 completion ownership of individual work-item deliveries. +/// +public class WorkItemCompletionTests +{ + static readonly TimeSpan Timeout = TimeSpan.FromSeconds(5); + + /// + /// 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); + Assert.False(nextExecution.IsCompleted); + await CompleteAsync(session.Server, next); + await nextExecution.WaitAsync(Timeout); + } + + /// + /// Duplicate completion cannot settle a replacement or an unrelated delivery. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task CompleteWorkItem_StaleTokenDoesNotSettleReplacementAsync(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); + 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(() => CompleteAsync(session.Server, first)); + + // Assert + Assert.Equal(StatusCode.NotFound, duplicate.StatusCode); + 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 completion tokens leave active deliveries untouched. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task CompleteWorkItem_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(); + P.WorkItem invalid = delivery.Clone(); + + // Act + invalid.CompletionToken = string.Empty; + RpcException empty = await Assert.ThrowsAsync(() => CompleteAsync(session.Server, invalid)); + invalid.CompletionToken = "unknown"; + RpcException unknown = await Assert.ThrowsAsync(() => CompleteAsync(session.Server, invalid)); + invalid.CompletionToken = other.CompletionToken; + RpcException wrongKind = await Assert.ThrowsAsync(() => CompleteAsync(session.Server, invalid)); + + // 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); + } + + /// + /// A mismatched logical identity does not consume a valid completion token. + /// + [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 invalid = delivery.Clone(); + if (activity) + { + invalid.ActivityRequest.OrchestrationInstance.InstanceId = "other"; + } + else + { + invalid.OrchestratorRequest.InstanceId = "other"; + } + + // Act + RpcException mismatch = await Assert.ThrowsAsync(() => CompleteAsync(session.Server, invalid)); + + // Assert + Assert.Equal(StatusCode.InvalidArgument, mismatch.StatusCode); + Assert.False(execution.IsCompleted); + if (activity) + { + invalid = delivery.Clone(); + invalid.ActivityRequest.TaskId++; + RpcException wrongTaskId = await Assert.ThrowsAsync(() => CompleteAsync(session.Server, invalid)); + Assert.Equal(StatusCode.InvalidArgument, wrongTaskId.StatusCode); + Assert.False(execution.IsCompleted); + } + + await CompleteAsync(session.Server, delivery); + await execution.WaitAsync(Timeout); + } + + /// + /// Concurrent final responses have exactly one winner and preserve that winner's result. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task CompleteWorkItem_OnlyOneResponseClaimsDeliveryAsync(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 first = Task.Run(() => CompleteAfterSignalAsync(start.Task, session.Server, delivery, "\"first\"")); + Task second = Task.Run(() => CompleteAfterSignalAsync(start.Task, session.Server, delivery, "\"second\"")); + + // Act + start.SetResult(); + StatusCode[] outcomes = await Task.WhenAll(first, second).WaitAsync(Timeout); + await execution.WaitAsync(Timeout); + + // Assert + Assert.Single(outcomes, status => status == StatusCode.OK); + Assert.Single(outcomes, status => status == StatusCode.NotFound); + string expected = outcomes[0] == StatusCode.OK ? "\"first\"" : "\"second\""; + if (activity) + { + ActivityExecutionResult result = await (Task)execution; + Assert.Equal(expected, Assert.IsType(result.ResponseEvent).Result); + } + else + { + GrpcOrchestratorExecutionResult result = await (Task)execution; + Assert.Equal(expected, result.CustomStatus); + } + } + + /// + /// Activity completion preserves 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.StartExecution(activity: true); + P.WorkItem delivery = await session.ReadAsync(); + + // Act + await session.Server.CompleteActivityTask(new() + { + InstanceId = delivery.ActivityRequest.OrchestrationInstance.InstanceId, + TaskId = delivery.ActivityRequest.TaskId, + CompletionToken = delivery.CompletionToken, + Result = "\"result\"", + FailureDetails = failed ? new() { ErrorType = "ExpectedFailure", ErrorMessage = "failed" } : null, + }, CreateContext()); + ActivityExecutionResult result = await ((Task)execution).WaitAsync(Timeout); + + // Assert + if (failed) + { + TaskFailedEvent failure = Assert.IsType(result.ResponseEvent); + Assert.Equal(delivery.ActivityRequest.TaskId, failure.TaskScheduledId); + Assert.Equal("ExpectedFailure", failure.FailureDetails?.ErrorType); + } + else + { + TaskCompletedEvent completed = Assert.IsType(result.ResponseEvent); + Assert.Equal(delivery.ActivityRequest.TaskId, completed.TaskScheduledId); + Assert.Equal("\"result\"", completed.Result); + } + } + + /// + /// Unsupported response fragments leave the delivery and its replay snapshot available for a full response. + /// + [Theory] + [InlineData(true, null)] + [InlineData(true, 0)] + [InlineData(false, 0)] + [InlineData(false, 1)] + public async Task CompleteOrchestrator_RejectsChunkedResponsesWithoutSettlingDeliveryAsync(bool partial, int? chunkIndex) + { + // Arrange + await using ServerSession session = new(); + Task execution = session.StartExecution(activity: false, historyPayload: 'x'); + P.WorkItem delivery = await session.ReadAsync(); + List snapshot = WorkerHistorySnapshotTestHelpers.GetSnapshots(session.Server)["instance"]; +#pragma warning disable CS0612 // Exercise rejection of the deprecated chunk fields. + P.OrchestratorResponse fragment = new() + { + InstanceId = delivery.OrchestratorRequest.InstanceId, + CompletionToken = delivery.CompletionToken, + IsPartial = partial, + ChunkIndex = chunkIndex, + Actions = { new P.OrchestratorAction { Id = 0, ScheduleTask = new() { Name = "Rejected" } } }, + }; +#pragma warning restore CS0612 + + // Act + RpcException unsupported = await Assert.ThrowsAsync(() => + session.Server.CompleteOrchestratorTask(fragment, CreateContext())); + + // Assert + Assert.Equal(StatusCode.InvalidArgument, unsupported.StatusCode); + Assert.Contains("Chunked orchestrator completion is not supported", unsupported.Status.Detail); + Assert.False(execution.IsCompleted); + Assert.Same(snapshot, WorkerHistorySnapshotTestHelpers.GetSnapshots(session.Server)["instance"]); + 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()); + 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)); + } + + /// + /// Failed-send cleanup cannot consume a replacement delivery or remove its replay snapshot. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task Dispatch_FailedSendDoesNotRemoveReplacementAsync(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 + { + Task nextExecution = session.StartExecution(activity, historyPayload: 'y'); + + // Act + releaseWrite.SetResult(); + await Assert.ThrowsAsync(() => firstExecution.WaitAsync(Timeout)); + P.WorkItem next = await session.ReadAsync(); + + // Assert + Assert.NotEqual(first.CompletionToken, next.CompletionToken); + RpcException lateCompletion = await Assert.ThrowsAsync(() => CompleteAsync(session.Server, first)); + Assert.Equal(StatusCode.NotFound, lateCompletion.StatusCode); + Assert.False(nextExecution.IsCompleted); + if (!activity) + { + Assert.Equal('y', WorkerHistorySnapshotTestHelpers.GetSnapshots(session.Server)["instance"][0].ExecutionStarted.Input![0]); + } + + await CompleteAsync(session.Server, next); + await nextExecution.WaitAsync(Timeout); + if (!activity) + { + GrpcOrchestratorExecutionResult result = await (Task)nextExecution; + Assert.Empty(result.Actions); + Assert.Equal("\"result\"", result.CustomStatus); + Assert.Empty(WorkerHistorySnapshotTestHelpers.GetSnapshots(session.Server)); + } + } + finally + { + releaseWrite.TrySetResult(); + } + } + + /// + /// 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(); + } + } + + /// + /// Disconnecting a work-item stream does not implicitly settle delivered work. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task Disconnect_DoesNotSettleDeliveredWorkAsync(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); + } + + /// + /// Abandon RPCs still only acknowledge requests without validating tokens or settling work. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task AbandonWorkItem_RemainsAcknowledgementOnlyAsync(bool activity) + { + // Arrange + await using ServerSession session = new(); + Task execution = session.StartExecution(activity); + P.WorkItem delivery = await session.ReadAsync(); + + // Act + foreach (string token in new[] { delivery.CompletionToken, string.Empty, "unknown", delivery.CompletionToken }) + { + if (activity) + { + await session.Server.AbandonTaskActivityWorkItem(new() { CompletionToken = token }, CreateContext()); + } + else + { + await session.Server.AbandonTaskOrchestratorWorkItem(new() { CompletionToken = token }, CreateContext()); + } + } + + // Assert + Assert.False(execution.IsCompleted); + await CompleteAsync(session.Server, delivery); + await execution.WaitAsync(Timeout); + } + + static Task CompleteAsync(TaskHubGrpcServer server, P.WorkItem delivery, string value = "\"result\"") => + delivery.ActivityRequest is { } activity + ? server.CompleteActivityTask(new() + { + InstanceId = activity.OrchestrationInstance.InstanceId, + TaskId = activity.TaskId, + CompletionToken = delivery.CompletionToken, + Result = value, + }, CreateContext()) + : server.CompleteOrchestratorTask(new() + { + InstanceId = delivery.OrchestratorRequest.InstanceId, + CompletionToken = delivery.CompletionToken, + CustomStatus = value, + }, CreateContext()); + + static async Task CompleteAfterSignalAsync(Task signal, TaskHubGrpcServer server, P.WorkItem delivery, string value) + { + await signal; + try + { + await CompleteAsync(server, delivery, value); + 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 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)), + ] + : []; + ITaskExecutor executor = this.Server; + Task execution = activity + ? executor.ExecuteActivity(instance, new TaskScheduledEvent(1, "Activity", string.Empty, null)) + : 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 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 removed after a send failure. + } + } + + foreach (Task execution in this.executions.Where(task => !task.IsCompleted)) + { + await execution.WaitAsync(Timeout); + } + } + 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..610af7c4 100644 --- a/test/InProcessTestHost.Tests/WorkerHistorySnapshotTests.cs +++ b/test/InProcessTestHost.Tests/WorkerHistorySnapshotTests.cs @@ -28,12 +28,10 @@ namespace InProcessTestHost.Tests; public class WorkerHistorySnapshotTests { /// - /// Keeps streamed history available through reads and partial responses, but not after the final response. + /// Keeps streamed history available through reads, but not after the full completion response. /// - [Theory] - [InlineData(false)] - [InlineData(true)] - public async Task ExecuteOrchestrator_FinalResponse_ReleasesSnapshot(bool partialResponse) + [Fact] + public async Task ExecuteOrchestrator_FinalResponse_ReleasesSnapshot() { // Arrange using CancellationTokenSource timeout = new(TimeSpan.FromSeconds(30)); @@ -49,26 +47,13 @@ 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), await ReadWorkerHistoryAsync(server, instance.InstanceId)); - if (partialResponse) - { -#pragma warning disable CS0612 // Exercise the legacy chunked-completion path. - await server.CompleteOrchestratorTask(new() - { - InstanceId = instance.InstanceId, - IsPartial = true, - Actions = { new P.OrchestratorAction { Id = 0, ScheduleTask = new() { Name = "First" } } }, - }, CreateContext()); -#pragma warning restore CS0612 - Assert.Equal(history.Select(ProtobufUtils.ToHistoryEventProto), - await ReadWorkerHistoryAsync(server, instance.InstanceId)); - } - Assert.False(episode.IsCompleted); Assert.Equal(1, WorkerHistorySnapshotTestHelpers.GetSnapshots(server).Count); @@ -76,13 +61,14 @@ 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()); GrpcOrchestratorExecutionResult result = await episode.WaitAsync(timeout.Token); // Assert - Assert.Equal(partialResponse ? 2 : 1, result.Actions.Count()); + Assert.Single(result.Actions); Assert.Equal("\"waiting for activity\"", result.CustomStatus); Assert.Equal(0, WorkerHistorySnapshotTestHelpers.GetSnapshots(server).Count); Assert.Empty(await ReadWorkerHistoryAsync(server, instance.InstanceId)); @@ -107,11 +93,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 +112,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 +140,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 +181,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 +204,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 +216,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 +232,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); }