diff --git a/src/InProcessTestHost/README.md b/src/InProcessTestHost/README.md index 8d996ea0..e583dea5 100644 --- a/src/InProcessTestHost/README.md +++ b/src/InProcessTestHost/README.md @@ -54,7 +54,10 @@ Only the current execution is retained. `ContinueAsNew` replaces the previous ge history when the new generation commits. The underlying in-memory service also accepts an execution ID: null or empty selects the current execution, and a different execution ID returns no history (gRPC `NotFound`). Worker history streaming continues to use the dispatched -episode's replay snapshot rather than this management snapshot. +episode's replay snapshot rather than this management snapshot. These temporary worker snapshots +remain available until the episode's final response and are released before the dispatcher commits +that episode or starts its next generation. A failed dispatch also releases its snapshot; committed +history is retained independently until purge or generation replacement. ### 5. Purge completed instances diff --git a/src/InProcessTestHost/Sidecar/Grpc/TaskHubGrpcServer.cs b/src/InProcessTestHost/Sidecar/Grpc/TaskHubGrpcServer.cs index fe1a9fa4..8331789b 100644 --- a/src/InProcessTestHost/Sidecar/Grpc/TaskHubGrpcServer.cs +++ b/src/InProcessTestHost/Sidecar/Grpc/TaskHubGrpcServer.cs @@ -822,6 +822,7 @@ async Task ITaskExecutor.ExecuteOrchestrator( // This must be done before we start the orchestrator execution. TaskCompletionSource tcs = this.CreateTaskCompletionSourceForOrchestrator(instance.InstanceId); + List? streamedPastEvents = null; try { @@ -845,7 +846,8 @@ async Task ITaskExecutor.ExecuteOrchestrator( if (this.supportsHistoryStreaming && totalBytes > HistoryStreamingThresholdBytes) { orkRequest.RequiresHistoryStreaming = true; - // Store past events to serve via StreamInstanceHistory + // Keep this episode's replay snapshot available until execution finishes. + streamedPastEvents = protoPastEvents; this.streamingPastEvents[instance.InstanceId] = protoPastEvents; } else @@ -858,6 +860,11 @@ await this.SendWorkItemToClientAsync(new P.WorkItem { OrchestratorRequest = orkRequest, }); + + // 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 { @@ -865,11 +872,14 @@ await this.SendWorkItemToClientAsync(new P.WorkItem this.RemoveOrchestratorTaskCompletionSource(instance.InstanceId); 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 to have a static timeout (e.g. 5 minutes). - return await tcs.Task; + finally + { + if (streamedPastEvents is not null) + { + this.streamingPastEvents.TryRemove( + new KeyValuePair>(instance.InstanceId, streamedPastEvents)); + } + } } async Task ITaskExecutor.ExecuteActivity(OrchestrationInstance instance, TaskScheduledEvent activityEvent) diff --git a/test/InProcessTestHost.Tests/OrchestrationHistoryTests.cs b/test/InProcessTestHost.Tests/OrchestrationHistoryTests.cs index 7dec7648..76ecb240 100644 --- a/test/InProcessTestHost.Tests/OrchestrationHistoryTests.cs +++ b/test/InProcessTestHost.Tests/OrchestrationHistoryTests.cs @@ -15,6 +15,7 @@ using Xunit; using Xunit.Abstractions; using P = Microsoft.DurableTask.Protobuf; +using PurgeResult = Microsoft.DurableTask.Client.PurgeResult; namespace InProcessTestHost.Tests; @@ -132,6 +133,76 @@ public async Task GetHistoryAsync_RunningAndCompleted_ReturnsFreshSnapshots(int Assert.Single(completed.OfType()); } + /// + /// Reuses one host without retaining completed episodes' replay snapshots or losing committed history. + /// + [Theory] + [InlineData(32, false)] + [InlineData(600 * 1024, false)] + [InlineData(600 * 1024, true)] + public async Task GetHistoryAsync_ReusedHost_ReleasesWorkerSnapshots(int payloadSize, bool continueAsNew) + { + // Arrange + using CancellationTokenSource timeout = new(TimeSpan.FromSeconds(30)); + HistoryRequestInterceptor interceptor = new(); + await using DurableTaskTestHost host = await DurableTaskTestHost.StartAsync(tasks => + { + tasks.AddOrchestratorFunc("PayloadLength", async (context, input) => + { + int length = await context.CallActivityAsync("Length", input); + if (continueAsNew && input[0] == 'x') + { + context.ContinueAsNew(new string('y', input.Length)); + } + + return length; + }); + tasks.AddActivityFunc("Length", (context, input) => input.Length); + }, new DurableTaskTestHostOptions + { + ConfigureServices = services => services.Configure( + options => options.Interceptors.Add(interceptor)), + }, timeout.Token); + string[] instanceIds = new string[2]; + + for (int i = 0; i < instanceIds.Length; i++) + { + // Act + string instanceId = await host.Client.ScheduleNewOrchestrationInstanceAsync( + "PayloadLength", new string('x', payloadSize), cancellation: timeout.Token); + instanceIds[i] = instanceId; + OrchestrationMetadata metadata = await host.Client.WaitForInstanceCompletionAsync( + instanceId, getInputsAndOutputs: true, cancellation: timeout.Token); + IList history = await host.Client.GetOrchestrationHistoryAsync(instanceId, timeout.Token); + + // Assert + Assert.Equal(OrchestrationRuntimeStatus.Completed, metadata.RuntimeStatus); + Assert.Equal(payloadSize, metadata.ReadOutputAs()); + Assert.Equal(8, history.Count); + Assert.StartsWith(continueAsNew ? "\"y" : "\"x", Assert.Single(history.OfType()).Input); + Assert.Single(history.OfType()); + Assert.Single(history.OfType()); + int pastEventBytes = history.Take(4).Sum(e => ProtobufUtils.ToHistoryEventProto(e).CalculateSize()); + bool streamsHistory = payloadSize > 32; + int streamsPerInstance = streamsHistory ? (continueAsNew ? 2 : 1) : 0; + this.output.WriteLine( + $"Instance {i + 1}: past-event protobuf size {pastEventBytes} bytes; worker history requests {interceptor.HistoryRequestCount}"); + Assert.Equal(streamsHistory, pastEventBytes > 1024 * 1024); + Assert.Equal(streamsPerInstance * (i + 1), interceptor.HistoryRequestCount); + Assert.Equal(0, WorkerHistorySnapshotTestHelpers.GetSnapshots(host).Count); + + PurgeResult purge = await host.Client.PurgeInstanceAsync(instanceId, cancellation: timeout.Token); + Assert.Equal(1, purge.PurgedInstanceCount); + Assert.Null(await host.Client.GetInstanceAsync(instanceId, cancellation: timeout.Token)); + ArgumentException missing = await Assert.ThrowsAsync(() => + host.Client.GetOrchestrationHistoryAsync(instanceId, timeout.Token)); + Assert.Equal(StatusCode.NotFound, Assert.IsType(missing.InnerException).StatusCode); + Assert.Equal(0, WorkerHistorySnapshotTestHelpers.GetSnapshots(host).Count); + } + + Assert.NotEqual(instanceIds[0], instanceIds[1]); + } + [Fact] public async Task GetHistoryAsync_ContinueAsNew_ReturnsOnlyCurrentGeneration() { diff --git a/test/InProcessTestHost.Tests/WorkerHistorySnapshotTests.cs b/test/InProcessTestHost.Tests/WorkerHistorySnapshotTests.cs new file mode 100644 index 00000000..0950c8cc --- /dev/null +++ b/test/InProcessTestHost.Tests/WorkerHistorySnapshotTests.cs @@ -0,0 +1,312 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Collections.Concurrent; +using System.Reflection; +using System.Threading.Channels; +using DurableTask.Core; +using DurableTask.Core.History; +using Grpc.Core; +using Microsoft.DurableTask.Testing; +using Microsoft.DurableTask.Testing.Sidecar; +using Microsoft.DurableTask.Testing.Sidecar.Dispatcher; +using Microsoft.DurableTask.Testing.Sidecar.Grpc; +using Microsoft.Extensions.DependencyInjection; +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 the lifetime of replay snapshots owned by dispatched orchestration episodes. +/// +public class WorkerHistorySnapshotTests +{ + /// + /// Keeps streamed history available through reads and partial responses, but not after the final response. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task ExecuteOrchestrator_FinalResponse_ReleasesSnapshot(bool partialResponse) + { + // Arrange + using CancellationTokenSource timeout = new(TimeSpan.FromSeconds(30)); + using TaskHubGrpcServer server = CreateServer(timeout.Token); + Channel workItems = Channel.CreateUnbounded(); + Mock> writer = CreateWorkItemWriter(workItems); + Task connection = server.GetWorkItems( + new() { Capabilities = { P.WorkerCapability.HistoryStreaming } }, + writer.Object, CreateContext(timeout.Token)); + OrchestrationInstance instance = new() { InstanceId = "instance", ExecutionId = "current" }; + HistoryEvent[] history = CreateHistory(instance); + Task episode = ((ITaskExecutor)server).ExecuteOrchestrator(instance, history, []); + + try + { + P.OrchestratorRequest request = (await workItems.Reader.ReadAsync(timeout.Token)).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); + + // Act + await server.CompleteOrchestratorTask(new() + { + InstanceId = instance.InstanceId, + 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.Equal("\"waiting for activity\"", result.CustomStatus); + Assert.Equal(0, WorkerHistorySnapshotTestHelpers.GetSnapshots(server).Count); + Assert.Empty(await ReadWorkerHistoryAsync(server, instance.InstanceId)); + } + finally + { + timeout.Cancel(); + await connection; + } + } + + /// + /// Releases a failed dispatch's snapshot without removing another active episode's history. + /// + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task ExecuteOrchestrator_SendFailure_ReleasesOnlyFailedSnapshot(bool streamClosed) + { + // Arrange + using CancellationTokenSource timeout = new(TimeSpan.FromSeconds(30)); + using TaskHubGrpcServer server = CreateServer(timeout.Token); + Channel workItems = Channel.CreateUnbounded(); + Mock> writer = new(); + 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()); + Task connection = server.GetWorkItems( + new() { Capabilities = { P.WorkerCapability.HistoryStreaming } }, + writer.Object, CreateContext(timeout.Token)); + ITaskExecutor executor = server; + OrchestrationInstance other = new() { InstanceId = "other", ExecutionId = "current" }; + HistoryEvent[] otherHistory = CreateHistory(other); + Task otherEpisode = executor.ExecuteOrchestrator(other, otherHistory, []); + await workItems.Reader.ReadAsync(timeout.Token); + OrchestrationInstance failed = new() { InstanceId = "failed", ExecutionId = "current" }; + + try + { + // Act + Task failedEpisode = executor.ExecuteOrchestrator(failed, CreateHistory(failed), []); + + // Assert + if (streamClosed) + { + await Assert.ThrowsAsync(() => failedEpisode); + RpcException disconnected = await Assert.ThrowsAsync(() => + executor.ExecuteOrchestrator(failed, CreateHistory(failed), [])); + Assert.Equal(StatusCode.Unavailable, disconnected.StatusCode); + } + else + { + await Assert.ThrowsAsync(() => failedEpisode); + } + + ConcurrentDictionary> snapshots = WorkerHistorySnapshotTestHelpers.GetSnapshots(server); + Assert.Equal(1, snapshots.Count); + Assert.True(snapshots.ContainsKey(other.InstanceId)); + Assert.Empty(await ReadWorkerHistoryAsync(server, failed.InstanceId)); + Assert.Equal(otherHistory.Select(ProtobufUtils.ToHistoryEventProto), + await ReadWorkerHistoryAsync(server, other.InstanceId)); + RpcException missing = await Assert.ThrowsAsync(() => + server.CompleteOrchestratorTask(new() { InstanceId = failed.InstanceId }, CreateContext())); + Assert.Equal(StatusCode.NotFound, missing.StatusCode); + Assert.False(otherEpisode.IsCompleted); + + await server.CompleteOrchestratorTask(new() { InstanceId = other.InstanceId }, CreateContext()); + await otherEpisode.WaitAsync(timeout.Token); + Assert.Equal(0, snapshots.Count); + } + finally + { + timeout.Cancel(); + await connection; + } + } + + /// + /// Lets a captured reader finish while a subsequent episode owns a different replay snapshot. + /// + [Fact] + public async Task StreamHistoryAsync_EpisodeCompletion_PreservesReaderAndNextSnapshot() + { + // Arrange + using CancellationTokenSource timeout = new(TimeSpan.FromSeconds(30)); + using TaskHubGrpcServer server = CreateServer(timeout.Token); + Channel workItems = Channel.CreateUnbounded(); + Mock> writer = CreateWorkItemWriter(workItems); + Task connection = server.GetWorkItems( + new() { Capabilities = { P.WorkerCapability.HistoryStreaming } }, + writer.Object, CreateContext(timeout.Token)); + ITaskExecutor executor = server; + OrchestrationInstance instance = new() { InstanceId = "instance", ExecutionId = "previous" }; + HistoryEvent[] history = CreateHistory(instance); + Task episode = executor.ExecuteOrchestrator(instance, history, []); + await workItems.Reader.ReadAsync(timeout.Token); + TaskCompletionSource firstChunkWritten = new(TaskCreationOptions.RunContinuationsAsynchronously); + TaskCompletionSource releaseReader = new(TaskCreationOptions.RunContinuationsAsynchronously); + List chunks = new(); + Mock> historyWriter = new(); + historyWriter.Setup(w => w.WriteAsync(It.IsAny())).Returns(async chunk => + { + chunks.Add(chunk.Clone()); + if (chunks.Count == 1) + { + firstChunkWritten.TrySetResult(); + await releaseReader.Task.WaitAsync(timeout.Token); + } + }); + Task reader = server.StreamInstanceHistory( + new() { InstanceId = instance.InstanceId, ForWorkItemProcessing = true }, + historyWriter.Object, CreateContext(timeout.Token)); + + try + { + await firstChunkWritten.Task.WaitAsync(timeout.Token); + + // Act + await server.CompleteOrchestratorTask(new() { InstanceId = instance.InstanceId }, CreateContext()); + await episode.WaitAsync(timeout.Token); + + // Assert + Assert.Equal(0, WorkerHistorySnapshotTestHelpers.GetSnapshots(server).Count); + 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; + Assert.Equal(next.ExecutionId, nextRequest.ExecutionId); + Assert.True(nextRequest.RequiresHistoryStreaming); + Assert.Equal(nextHistory.Select(ProtobufUtils.ToHistoryEventProto), + await ReadWorkerHistoryAsync(server, next.InstanceId)); + + releaseReader.TrySetResult(); + await reader.WaitAsync(timeout.Token); + Assert.Equal(2, chunks.Count); + Assert.Equal(history.Select(ProtobufUtils.ToHistoryEventProto), chunks.SelectMany(chunk => chunk.Events)); + Assert.Equal(1, WorkerHistorySnapshotTestHelpers.GetSnapshots(server).Count); + Assert.Equal(nextHistory.Select(ProtobufUtils.ToHistoryEventProto), + WorkerHistorySnapshotTestHelpers.GetSnapshots(server)[next.InstanceId]); + Assert.False(nextEpisode.IsCompleted); + + await server.CompleteOrchestratorTask(new() { InstanceId = next.InstanceId }, CreateContext()); + await nextEpisode.WaitAsync(timeout.Token); + Assert.Equal(0, WorkerHistorySnapshotTestHelpers.GetSnapshots(server).Count); + } + finally + { + releaseReader.TrySetResult(); + await reader; + timeout.Cancel(); + await connection; + } + } + + static HistoryEvent[] CreateHistory(OrchestrationInstance instance, char payload = 'x') => + [ + new ExecutionStartedEvent(-1, new string(payload, 600 * 1024)) + { + Name = "PayloadLength", + Version = string.Empty, + OrchestrationInstance = instance, + }, + new TaskScheduledEvent(0) + { + Name = "Length", + Version = string.Empty, + Input = new string(payload, 600 * 1024), + }, + ]; + + static Mock> CreateWorkItemWriter(Channel workItems) + { + Mock> writer = new(); + writer.Setup(w => w.WriteAsync(It.IsAny())) + .Returns(workItem => workItems.Writer.WriteAsync(workItem).AsTask()); + return writer; + } + + static async Task> ReadWorkerHistoryAsync(TaskHubGrpcServer server, string instanceId) + { + List events = new(); + Mock> writer = new(); + writer.Setup(w => w.WriteAsync(It.IsAny())).Returns(chunk => + { + events.AddRange(chunk.Events); + return Task.CompletedTask; + }); + await server.StreamInstanceHistory( + new() { InstanceId = instanceId, ForWorkItemProcessing = true }, writer.Object, CreateContext()); + return events; + } + + static TaskHubGrpcServer CreateServer(CancellationToken stopping) + { + InMemoryOrchestrationService service = new(); + return new( + Mock.Of(lifetime => lifetime.ApplicationStopping == stopping), + NullLoggerFactory.Instance, service, service, Options.Create(new TaskHubGrpcServerOptions())); + } + + static ServerCallContext CreateContext(CancellationToken cancellation = default) + { + Mock context = new(); + context.Protected().SetupGet("CancellationTokenCore").Returns(cancellation); + return context.Object; + } +} + +static class WorkerHistorySnapshotTestHelpers +{ + internal static ConcurrentDictionary> GetSnapshots(DurableTaskTestHost host) + { + FieldInfo field = typeof(DurableTaskTestHost).GetField("sidecarHost", BindingFlags.Instance | BindingFlags.NonPublic) + ?? throw new MissingFieldException(nameof(DurableTaskTestHost), "sidecarHost"); + IHost sidecar = Assert.IsAssignableFrom(field.GetValue(host)); + return GetSnapshots(sidecar.Services.GetRequiredService()); + } + + internal static ConcurrentDictionary> GetSnapshots(TaskHubGrpcServer server) + { + FieldInfo field = typeof(TaskHubGrpcServer).GetField("streamingPastEvents", BindingFlags.Instance | BindingFlags.NonPublic) + ?? throw new MissingFieldException(nameof(TaskHubGrpcServer), "streamingPastEvents"); + return Assert.IsType>>(field.GetValue(server)); + } +}