diff --git a/src/Orleans.Runtime/Activation/ActivationDataActivatorProvider.cs b/src/Orleans.Runtime/Activation/ActivationDataActivatorProvider.cs index 4d14d65ce1c..26a39c15939 100644 --- a/src/Orleans.Runtime/Activation/ActivationDataActivatorProvider.cs +++ b/src/Orleans.Runtime/Activation/ActivationDataActivatorProvider.cs @@ -49,14 +49,13 @@ public bool TryGet(GrainType grainType, [NotNullWhen(true)] out IGrainContextAct return true; } - private partial class ActivationDataActivator : IGrainContextActivator + private partial class ActivationDataActivator : IPreparedGrainContextActivator { private readonly IOptions _schedulingOptions; private readonly IGrainActivator _grainActivator; private readonly IServiceProvider _serviceProvider; private readonly GrainTypeSharedContext _sharedComponents; private readonly Func _createWorkItemGroup; - private readonly Action _startActivation; public ActivationDataActivator( IGrainActivator grainActivator, @@ -73,30 +72,40 @@ public ActivationDataActivator( context, _schedulingOptions, schedulerInstruments); - _startActivation = state => ((ActivationData)state!).Start(_grainActivator); } public IGrainContext CreateContext(GrainAddress activationAddress, IConfigureGrainContext[] configureActions) + { + var preparedContext = CreatePreparedContext(activationAddress, configureActions); + using var startup = preparedContext.Start(); + return preparedContext.Context; + } + + public PreparedGrainContext CreatePreparedContext( + GrainAddress activationAddress, + IConfigureGrainContext[] configureActions) { var context = new ActivationData( activationAddress, _createWorkItemGroup, _serviceProvider, - _sharedComponents); + _sharedComponents, + _grainActivator); - foreach (var configure in configureActions) + try { - configure.Configure(context); - } + foreach (var configure in configureActions) + { + configure.Configure(context); + } - using var ecSuppressor = ExecutionContext.SuppressFlow(); - _ = Task.Factory.StartNew( - _startActivation, - context, - CancellationToken.None, - TaskCreationOptions.DenyChildAttach, - context.ActivationTaskScheduler); - return context; + return new(context, (IGrainContextStartup)context); + } + catch + { + ((IGrainContextStartup)context).Abort(); + throw; + } } } } diff --git a/src/Orleans.Runtime/Activation/IGrainContextActivator.cs b/src/Orleans.Runtime/Activation/IGrainContextActivator.cs index ab496fbccb9..1a520072617 100644 --- a/src/Orleans.Runtime/Activation/IGrainContextActivator.cs +++ b/src/Orleans.Runtime/Activation/IGrainContextActivator.cs @@ -57,6 +57,13 @@ public GrainContextActivator( /// The grain address. /// The grain context. public IGrainContext CreateInstance(GrainAddress address) + { + var preparedContext = CreatePreparedContext(address); + using var startup = preparedContext.Start(); + return preparedContext.Context; + } + + internal PreparedGrainContext CreatePreparedContext(GrainAddress address) { var grainId = address.GrainId; if (!_activators.TryGetValue(grainId.Type, out var activator)) @@ -64,7 +71,7 @@ public IGrainContext CreateInstance(GrainAddress address) activator = this.CreateActivator(grainId.Type); } - return activator.Activator.CreateContext(address, activator.ConfigureActions); + return PreparedGrainContext.Create(activator.Activator, address, activator.ConfigureActions); } private (IGrainContextActivator, IConfigureGrainContext[]) CreateActivator(GrainType grainType) @@ -134,6 +141,79 @@ public interface IGrainContextActivator public IGrainContext CreateContext(GrainAddress address, IConfigureGrainContext[] configureActions); } + internal interface IPreparedGrainContextActivator : IGrainContextActivator + { + PreparedGrainContext CreatePreparedContext( + GrainAddress address, + IConfigureGrainContext[] configureActions); + } + + internal readonly struct PreparedGrainContext(IGrainContext context, IGrainContextStartup? startup) + { + private readonly IGrainContext? _context = context; + private readonly IGrainContextStartup? _startup = startup; + + public IGrainContext Context + => _context ?? throw new InvalidOperationException("The grain context activation is not initialized."); + + public bool HasStartup => _startup is not null; + + public static PreparedGrainContext Create( + IGrainContextActivator activator, + GrainAddress address, + IConfigureGrainContext[] configureActions) => + activator is IPreparedGrainContextActivator preparedActivator + ? preparedActivator.CreatePreparedContext(address, configureActions) + : new(activator.CreateContext(address, configureActions), startup: null); + + public IDisposable Start() + { + if (_startup is not { } startup) + { + return NoopDisposable.Instance; + } + + try + { + return startup.Start(); + } + catch (Exception startupException) + { + try + { + startup.Abort(); + } + catch (Exception abortException) + { + throw new AggregateException( + "Grain context startup failed and aborting the startup also failed.", + startupException, + abortException); + } + + throw; + } + } + + public void Abort() => _startup?.Abort(); + + private sealed class NoopDisposable : IDisposable + { + public static NoopDisposable Instance { get; } = new(); + + public void Dispose() + { + } + } + } + + internal interface IGrainContextStartup + { + IDisposable Start(); + + void Abort(); + } + /// /// Provides a instance for the provided grain type. /// diff --git a/src/Orleans.Runtime/Catalog/ActivationData.cs b/src/Orleans.Runtime/Catalog/ActivationData.cs index 0ce9b85710e..03e95744d00 100644 --- a/src/Orleans.Runtime/Catalog/ActivationData.cs +++ b/src/Orleans.Runtime/Catalog/ActivationData.cs @@ -13,6 +13,7 @@ using Orleans.Internal; using Orleans.Runtime.Diagnostics; using Orleans.Runtime.GrainDirectory; +using Orleans.Runtime.Internal; using Orleans.Runtime.Placement; using Orleans.Runtime.Scheduler; using Orleans.Serialization.Invocation; @@ -37,13 +38,16 @@ internal sealed partial class ActivationData : IGrainManagementExtension, IGrainCallCancellationExtension, ICallChainReentrantGrainContext, + IGrainContextStartup, IAsyncDisposable, IDisposable { private const string GrainAddressMigrationContextKey = "sys.addr"; private readonly GrainTypeSharedContext _shared; + private readonly IGrainActivator _grainActivator; private readonly IServiceScope _serviceScope; private readonly WorkItemGroup _workItemGroup; + private WorkItemGroup.ActivationStartup? _startup; private readonly List<(Message Message, CoarseStopwatch QueuedTime)> _waitingRequests = new(); private readonly Dictionary _runningRequests = new(); private readonly SingleWaiterAutoResetEvent _workSignal = new() { RunContinuationsAsynchronously = true }; @@ -51,6 +55,7 @@ internal sealed partial class ActivationData : private Queue? _pendingOperations; private Message? _blockingRequest; private bool _isInWorkingSet = true; + private bool _wasActivated; private CoarseStopwatch _busyDuration; private CoarseStopwatch _idleDuration; private GrainReference? _selfReference; @@ -65,6 +70,7 @@ internal sealed partial class ActivationData : #pragma warning disable IDE0052 // Remove unread private members private Task? _messageLoopTask; #pragma warning restore IDE0052 // Remove unread private members + private int _started; private Activity? _activationActivity; @@ -88,18 +94,29 @@ public ActivationData( GrainAddress grainAddress, Func createWorkItemGroup, IServiceProvider applicationServices, - GrainTypeSharedContext shared) + GrainTypeSharedContext shared, + IGrainActivator grainActivator) { ArgumentNullException.ThrowIfNull(grainAddress); ArgumentNullException.ThrowIfNull(createWorkItemGroup); ArgumentNullException.ThrowIfNull(applicationServices); ArgumentNullException.ThrowIfNull(shared); + ArgumentNullException.ThrowIfNull(grainActivator); _shared = shared; + _grainActivator = grainActivator; Address = grainAddress; _serviceScope = applicationServices.CreateScope(); Debug.Assert(_serviceScope != null, "_serviceScope must not be null."); _workItemGroup = createWorkItemGroup(this); Debug.Assert(_workItemGroup != null, "_workItemGroup must not be null."); + _startup = _workItemGroup.BeginActivationStartup(); + _workItemGroup.QueueAction( + static state => + { + var context = (ActivationData)state; + context._messageLoopTask = context.RunMessageLoop(); + }, + this); } internal void SetActivationActivity(Activity activity) @@ -116,15 +133,115 @@ internal void SetActivationActivity(Activity activity) return _activationActivity?.Context; } - public void Start(IGrainActivator grainActivator) + IDisposable IGrainContextStartup.Start() + { + if (Interlocked.Exchange(ref _started, 1) != 0 || _startup is not { } startup) + { + throw new InvalidOperationException("The activation has already been started."); + } + + try + { + ExecutionContext.Run( + DefaultExecutionContext.Instance, + static state => + { + var context = (ActivationData)state!; + var task = new Task( + static state => ((ActivationData)state!).StartCore(), + context, + CancellationToken.None, + TaskCreationOptions.DenyChildAttach); + context._startup!.RunConstructor(task); + task.GetAwaiter().GetResult(); + }, + this); + _startup = null; + return startup; + } + catch (Exception exception) + { + _startup = null; + try + { + Deactivate( + new DeactivationReason( + DeactivationReasonCode.ActivationFailed, + exception, + "Error starting grain construction."), + CancellationToken.None); + } + finally + { + startup.Dispose(); + } + + throw; + } + } + + void IGrainContextStartup.Abort() + { + if (Interlocked.Exchange(ref _startup, null) is { } startup) + { + startup.Abort(); + using var suppressExecutionContext = new ExecutionContextSuppressor(); + Task.Factory.StartNew( + static state => ((ActivationData)state!).AbortStartupAsync().AsTask(), + this, + CancellationToken.None, + TaskCreationOptions.DenyChildAttach, + TaskScheduler.Default) + .Unwrap() + .Ignore(); + } + } + + private async ValueTask AbortStartupAsync() + { + var exception = new InvalidOperationException("Activation startup was aborted."); + DeactivationReason = new( + DeactivationReasonCode.ActivationFailed, + exception, + exception.Message); + + lock (this) + { + SetState(ActivationState.Invalid); + } + + try + { + RejectAllQueuedMessages(); + } + finally + { + try + { + await DisposeAsync(); + } + finally + { + GetDeactivationCompletionSource().TrySetResult(true); + _workSignal.Signal(); + } + } + } + + private void StartCore() { Debug.Assert(Equals(ActivationTaskScheduler, TaskScheduler.Current)); // locking on `this` is intentional as there are other places in the codebase taking locks on ActivationData instances lock (this) { + if (State is not ActivationState.Creating) + { + return; + } + try { - var instance = grainActivator.CreateInstance(this); + var instance = _grainActivator.CreateInstance(this); SetGrainInstance(instance); _activationActivity?.AddEvent(new ActivityEvent("instance-created")); @@ -137,7 +254,6 @@ public void Start(IGrainActivator grainActivator) Deactivate(new(DeactivationReasonCode.ActivationFailed, exception, "Error constructing grain instance."), _activationActivity?.Context, CancellationToken.None); } - _messageLoopTask = RunMessageLoop(); } } @@ -575,7 +691,7 @@ internal bool TryStartMigration(Dictionary? requestContext, Canc { lock (this) { - if (State is not (ActivationState.Activating or ActivationState.Valid or ActivationState.Deactivating)) + if (State is not (ActivationState.Valid or ActivationState.Deactivating)) { return false; } @@ -868,7 +984,11 @@ public async ValueTask DisposeAsync() lock (this) { - _shared.InternalRuntime.ActivationWorkingSet.OnDeactivated(this); + if (_wasActivated) + { + _shared.InternalRuntime.ActivationWorkingSet.OnDeactivated(this); + } + SetState(ActivationState.Invalid); } @@ -888,7 +1008,11 @@ public async ValueTask DisposeAsync() try { - _shared.OnDestroyActivation(this); + if (GrainInstance is not null) + { + _shared.OnDestroyActivation(this); + } + GetComponent()?.OnDestroyActivation(this); } catch (ObjectDisposedException) @@ -1554,9 +1678,29 @@ private void ReceiveRequest(Message message) return; } + var invalid = false; lock (this) { - _waitingRequests.Add((message, CoarseStopwatch.StartNew())); + if (State is ActivationState.Invalid) + { + invalid = true; + } + else + { + _waitingRequests.Add((message, CoarseStopwatch.StartNew())); + } + } + + if (invalid) + { + _shared.InternalRuntime.MessageCenter.ProcessRequestsToInvalidActivation( + [message], + Address, + forwardingAddress: ForwardingAddress, + failedOperation: DeactivationReason.Description, + exc: DeactivationException, + rejectMessages: DeactivationException is not null && ForwardingAddress is null); + return; } _workSignal.Signal(); @@ -1849,6 +1993,7 @@ private async Task ActivateAsync(Dictionary? requestContextData, if (State is ActivationState.Activating) { SetState(ActivationState.Valid); + _wasActivated = true; _shared.InternalRuntime.ActivationWorkingSet.OnActivated(this); } } diff --git a/src/Orleans.Runtime/Catalog/Catalog.cs b/src/Orleans.Runtime/Catalog/Catalog.cs index b7902bf98a7..5c2dad4541b 100644 --- a/src/Orleans.Runtime/Catalog/Catalog.cs +++ b/src/Orleans.Runtime/Catalog/Catalog.cs @@ -135,6 +135,7 @@ internal int UnregisterGrainForTesting(GrainId grain) Dictionary? requestContextData, MigrationContext? rehydrationContext) { + PreparedGrainContext preparedContext = default; if (TryGetGrainContext(grainId, out var result)) { rehydrationContext?.Dispose(); @@ -146,28 +147,48 @@ internal int UnregisterGrainForTesting(GrainId grain) return null; } - lock (GetStripedLock(grainId)) + try + { + lock (GetStripedLock(grainId)) + { + if (TryGetGrainContext(grainId, out result)) + { + rehydrationContext?.Dispose(); + return result; + } + + if (_siloStatusOracle.CurrentStatus == SiloStatus.Active) + { + var address = new GrainAddress + { + SiloAddress = Silo, + GrainId = grainId, + ActivationId = ActivationId.NewId(), + MembershipVersion = MembershipVersion.MinValue, + }; + + preparedContext = this.grainActivator.CreatePreparedContext(address); + result = preparedContext.Context; + activations.RecordNewTarget(result); + } + } + } + catch (Exception exception) when (result is null) { - if (TryGetGrainContext(grainId, out result)) + try { rehydrationContext?.Dispose(); - return result; } - - if (_siloStatusOracle.CurrentStatus == SiloStatus.Active) + catch (Exception cleanupException) { - var address = new GrainAddress - { - SiloAddress = Silo, - GrainId = grainId, - ActivationId = ActivationId.NewId(), - MembershipVersion = MembershipVersion.MinValue, - }; - - result = this.grainActivator.CreateInstance(address); - activations.RecordNewTarget(result); + throw new AggregateException( + "Error creating grain activation and disposing its rehydration context.", + exception, + cleanupException); } - } // End lock + + throw; + } if (result is null) { @@ -175,36 +196,103 @@ internal int UnregisterGrainForTesting(GrainId grain) return UnableToCreateActivation(this, grainId); } - // Start activation span with parent context from request if available - var parentContext = requestContextData.TryGetActivityContext(); - var activationActivity = parentContext.HasValue - ? ActivitySources.LifecycleGrainSource.StartActivity(ActivityNames.ActivateGrain, ActivityKind.Internal, parentContext.Value) - : ActivitySources.LifecycleGrainSource.StartActivity(ActivityNames.ActivateGrain, ActivityKind.Internal); - if (activationActivity is not null) + IDisposable activationStartup; + var startAttempted = false; + try { - activationActivity.SetTag(ActivityTagKeys.GrainId, grainId.ToString()); - activationActivity.SetTag(ActivityTagKeys.GrainType, grainId.Type.ToString()); - activationActivity.SetTag(ActivityTagKeys.SiloId, Silo.ToString()); - activationActivity.SetTag(ActivityTagKeys.ActivationCause, rehydrationContext is null ? "new" : "rehydrate"); - if (result is ActivationData act) + // Start activation span with parent context from request if available + var parentContext = requestContextData.TryGetActivityContext(); + var activationActivity = parentContext.HasValue + ? ActivitySources.LifecycleGrainSource.StartActivity(ActivityNames.ActivateGrain, ActivityKind.Internal, parentContext.Value) + : ActivitySources.LifecycleGrainSource.StartActivity(ActivityNames.ActivateGrain, ActivityKind.Internal); + if (activationActivity is not null) { - activationActivity.SetTag(ActivityTagKeys.ActivationId, act.ActivationId.ToString()); - act.SetActivationActivity(activationActivity); - activationActivity.AddEvent(new ActivityEvent("creating")); + activationActivity.SetTag(ActivityTagKeys.GrainId, grainId.ToString()); + activationActivity.SetTag(ActivityTagKeys.GrainType, grainId.Type.ToString()); + activationActivity.SetTag(ActivityTagKeys.SiloId, Silo.ToString()); + activationActivity.SetTag(ActivityTagKeys.ActivationCause, rehydrationContext is null ? "new" : "rehydrate"); + if (result is ActivationData act) + { + activationActivity.SetTag(ActivityTagKeys.ActivationId, act.ActivationId.ToString()); + act.SetActivationActivity(activationActivity); + activationActivity.AddEvent(new ActivityEvent("creating")); + } } + + startAttempted = true; + activationStartup = StartPreparedContext(preparedContext, result, activations); } + catch (Exception exception) + { + List? cleanupExceptions = null; + if (!startAttempted) + { + if (activations.RemoveTarget(result)) + { + LogTraceUnregisteredActivation(result); + } - _catalogInstruments.OnActivationCreated(); + try + { + AbortPreparedContext( + preparedContext, + result, + new DeactivationReason( + DeactivationReasonCode.ActivationFailed, + exception, + "Error preparing grain activation.")); + } + catch (Exception cleanupException) + { + (cleanupExceptions ??= []).Add(cleanupException); + } + } - // Rehydration occurs before activation. - if (rehydrationContext is not null) - { - result.Rehydrate(rehydrationContext); + try + { + rehydrationContext?.Dispose(); + } + catch (Exception cleanupException) + { + (cleanupExceptions ??= []).Add(cleanupException); + } + + if (cleanupExceptions is not null) + { + cleanupExceptions.Insert(0, exception); + throw new AggregateException("Error preparing grain activation and cleaning up the failed context.", cleanupExceptions); + } + + throw; } - // Initialize the new activation asynchronously. - result.Activate(requestContextData); - return result; + using (activationStartup) + { + _catalogInstruments.OnActivationCreated(); + + try + { + // Rehydration occurs before activation. + if (rehydrationContext is not null) + { + result.Rehydrate(rehydrationContext); + } + + // Initialize the new activation asynchronously. + result.Activate(requestContextData); + return result; + } + catch (Exception exception) + { + result.Deactivate( + new DeactivationReason( + DeactivationReasonCode.ActivationFailed, + exception, + "Error starting grain activation."), + CancellationToken.None); + throw; + } + } [MethodImpl(MethodImplOptions.NoInlining)] static IGrainContext? UnableToCreateActivation(Catalog self, GrainId grainId) @@ -240,6 +328,37 @@ internal int UnregisterGrainForTesting(GrainId grain) } } + internal static IDisposable StartPreparedContext( + PreparedGrainContext preparedContext, + IGrainContext context, + ActivationDirectory activations) + { + try + { + return preparedContext.Start(); + } + catch + { + activations.RemoveTarget(context); + throw; + } + } + + internal static void AbortPreparedContext( + PreparedGrainContext preparedContext, + IGrainContext context, + DeactivationReason reason) + { + if (preparedContext.HasStartup) + { + preparedContext.Abort(); + } + else + { + context.Deactivate(reason, CancellationToken.None); + } + } + private async Task UnregisterNonExistentActivation(GrainAddress address) { try diff --git a/src/Orleans.Runtime/Catalog/StatelessWorkerGrainContext.cs b/src/Orleans.Runtime/Catalog/StatelessWorkerGrainContext.cs index b244328caf0..6272d058c2f 100644 --- a/src/Orleans.Runtime/Catalog/StatelessWorkerGrainContext.cs +++ b/src/Orleans.Runtime/Catalog/StatelessWorkerGrainContext.cs @@ -326,20 +326,57 @@ private ActivationData CreateWorker(object? message) { Debug.Assert(!_terminated, "CreateWorker must not be called on a terminated stateless worker context."); var address = GrainAddress.GetAddress(Address.SiloAddress, Address.GrainId, ActivationId.NewId()); - var newWorker = (ActivationData)_innerActivator.CreateContext(address, []); - - // Observe the create/destroy lifecycle of the activation - newWorker.SetComponent(this); + var preparedContext = PreparedGrainContext.Create(_innerActivator, address, []); + var newWorker = (ActivationData)preparedContext.Context; + IDisposable activationStartup; + var startInvoked = false; + try + { + // Observe the create/destroy lifecycle of the activation + newWorker.SetComponent(this); + _workers.Add(newWorker); + startInvoked = true; + activationStartup = preparedContext.Start(); + } + catch + { + try + { + if (!startInvoked) + { + preparedContext.Abort(); + } + } + finally + { + _workers.Remove(newWorker); + } - // If this is a new worker and there is a message in scope, try to get the request context and activate the worker - var requestContext = (message as Message)?.RequestContextData ?? []; - var cancellation = new CancellationTokenSource(_shared.Shared.InternalRuntime.CollectionOptions.Value.ActivationTimeout); + throw; + } - newWorker.Activate(requestContext, cancellation.Token); - _workers.Add(newWorker); - StatelessWorkerEvents.EmitWorkerCreated(this, newWorker, _workers.Count); + using (activationStartup) + { + try + { + // If this is a new worker and there is a message in scope, try to get the request context and activate the worker + var requestContext = (message as Message)?.RequestContextData ?? []; + newWorker.Activate(requestContext, CancellationToken.None); + StatelessWorkerEvents.EmitWorkerCreated(this, newWorker, _workers.Count); - return newWorker; + return newWorker; + } + catch (Exception exception) + { + newWorker.Deactivate( + new DeactivationReason( + DeactivationReasonCode.ActivationFailed, + exception, + "Error starting stateless worker activation."), + CancellationToken.None); + throw; + } + } } private void DeactivateInternal(DeactivationReason reason, CancellationToken cancellationToken) diff --git a/src/Orleans.Runtime/Scheduler/ActivationTaskScheduler.cs b/src/Orleans.Runtime/Scheduler/ActivationTaskScheduler.cs index a854e04b108..5eaad97fbc0 100644 --- a/src/Orleans.Runtime/Scheduler/ActivationTaskScheduler.cs +++ b/src/Orleans.Runtime/Scheduler/ActivationTaskScheduler.cs @@ -56,10 +56,24 @@ internal void RunTaskFromWorkItemGroup(Task task) } } + internal void RunTaskSynchronously(Task task) + { + task.Start(this); + if (!TryExecuteTask(task)) + { + throw new InvalidOperationException($"Unable to execute synchronous task {task.Id}."); + } + } + /// Queues a task to the scheduler. /// The task to be queued. protected override void QueueTask(Task task) { + if (workerGroup.IsCurrentTask(task)) + { + return; + } + #if DEBUG LogTraceQueueTask(myId, task.Id); #endif diff --git a/src/Orleans.Runtime/Scheduler/WorkItemGroup.cs b/src/Orleans.Runtime/Scheduler/WorkItemGroup.cs index 8320c332a8d..1ba689cd795 100644 --- a/src/Orleans.Runtime/Scheduler/WorkItemGroup.cs +++ b/src/Orleans.Runtime/Scheduler/WorkItemGroup.cs @@ -56,6 +56,8 @@ internal int ExternalWorkItemCount get { lock (_lockObj) { return _workItems.Count; } } } + internal bool IsCurrentTask(Task task) => ReferenceEquals(_currentTask, task); + public WorkItemGroup( IGrainContext grainContext, IOptions schedulingOptions, @@ -119,6 +121,124 @@ public void EnqueueTask(Task task) } } + internal ActivationStartup BeginActivationStartup() + { + lock (_lockObj) + { + if (_state != WorkGroupStatus.Waiting) + { + throw new InvalidOperationException($"Cannot reserve execution while {this} is {_state}."); + } + + _state = WorkGroupStatus.Running; + return new(this); + } + } + + private void RunTaskSynchronously(Task task) + { + long taskStart; + lock (_lockObj) + { + if (_state != WorkGroupStatus.Running || _currentTask is not null) + { + throw new InvalidOperationException($"Synchronous execution requires a reserved {this}."); + } + + _currentTask = task; + _currentTaskStarted = taskStart = Environment.TickCount64; + _totalItemsEnqueued++; + } + + RuntimeContext.SetExecutionContext(GrainContext, out var originalContext); + try + { +#if DEBUG + LogTaskStart(task); +#endif + TaskScheduler.RunTaskSynchronously(task); + } + finally + { + RuntimeContext.ResetExecutionContext(originalContext); + _totalItemsProcessed++; + var taskDurationMs = Environment.TickCount64 - taskStart; + if (taskDurationMs > (long)Math.Ceiling(_schedulingOptions.TurnWarningLengthThreshold.TotalMilliseconds)) + { + _schedulerInstruments.OnLongRunningTurn(); + LogLongRunningTurn(task, taskDurationMs); + } + + _currentTask = null; + } + } + + private void ReleaseExecution() + { + lock (_lockObj) + { + Debug.Assert(_state == WorkGroupStatus.Running); + Debug.Assert(_currentTask is null); + if (_workItems.Count > 0) + { + _state = WorkGroupStatus.Runnable; + ScheduleExecution(this); + } + else + { + _state = WorkGroupStatus.Waiting; + } + } + } + + private void AbortExecution() + { + lock (_lockObj) + { + if (_state != WorkGroupStatus.Running || _currentTask is not null) + { + throw new InvalidOperationException($"Cannot abort execution while {this} is {_state}."); + } + + _workItems.Clear(); + _state = WorkGroupStatus.Waiting; + } + } + // One-shot scheduler reservation used while an activation is published and constructed. + internal sealed class ActivationStartup : IDisposable + { + private WorkItemGroup? _owner; + private int _constructorStarted; + + internal ActivationStartup(WorkItemGroup owner) + { + _owner = owner; + } + + public void RunConstructor(Task task) + { + if (Interlocked.Exchange(ref _constructorStarted, 1) != 0) + { + throw new InvalidOperationException("The activation constructor has already been run."); + } + + (_owner ?? throw new ObjectDisposedException(nameof(ActivationStartup))) + .RunTaskSynchronously(task); + } + + public void Dispose() => Interlocked.Exchange(ref _owner, null)?.ReleaseExecution(); + + public void Abort() + { + if (Volatile.Read(ref _constructorStarted) != 0) + { + throw new InvalidOperationException("An activation startup cannot be aborted after construction has begun."); + } + + Interlocked.Exchange(ref _owner, null)?.AbortExecution(); + } + } + [MethodImpl(MethodImplOptions.NoInlining)] private void LogTooManyTasksInQueue(int count, int maxPendingItemsLimit) { diff --git a/src/Orleans.Runtime/Utils/DefaultExecutionContext.cs b/src/Orleans.Runtime/Utils/DefaultExecutionContext.cs new file mode 100644 index 00000000000..97708c831e1 --- /dev/null +++ b/src/Orleans.Runtime/Utils/DefaultExecutionContext.cs @@ -0,0 +1,38 @@ +using System.Threading; +using System.Threading.Tasks; + +namespace Orleans.Runtime; + +/// +/// Provides an which contains no ambient state. +/// +internal static class DefaultExecutionContext +{ + public static ExecutionContext Instance { get; } = CaptureDefault(); + + internal static ExecutionContext CaptureDefault() + { + var completion = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + if (!ThreadPool.UnsafeQueueUserWorkItem( + static completion => + { + try + { + completion.SetResult( + ExecutionContext.Capture() + ?? throw new InvalidOperationException("Could not capture the default execution context.")); + } + catch (Exception exception) + { + completion.SetException(exception); + } + }, + completion, + preferLocal: false)) + { + throw new InvalidOperationException("Could not queue work to capture the default execution context."); + } + + return completion.Task.GetAwaiter().GetResult(); + } +} diff --git a/test/Grains/TestInternalGrains/ActivationStartupTestGrain.cs b/test/Grains/TestInternalGrains/ActivationStartupTestGrain.cs new file mode 100644 index 00000000000..308e016b412 --- /dev/null +++ b/test/Grains/TestInternalGrains/ActivationStartupTestGrain.cs @@ -0,0 +1,121 @@ +using Orleans.Runtime; + +namespace UnitTests.Grains; + +public interface IActivationStartupTestGrain : IGrainWithStringKey +{ + Task Invoke(string payload); +} + +public sealed class ActivationStartupTestGrain : + IActivationStartupTestGrain, + IGrainBase, + ILifecycleParticipant +{ + private readonly ActivationStartupScenario _scenario; + + public ActivationStartupTestGrain(IGrainContext grainContext, ActivationStartupTestHooks hooks) + { + GrainContext = grainContext; + _scenario = hooks.GetRequiredScenario(grainContext.GrainId); + _scenario.ObserveConstructor(grainContext); + } + + public IGrainContext GrainContext { get; } + + public void Participate(IGrainLifecycle lifecycle) + { + lifecycle.Subscribe( + "ActivationStartup-Low", + GrainLifecycleStage.SetupState, + _ => + { + _scenario.Record("LifecycleStartLow", GrainContext); + return Task.CompletedTask; + }, + _ => + { + _scenario.Record("LifecycleStopLow", GrainContext); + return Task.CompletedTask; + }); + lifecycle.Subscribe( + "ActivationStartup-High", + GrainLifecycleStage.Activate, + _ => + { + _scenario.Record("LifecycleStartHigh", GrainContext); + return Task.CompletedTask; + }, + _ => + { + _scenario.Record("LifecycleStopHigh", GrainContext); + return Task.CompletedTask; + }); + } + + public Task OnActivateAsync(CancellationToken cancellationToken) + { + _scenario.ObserveOnActivate(GrainContext); + return _scenario.Completion switch + { + ActivationStartupCompletion.ImmediateSuccess => CompleteActivation(), + ActivationStartupCompletion.ImmediateFailure => FailActivation(), + ActivationStartupCompletion.AsynchronousSuccess => CompleteActivationAfterRelease(cancellationToken), + ActivationStartupCompletion.AsynchronousFailure => FailActivationAfterRelease(cancellationToken), + ActivationStartupCompletion.Cancellation => CompleteActivationAfterRelease(cancellationToken), + _ => throw new ArgumentOutOfRangeException(), + }; + } + + public Task OnDeactivateAsync(DeactivationReason reason, CancellationToken cancellationToken) + { + _scenario.ObserveOnDeactivate(GrainContext); + _scenario.Record("OnDeactivateCompleted", GrainContext); + return Task.CompletedTask; + } + + public Task Invoke(string payload) + { + _scenario.ObserveRequest(GrainContext); + return Task.FromResult($"{payload}:{_scenario.RequestContextValue}"); + } + + private Task CompleteActivation() + { + _scenario.Record("OnActivateCompleted", GrainContext); + return Task.CompletedTask; + } + + private Task FailActivation() + { + _scenario.Record("OnActivateFailed", GrainContext); + throw _scenario.ActivationException; + } + + private async Task CompleteActivationAfterRelease(CancellationToken cancellationToken) + { + using var registration = cancellationToken.Register( + static state => + { + var (scenario, context) = ((ActivationStartupScenario, IGrainContext))state!; + scenario.Record("CancellationObserved", context); + }, + (_scenario, GrainContext)); + await _scenario.ActivationRelease.WaitAsync(cancellationToken); + _scenario.Record("OnActivateCompleted", GrainContext); + } + + private async Task FailActivationAfterRelease(CancellationToken cancellationToken) + { + using var registration = cancellationToken.Register( + static state => + { + var (scenario, context) = ((ActivationStartupScenario, IGrainContext))state!; + scenario.Record("CancellationObserved", context); + }, + (_scenario, GrainContext)); + await _scenario.ActivationRelease.WaitAsync(cancellationToken); + _scenario.Record("OnActivateFailed", GrainContext); + throw _scenario.ActivationException; + } +} diff --git a/test/Grains/TestInternalGrains/ActivationStartupTestHooks.cs b/test/Grains/TestInternalGrains/ActivationStartupTestHooks.cs new file mode 100644 index 00000000000..28a10371f8d --- /dev/null +++ b/test/Grains/TestInternalGrains/ActivationStartupTestHooks.cs @@ -0,0 +1,381 @@ +using System.Collections.Concurrent; +using Microsoft.Extensions.DependencyInjection; +using Orleans.Runtime; + +namespace UnitTests.Grains; + +public enum ActivationStartupCompletion +{ + ImmediateSuccess, + AsynchronousSuccess, + ImmediateFailure, + AsynchronousFailure, + Cancellation, +} + +public enum ActivationStartupDisposal +{ + Synchronous, + Asynchronous, +} + +public readonly record struct ActivationStartupEvent(string Name, GrainId GrainId, ActivationId ActivationId); + +public sealed class ActivationStartupTestHooks +{ + public const string RequestContextKey = "activation-startup"; + + private readonly AsyncLocal _ambientValue = new(); + private readonly ConcurrentDictionary _scenarios = new(); + + public string? AmbientValue + { + get => _ambientValue.Value; + set => _ambientValue.Value = value; + } + + public ActivationStartupScenario CreateScenario( + GrainId grainId, + ActivationStartupCompletion completion, + ActivationStartupDisposal disposal) + { + var scenario = new ActivationStartupScenario(this, grainId, completion, disposal); + if (!_scenarios.TryAdd(grainId, scenario)) + { + throw new InvalidOperationException($"A startup scenario already exists for grain '{grainId}'."); + } + + return scenario; + } + + public ActivationStartupScenario GetRequiredScenario(GrainId grainId) => + _scenarios.TryGetValue(grainId, out var scenario) + ? scenario + : throw new InvalidOperationException($"No startup scenario exists for grain '{grainId}'."); + + public bool TryGetScenario(GrainId grainId, out ActivationStartupScenario? scenario) => + _scenarios.TryGetValue(grainId, out scenario); + + public void RemoveScenario(GrainId grainId) + { + _scenarios.TryRemove(grainId, out _); + } +} + +public sealed class ActivationStartupScenario +{ + private readonly ActivationStartupTestHooks _hooks; + private readonly TaskCompletionSource _activationRelease = CreateCompletionSource(); + private readonly TaskCompletionSource _disposalRelease = CreateCompletionSource(); + private readonly ConcurrentQueue _events = new(); + private readonly ConcurrentDictionary> _eventWaiters = new(); + private IGrainContext? _context; + private string? _constructorAmbientValue; + private string? _constructorRequestContextValue; + private string? _onActivateAmbientValue; + private string? _onActivateRequestContextValue; + private string? _requestAmbientValue; + private string? _requestContextValue; + private int _constructorCount; + private int _createCount; + private int _disposeCompletedCount; + private int _disposeStartedCount; + private int _onActivateCount; + private int _onDeactivateCount; + private int _requestInvocationCount; + private int _scopeDisposeCount; + + internal ActivationStartupScenario( + ActivationStartupTestHooks hooks, + GrainId grainId, + ActivationStartupCompletion completion, + ActivationStartupDisposal disposal) + { + _hooks = hooks; + GrainId = grainId; + Completion = completion; + Disposal = disposal; + ActivationException = new InvalidOperationException("activate-fault"); + } + + public GrainId GrainId { get; } + + public ActivationStartupCompletion Completion { get; } + + public ActivationStartupDisposal Disposal { get; } + + public InvalidOperationException ActivationException { get; } + + public IReadOnlyList Events => _events.ToArray(); + + public int ConstructorCount => Volatile.Read(ref _constructorCount); + + public int CreateCount => Volatile.Read(ref _createCount); + + public int DisposeCompletedCount => Volatile.Read(ref _disposeCompletedCount); + + public int DisposeStartedCount => Volatile.Read(ref _disposeStartedCount); + + public int OnActivateCount => Volatile.Read(ref _onActivateCount); + + public int OnDeactivateCount => Volatile.Read(ref _onDeactivateCount); + + public int RequestInvocationCount => Volatile.Read(ref _requestInvocationCount); + + public int ScopeDisposeCount => Volatile.Read(ref _scopeDisposeCount); + + public string? ConstructorAmbientValue => Volatile.Read(ref _constructorAmbientValue); + + public string? ConstructorRequestContextValue => Volatile.Read(ref _constructorRequestContextValue); + + public string? OnActivateAmbientValue => Volatile.Read(ref _onActivateAmbientValue); + + public string? OnActivateRequestContextValue => Volatile.Read(ref _onActivateRequestContextValue); + + public string? RequestAmbientValue => Volatile.Read(ref _requestAmbientValue); + + public string? RequestContextValue => Volatile.Read(ref _requestContextValue); + + internal IGrainContext Context => + Volatile.Read(ref _context) + ?? throw new InvalidOperationException($"The startup scenario for grain '{GrainId}' has no activation context."); + + internal Task ActivationRelease => _activationRelease.Task; + + internal Task DisposalRelease => _disposalRelease.Task; + + public Task WaitForEvent(string name) => + _eventWaiters.GetOrAdd(name, static _ => CreateEventCompletionSource()).Task; + + public void ReleaseActivation() => _activationRelease.TrySetResult(); + + public void ReleaseDisposal() => _disposalRelease.TrySetResult(); + + public int GetEventCount(string name) => _events.Count(entry => entry.Name == name); + + public void Record(string name, IGrainContext context) + { + Attach(context); + var entry = new ActivationStartupEvent(name, context.GrainId, context.ActivationId); + _events.Enqueue(entry); + _eventWaiters.GetOrAdd(name, static _ => CreateEventCompletionSource()).TrySetResult(entry); + } + + internal void ObserveCreate(IGrainContext context) + { + Attach(context); + Interlocked.Increment(ref _createCount); + } + + internal void ObserveConstructor(IGrainContext context) + { + Attach(context); + Interlocked.Increment(ref _constructorCount); + Volatile.Write(ref _constructorAmbientValue, _hooks.AmbientValue); + Volatile.Write( + ref _constructorRequestContextValue, + RequestContext.Get(ActivationStartupTestHooks.RequestContextKey) as string); + } + + internal void ObserveOnActivate(IGrainContext context) + { + Interlocked.Increment(ref _onActivateCount); + Volatile.Write(ref _onActivateAmbientValue, _hooks.AmbientValue); + Volatile.Write( + ref _onActivateRequestContextValue, + RequestContext.Get(ActivationStartupTestHooks.RequestContextKey) as string); + Record("OnActivateEntered", context); + } + + internal void ObserveOnDeactivate(IGrainContext context) + { + Interlocked.Increment(ref _onDeactivateCount); + Record("OnDeactivateEntered", context); + } + + internal void ObserveRequest(IGrainContext context) + { + Interlocked.Increment(ref _requestInvocationCount); + Volatile.Write(ref _requestAmbientValue, _hooks.AmbientValue); + Volatile.Write( + ref _requestContextValue, + RequestContext.Get(ActivationStartupTestHooks.RequestContextKey) as string); + Record("RequestInvoked", context); + } + + internal void ObserveDisposeStarted(IGrainContext context) + { + Interlocked.Increment(ref _disposeStartedCount); + Record("ActivatorDisposeStarted", context); + } + + internal void ObserveDisposeCompleted(IGrainContext context) + { + Interlocked.Increment(ref _disposeCompletedCount); + Record("ActivatorDisposeCompleted", context); + } + + internal void ObserveScopeDisposed(IGrainContext context) + { + Interlocked.Increment(ref _scopeDisposeCount); + Record("ScopeDisposed", context); + } + + private void Attach(IGrainContext context) + { + if (!context.GrainId.Equals(GrainId)) + { + throw new InvalidOperationException( + $"Scenario grain '{GrainId}' cannot observe activation '{context.GrainId}'."); + } + + var existing = Interlocked.CompareExchange(ref _context, context, null); + if (existing is not null && !ReferenceEquals(existing, context)) + { + throw new InvalidOperationException($"Scenario grain '{GrainId}' observed multiple activation contexts."); + } + } + + private static TaskCompletionSource CreateCompletionSource() => + new(TaskCreationOptions.RunContinuationsAsynchronously); + + private static TaskCompletionSource CreateEventCompletionSource() => + new(TaskCreationOptions.RunContinuationsAsynchronously); +} + +public sealed class ActivationStartupScopedResource(ActivationStartupTestHooks hooks) : IDisposable +{ + private IGrainContext? _context; + private int _disposed; + + public void Attach(IGrainContext context) => _context = context; + + public void Dispose() + { + if (Interlocked.Exchange(ref _disposed, 1) == 0 + && _context is { } context + && hooks.TryGetScenario(context.GrainId, out var scenario)) + { + scenario!.ObserveScopeDisposed(context); + } + } +} + +public enum CatalogActivationStartupOutcome +{ + Success, + ConstructorFailure, +} + +public sealed class CatalogActivationStartupTestHooks +{ + private readonly ConcurrentDictionary _gates = new(); + + public CatalogActivationStartupGate CreateGate( + GrainId grainId, + CatalogActivationStartupOutcome outcome) + { + var gate = new CatalogActivationStartupGate(grainId, outcome); + if (!_gates.TryAdd(grainId, gate)) + { + throw new InvalidOperationException($"A catalog startup gate already exists for grain '{grainId}'."); + } + + return gate; + } + + public bool TryGetGate(GrainId grainId, out CatalogActivationStartupGate? gate) => + _gates.TryGetValue(grainId, out gate); + + public void RemoveGate(GrainId grainId) => _gates.TryRemove(grainId, out _); +} + +public sealed class CatalogActivationStartupGate( + GrainId grainId, + CatalogActivationStartupOutcome outcome) +{ + private static readonly TimeSpan Timeout = TimeSpan.FromSeconds(30); + private readonly TaskCompletionSource _entered = + new(TaskCreationOptions.RunContinuationsAsynchronously); + private readonly TaskCompletionSource _release = + new(TaskCreationOptions.RunContinuationsAsynchronously); + private int _entryCount; + + public GrainId GrainId { get; } = grainId; + + public CatalogActivationStartupOutcome Outcome { get; } = outcome; + + public InvalidOperationException ConstructorException { get; } = + new("constructor-fault"); + + public Task Entered => _entered.Task; + + public int EntryCount => Volatile.Read(ref _entryCount); + + public void Release() => _release.TrySetResult(); + + internal void EnterAndWait(IGrainContext context, ActivationStartupScenario scenario) + { + if (!context.GrainId.Equals(GrainId)) + { + throw new InvalidOperationException( + $"Catalog startup gate for grain '{GrainId}' cannot observe activation '{context.GrainId}'."); + } + + if (Interlocked.Increment(ref _entryCount) != 1) + { + throw new InvalidOperationException( + $"Catalog startup gate for grain '{GrainId}' was entered more than once."); + } + + scenario.Record("ConstructorBlocked", context); + _entered.TrySetResult(context); + try + { + _release.Task.WaitAsync(Timeout).GetAwaiter().GetResult(); + } + catch (TimeoutException exception) + { + throw new TimeoutException( + $"Timed out releasing catalog startup for grain '{context.GrainId}', activation '{context.ActivationId}', outcome '{Outcome}'.", + exception); + } + } +} + +public sealed class CatalogActivationStartupTestActivator( + ActivationStartupTestHooks hooks, + CatalogActivationStartupTestHooks catalogHooks) : IGrainActivator +{ + public object CreateInstance(IGrainContext context) + { + var scenario = hooks.GetRequiredScenario(context.GrainId); + scenario.ObserveCreate(context); + context.ActivationServices.GetRequiredService().Attach(context); + var instance = new ActivationStartupTestGrain(context, hooks); + + if (catalogHooks.TryGetGate(context.GrainId, out var gate)) + { + gate!.EnterAndWait(context, scenario); + if (gate.Outcome is CatalogActivationStartupOutcome.ConstructorFailure) + { + scenario.Record("ConstructorFailed", context); + throw gate.ConstructorException; + } + } + + return instance; + } + + public async ValueTask DisposeInstance(IGrainContext context, object instance) + { + var scenario = hooks.GetRequiredScenario(context.GrainId); + scenario.ObserveDisposeStarted(context); + if (scenario.Disposal is ActivationStartupDisposal.Asynchronous) + { + await scenario.DisposalRelease; + } + + scenario.ObserveDisposeCompleted(context); + } +} diff --git a/test/Grains/TestInternalGrains/StatelessWorkerActivationStartupTestGrain.cs b/test/Grains/TestInternalGrains/StatelessWorkerActivationStartupTestGrain.cs new file mode 100644 index 00000000000..22383300bc2 --- /dev/null +++ b/test/Grains/TestInternalGrains/StatelessWorkerActivationStartupTestGrain.cs @@ -0,0 +1,127 @@ +using Orleans.Concurrency; +using Orleans.Runtime; + +namespace UnitTests.Grains; + +public interface IStatelessWorkerActivationStartupTestGrain : IGrainWithStringKey +{ + Task Invoke(string payload); +} + +[StatelessWorker(1)] +public sealed class StatelessWorkerActivationStartupTestGrain : + IStatelessWorkerActivationStartupTestGrain, + IGrainBase, + ILifecycleParticipant +{ + private readonly ActivationStartupScenario _scenario; + + public StatelessWorkerActivationStartupTestGrain( + IGrainContext grainContext, + ActivationStartupTestHooks hooks) + { + GrainContext = grainContext; + _scenario = hooks.GetRequiredScenario(grainContext.GrainId); + _scenario.ObserveConstructor(grainContext); + _scenario.Record("ConstructorEntered", grainContext); + } + + public IGrainContext GrainContext { get; } + + public void Participate(IGrainLifecycle lifecycle) + { + lifecycle.Subscribe( + "StatelessWorkerActivationStartup-Low", + GrainLifecycleStage.SetupState, + _ => + { + _scenario.Record("LifecycleStartLow", GrainContext); + return Task.CompletedTask; + }, + _ => + { + _scenario.Record("LifecycleStopLow", GrainContext); + return Task.CompletedTask; + }); + lifecycle.Subscribe( + "StatelessWorkerActivationStartup-High", + GrainLifecycleStage.Activate, + _ => + { + _scenario.Record("LifecycleStartHigh", GrainContext); + return Task.CompletedTask; + }, + _ => + { + _scenario.Record("LifecycleStopHigh", GrainContext); + return Task.CompletedTask; + }); + } + + public Task OnActivateAsync(CancellationToken cancellationToken) + { + _scenario.ObserveOnActivate(GrainContext); + return _scenario.Completion switch + { + ActivationStartupCompletion.ImmediateSuccess => CompleteActivation(), + ActivationStartupCompletion.ImmediateFailure => FailActivation(), + ActivationStartupCompletion.AsynchronousSuccess => CompleteActivationAfterRelease(cancellationToken), + ActivationStartupCompletion.AsynchronousFailure => FailActivationAfterRelease(cancellationToken), + ActivationStartupCompletion.Cancellation => CompleteActivationAfterRelease(cancellationToken), + _ => throw new ArgumentOutOfRangeException(), + }; + } + + public Task OnDeactivateAsync(DeactivationReason reason, CancellationToken cancellationToken) + { + _scenario.ObserveOnDeactivate(GrainContext); + _scenario.Record("OnDeactivateCompleted", GrainContext); + return Task.CompletedTask; + } + + public Task Invoke(string payload) + { + _scenario.ObserveRequest(GrainContext); + return Task.FromResult( + $"{payload}:{GrainContext.ActivationId}:{_scenario.RequestContextValue}"); + } + + private Task CompleteActivation() + { + _scenario.Record("OnActivateCompleted", GrainContext); + return Task.CompletedTask; + } + + private Task FailActivation() + { + _scenario.Record("OnActivateFailed", GrainContext); + throw _scenario.ActivationException; + } + + private async Task CompleteActivationAfterRelease(CancellationToken cancellationToken) + { + using var registration = cancellationToken.Register( + static state => + { + var (scenario, context) = ((ActivationStartupScenario, IGrainContext))state!; + scenario.Record("CancellationObserved", context); + }, + (_scenario, GrainContext)); + await _scenario.ActivationRelease.WaitAsync(cancellationToken); + _scenario.Record("OnActivateCompleted", GrainContext); + } + + private async Task FailActivationAfterRelease(CancellationToken cancellationToken) + { + using var registration = cancellationToken.Register( + static state => + { + var (scenario, context) = ((ActivationStartupScenario, IGrainContext))state!; + scenario.Record("CancellationObserved", context); + }, + (_scenario, GrainContext)); + await _scenario.ActivationRelease.WaitAsync(cancellationToken); + _scenario.Record("OnActivateFailed", GrainContext); + throw _scenario.ActivationException; + } +} diff --git a/test/Orleans.Core.Tests/Runtime/DefaultExecutionContextTests.cs b/test/Orleans.Core.Tests/Runtime/DefaultExecutionContextTests.cs new file mode 100644 index 00000000000..078855372f1 --- /dev/null +++ b/test/Orleans.Core.Tests/Runtime/DefaultExecutionContextTests.cs @@ -0,0 +1,75 @@ +using System.Threading; +using System.Threading.Tasks; +using Orleans.Runtime; +using TestExtensions; +using Xunit; + +namespace UnitTests.Runtime; + +[TestSuite("BVT")] +[TestProvider("None")] +[TestArea("Runtime")] +public class DefaultExecutionContextTests +{ + [Fact] + public void InstanceDoesNotContainAmbientState() => AssertDoesNotContainAmbientState(DefaultExecutionContext.Instance); + + [Fact] + public void FallbackDoesNotContainAmbientState() => AssertDoesNotContainAmbientState(DefaultExecutionContext.CaptureDefault()); + + [Fact] + public async Task InstanceSupportsConcurrentExecution() + { + var tasks = new Task[Math.Max(4, Environment.ProcessorCount)]; + for (var i = 0; i < tasks.Length; i++) + { + var expected = new object(); + tasks[i] = Task.Run( + () => + { + var ambientState = new AsyncLocal { Value = expected }; + object? observed = expected; + + ExecutionContext.Run( + DefaultExecutionContext.Instance, + _ => observed = ambientState.Value, + null); + + Assert.Null(observed); + Assert.Same(expected, ambientState.Value); + }, + TestContext.Current.CancellationToken); + } + + await Task.WhenAll(tasks); + } + + private static void AssertDoesNotContainAmbientState(ExecutionContext executionContext) + { + var ambientState = new AsyncLocal(); + var expected = new object(); + ambientState.Value = expected; + object? observed = expected; + var flowSuppressed = true; + + try + { + ExecutionContext.Run( + executionContext, + _ => + { + observed = ambientState.Value; + flowSuppressed = ExecutionContext.IsFlowSuppressed(); + }, + null); + + Assert.Null(observed); + Assert.False(flowSuppressed); + Assert.Same(expected, ambientState.Value); + } + finally + { + ambientState.Value = null; + } + } +} diff --git a/test/Orleans.Core.Tests/Runtime/GrainContextActivatorTests.cs b/test/Orleans.Core.Tests/Runtime/GrainContextActivatorTests.cs index 5ace698222b..5cb03352cf8 100644 --- a/test/Orleans.Core.Tests/Runtime/GrainContextActivatorTests.cs +++ b/test/Orleans.Core.Tests/Runtime/GrainContextActivatorTests.cs @@ -1,5 +1,6 @@ using System.Collections.Generic; using System.Diagnostics.CodeAnalysis; +using Microsoft.Extensions.DependencyInjection; using NSubstitute; using Orleans.Metadata; using Orleans.Runtime; @@ -8,6 +9,9 @@ namespace UnitTests.Runtime; +[TestSuite("BVT")] +[TestProvider("None")] +[TestArea("Runtime")] public class GrainContextActivatorTests { [Fact, TestCategory("BVT")] @@ -26,6 +30,154 @@ [new TestConfigureGrainContextProvider(events)], Assert.Equal(["configure", "activate"], events); } + [Fact] + public void PreparedContextCreate_CustomActivatorRemainsEager() + { + var events = new List(); + var context = Substitute.For(); + var activator = new TestGrainContextActivator(context, events); + var address = new GrainAddress { GrainId = GrainId.Create("test", "grain") }; + IConfigureGrainContext[] configureActions = [new TestConfigureGrainContext(events)]; + + var preparedContext = PreparedGrainContext.Create(activator, address, configureActions); + using var startup = preparedContext.Start(); + preparedContext.Abort(); + + Assert.Same(context, preparedContext.Context); + Assert.Equal(["configure", "activate"], events); + } + + [Fact] + public void CreateInstance_PreparedContextStartsAndReleasesExactlyOnce() + { + var events = new List(); + var context = Substitute.For(); + var activator = CreateActivator(new TestPreparedGrainContextActivator(context, events), events); + var address = new GrainAddress { GrainId = GrainId.Create("test", "grain") }; + + Assert.Same(context, activator.CreateInstance(address)); + Assert.Equal(["configure", "create", "start", "release"], events); + } + + [Fact] + public void PreparedContext_AbortDoesNotStart() + { + var events = new List(); + var context = Substitute.For(); + var preparedContext = new PreparedGrainContext(context, new TestGrainContextStartup(events)); + + preparedContext.Abort(); + + Assert.Equal(["abort"], events); + } + + [Fact] + public void PreparedContext_StartFailureAbortsExactlyOnce() + { + var events = new List(); + var context = Substitute.For(); + var expected = new InvalidOperationException("start-fault"); + var preparedContext = new PreparedGrainContext( + context, + new ThrowingGrainContextStartup(events, expected)); + + var actual = Assert.Throws(preparedContext.Start); + + Assert.Same(expected, actual); + Assert.Equal(["start", "abort"], events); + } + + [Fact] + public void PreparedContext_StartAndAbortFailureReportsBothExceptions() + { + var events = new List(); + var context = Substitute.For(); + var startupException = new InvalidOperationException("start-fault"); + var abortException = new InvalidOperationException("abort-fault"); + var preparedContext = new PreparedGrainContext( + context, + new ThrowingStartAndAbortGrainContextStartup( + events, + startupException, + abortException)); + + var actual = Assert.Throws(preparedContext.Start); + + Assert.Equal( + "Grain context startup failed and aborting the startup also failed.", + actual.Message.Split(" (")[0]); + Assert.Equal([startupException, abortException], actual.InnerExceptions); + Assert.Equal(["start", "abort"], events); + } + + [Fact] + public void CatalogStartPreparedContext_StartFailureRemovesRecordedTarget() + { + var events = new List(); + var grainId = GrainId.Create("test", "grain"); + var context = Substitute.For(); + context.GrainId.Returns(grainId); + context.Equals(Arg.Any()) + .Returns(call => ReferenceEquals(context, call.Arg())); + var expected = new InvalidOperationException("start-fault"); + var preparedContext = new PreparedGrainContext( + context, + new ThrowingGrainContextStartup(events, expected)); + using var services = new ServiceCollection() + .AddMetrics() + .AddSingleton() + .AddSingleton() + .BuildServiceProvider(); + var activations = new ActivationDirectory(services.GetRequiredService()); + activations.RecordNewTarget(context); + + var actual = Assert.Throws( + () => Catalog.StartPreparedContext(preparedContext, context, activations)); + + Assert.Same(expected, actual); + Assert.Null(activations.FindTarget(grainId)); + Assert.Equal(0, activations.Count); + Assert.Equal(["start", "abort"], events); + } + + [Fact] + public void CatalogAbortPreparedContext_EagerCustomContextIsDeactivated() + { + var context = Substitute.For(); + var reason = new DeactivationReason( + DeactivationReasonCode.ActivationFailed, + "pre-start-failure"); + var preparedContext = new PreparedGrainContext(context, startup: null); + + Catalog.AbortPreparedContext(preparedContext, context, reason); + + context.Received(1).Deactivate(reason, CancellationToken.None); + } + + [Fact] + public void CatalogAbortPreparedContext_PreparedContextIsAborted() + { + var events = new List(); + var context = Substitute.For(); + var reason = new DeactivationReason( + DeactivationReasonCode.ActivationFailed, + "pre-start-failure"); + var preparedContext = new PreparedGrainContext(context, new TestGrainContextStartup(events)); + + Catalog.AbortPreparedContext(preparedContext, context, reason); + + Assert.Equal(["abort"], events); + context.DidNotReceive().Deactivate(Arg.Any(), Arg.Any()); + } + + private static GrainContextActivator CreateActivator( + IGrainContextActivator contextActivator, + List events) => + new( + [new TestGrainContextActivatorProvider(contextActivator)], + [new TestConfigureGrainContextProvider(events)], + new GrainPropertiesResolver(Substitute.For())); + private sealed class TestGrainContextActivatorProvider(IGrainContextActivator activator) : IGrainContextActivatorProvider { public bool TryGet(GrainType grainType, [NotNullWhen(true)] out IGrainContextActivator? result) @@ -65,4 +217,84 @@ public IGrainContext CreateContext(GrainAddress address, IConfigureGrainContext[ return context; } } + + private sealed class TestPreparedGrainContextActivator( + IGrainContext context, + List events) : IPreparedGrainContextActivator + { + public IGrainContext CreateContext(GrainAddress address, IConfigureGrainContext[] configureActions) + { + var preparedContext = CreatePreparedContext(address, configureActions); + using var startup = preparedContext.Start(); + return preparedContext.Context; + } + + public PreparedGrainContext CreatePreparedContext( + GrainAddress address, + IConfigureGrainContext[] configureActions) + { + foreach (var configure in configureActions) + { + configure.Configure(context); + } + + events.Add("create"); + return new(context, new TestGrainContextStartup(events)); + } + } + + private sealed class TestGrainContextStartup(List events) : IGrainContextStartup + { + public IDisposable Start() + { + events.Add("start"); + return new TestStartupLease(events); + } + + public void Abort() => events.Add("abort"); + } + + private sealed class TestStartupLease(List events) : IDisposable + { + private int _disposed; + + public void Dispose() + { + if (Interlocked.Exchange(ref _disposed, 1) == 0) + { + events.Add("release"); + } + } + } + + private sealed class ThrowingGrainContextStartup( + List events, + Exception exception) : IGrainContextStartup + { + public IDisposable Start() + { + events.Add("start"); + throw exception; + } + + public void Abort() => events.Add("abort"); + } + + private sealed class ThrowingStartAndAbortGrainContextStartup( + List events, + Exception startupException, + Exception abortException) : IGrainContextStartup + { + public IDisposable Start() + { + events.Add("start"); + throw startupException; + } + + public void Abort() + { + events.Add("abort"); + throw abortException; + } + } } diff --git a/test/Orleans.Core.Tests/SchedulerTests/OrleansTaskSchedulerBasicTests.cs b/test/Orleans.Core.Tests/SchedulerTests/OrleansTaskSchedulerBasicTests.cs index cce2d6e5aa2..9800c25768a 100644 --- a/test/Orleans.Core.Tests/SchedulerTests/OrleansTaskSchedulerBasicTests.cs +++ b/test/Orleans.Core.Tests/SchedulerTests/OrleansTaskSchedulerBasicTests.cs @@ -116,6 +116,181 @@ public async Task Async_Task_Start_ActivationTaskScheduler() Assert.Equal(expected, received); } + [Fact] + public async Task Sched_ActivationStartup_DefersQueuedWorkUntilDisposed() + { + Task? queuedTask = null; + var startCompleted = 0; + var queuedTaskObservedStartCompletion = 0; + var startTask = new Task(() => + { + Assert.Same(_rootContext, RuntimeContext.Current); + Assert.Same(_rootContext.WorkItemGroup.TaskScheduler, TaskScheduler.Current); + queuedTask = new Task( + () => + { + Volatile.Write( + ref queuedTaskObservedStartCompletion, + Volatile.Read(ref startCompleted)); + }); + queuedTask.Start(_rootContext.WorkItemGroup.TaskScheduler); + + Assert.Equal(1, _rootContext.WorkItemGroup.ExternalWorkItemCount); + Assert.False(queuedTask.IsCompleted); + Volatile.Write(ref startCompleted, 1); + }); + + using (var startup = _rootContext.WorkItemGroup.BeginActivationStartup()) + { + startup.RunConstructor(startTask); + Assert.True(startTask.IsCompletedSuccessfully, startTask.Exception?.ToString()); + Assert.False(queuedTask!.IsCompleted); + Assert.Throws( + () => startup.RunConstructor(new Task(static () => { }))); + Assert.Throws(startup.Abort); + } + + await queuedTask.WaitAsync(TimeSpan.FromSeconds(5), TestContext.Current.CancellationToken); + Assert.Equal(1, queuedTaskObservedStartCompletion); + } + + [Fact] + public async Task Sched_ActivationStartup_PreservesContextForReleasedAsynchronousWork() + { + const int IterationCount = 1_000; + var signal = new SingleWaiterAutoResetEvent { RunContinuationsAsynchronously = true }; + var observations = new TaskCompletionSource[IterationCount]; + for (var i = 0; i < observations.Length; i++) + { + observations[i] = new(TaskCreationOptions.RunContinuationsAsynchronously); + } + + Task? observationLoop = null; + var initialObservation = 0; + var asyncInitialObservation = 0; + var loopStarted = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var loopStarter = new Task(() => + { + observationLoop = ObserveSignals(); + loopStarted.SetResult(); + }); + var startTask = new Task(() => + { + if (ReferenceEquals(RuntimeContext.Current, _rootContext)) + { + initialObservation |= 1; + } + + if (ReferenceEquals(TaskScheduler.Current, _rootContext.WorkItemGroup.TaskScheduler)) + { + initialObservation |= 2; + } + + loopStarter.Start(_rootContext.WorkItemGroup.TaskScheduler); + }); + using (var startup = _rootContext.WorkItemGroup.BeginActivationStartup()) + { + startup.RunConstructor(startTask); + Assert.True(startTask.IsCompletedSuccessfully, startTask.Exception?.ToString()); + Assert.Equal(3, initialObservation); + } + await loopStarted.Task.WaitAsync(TimeSpan.FromSeconds(5), TestContext.Current.CancellationToken); + Assert.Equal(3, asyncInitialObservation); + + for (var i = 0; i < observations.Length; i++) + { + signal.Signal(); + Assert.Equal( + 3, + await observations[i].Task.WaitAsync( + TimeSpan.FromSeconds(5), + TestContext.Current.CancellationToken)); + } + + await observationLoop!.WaitAsync(TimeSpan.FromSeconds(5), TestContext.Current.CancellationToken); + + async Task ObserveSignals() + { + if (ReferenceEquals(RuntimeContext.Current, _rootContext)) + { + asyncInitialObservation |= 1; + } + + if (ReferenceEquals(TaskScheduler.Current, _rootContext.WorkItemGroup.TaskScheduler)) + { + asyncInitialObservation |= 2; + } + + for (var i = 0; i < observations.Length; i++) + { + await signal.WaitAsync(); + var observation = 0; + if (ReferenceEquals(RuntimeContext.Current, _rootContext)) + { + observation |= 1; + } + + if (ReferenceEquals(TaskScheduler.Current, _rootContext.WorkItemGroup.TaskScheduler)) + { + observation |= 2; + } + + observations[i].SetResult(observation); + } + } + } + + [Fact] + public async Task Sched_ActivationStartup_DisposeIsIdempotent() + { + var queuedWorkCount = 0; + var subsequentWorkCount = 0; + var queuedTask = new Task(() => Interlocked.Increment(ref queuedWorkCount)); + var startup = _rootContext.WorkItemGroup.BeginActivationStartup(); + queuedTask.Start(_rootContext.WorkItemGroup.TaskScheduler); + + startup.Dispose(); + await queuedTask.WaitAsync(TimeSpan.FromSeconds(5), TestContext.Current.CancellationToken); + Assert.Equal(1, Volatile.Read(ref queuedWorkCount)); + Assert.Equal(0, _rootContext.WorkItemGroup.ExternalWorkItemCount); + + startup.Dispose(); + Assert.Equal(1, Volatile.Read(ref queuedWorkCount)); + Assert.Equal(0, _rootContext.WorkItemGroup.ExternalWorkItemCount); + + var subsequentTask = new Task(() => Interlocked.Increment(ref subsequentWorkCount)); + subsequentTask.Start(_rootContext.WorkItemGroup.TaskScheduler); + await subsequentTask.WaitAsync(TimeSpan.FromSeconds(5), TestContext.Current.CancellationToken); + + Assert.Equal(1, Volatile.Read(ref queuedWorkCount)); + Assert.Equal(1, Volatile.Read(ref subsequentWorkCount)); + Assert.Equal(0, _rootContext.WorkItemGroup.ExternalWorkItemCount); + } + + [Fact] + public async Task Sched_ActivationStartup_AbortDiscardsQueuedWorkAndAllowsReuse() + { + var discardedWorkCount = 0; + var subsequentWorkCount = 0; + var discardedTask = new Task(() => Interlocked.Increment(ref discardedWorkCount)); + var startup = _rootContext.WorkItemGroup.BeginActivationStartup(); + discardedTask.Start(_rootContext.WorkItemGroup.TaskScheduler); + Assert.Equal(1, _rootContext.WorkItemGroup.ExternalWorkItemCount); + + startup.Abort(); + + Assert.Equal(0, Volatile.Read(ref discardedWorkCount)); + Assert.Equal(0, _rootContext.WorkItemGroup.ExternalWorkItemCount); + + var subsequentTask = new Task(() => Interlocked.Increment(ref subsequentWorkCount)); + subsequentTask.Start(_rootContext.WorkItemGroup.TaskScheduler); + await subsequentTask.WaitAsync(TimeSpan.FromSeconds(5), TestContext.Current.CancellationToken); + + Assert.Equal(0, Volatile.Read(ref discardedWorkCount)); + Assert.Equal(1, Volatile.Read(ref subsequentWorkCount)); + Assert.Equal(0, _rootContext.WorkItemGroup.ExternalWorkItemCount); + } + [Fact] public async Task Sched_SimpleFifoTest() { @@ -443,6 +618,58 @@ internal static ILoggerFactory InitSchedulerLogging() var loggerFactory = TestingUtils.CreateDefaultLoggerFactory(TestingUtils.CreateTraceFileName("Silo", DateTime.UtcNow.ToString("yyyyMMdd_hhmmss")), filters); return loggerFactory; } + + [Fact] + public async Task Sched_ActivationStartup_ConcurrentReleaseExecutesQueuedWorkExactlyOnce() + { + const int ParticipantCount = 8; + var queuedWorkCount = 0; + var subsequentWorkCount = 0; + var queuedWork = new Task(() => Interlocked.Increment(ref queuedWorkCount)); + var startup = _rootContext.WorkItemGroup.BeginActivationStartup(); + queuedWork.Start(_rootContext.WorkItemGroup.TaskScheduler); + + var releaseParticipants = new Task[ParticipantCount]; + var participantReady = new TaskCompletionSource[ParticipantCount]; + var releaseParticipantsBarrier = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + for (var i = 0; i < ParticipantCount; i++) + { + var participant = participantReady[i] = new(TaskCreationOptions.RunContinuationsAsynchronously); + releaseParticipants[i] = Task.Run( + async () => + { + participant.SetResult(); + await releaseParticipantsBarrier.Task; + startup.Dispose(); + }, + TestContext.Current.CancellationToken); + } + + try + { + await Task.WhenAll(Array.ConvertAll(participantReady, static participant => participant.Task)) + .WaitAsync(TimeSpan.FromSeconds(5), TestContext.Current.CancellationToken); + Assert.False(queuedWork.IsCompleted); + + releaseParticipantsBarrier.SetResult(); + await Task.WhenAll(releaseParticipants) + .WaitAsync(TimeSpan.FromSeconds(5), TestContext.Current.CancellationToken); + await queuedWork.WaitAsync(TimeSpan.FromSeconds(5), TestContext.Current.CancellationToken); + Assert.Equal(1, Volatile.Read(ref queuedWorkCount)); + + var subsequentWork = new Task(() => Interlocked.Increment(ref subsequentWorkCount)); + subsequentWork.Start(_rootContext.WorkItemGroup.TaskScheduler); + await subsequentWork.WaitAsync(TimeSpan.FromSeconds(5), TestContext.Current.CancellationToken); + Assert.Equal(1, Volatile.Read(ref subsequentWorkCount)); + } + finally + { + releaseParticipantsBarrier.TrySetResult(); + startup.Dispose(); + await Task.WhenAll(releaseParticipants) + .WaitAsync(TimeSpan.FromSeconds(5), TestContext.Current.CancellationToken); + } + } } } diff --git a/test/Orleans.Runtime.Internal.Tests/ActivationsLifeCycleTests/ActivationDataMigrationTests.cs b/test/Orleans.Runtime.Internal.Tests/ActivationsLifeCycleTests/ActivationDataMigrationTests.cs index bdc810b4df0..41d7256a7b9 100644 --- a/test/Orleans.Runtime.Internal.Tests/ActivationsLifeCycleTests/ActivationDataMigrationTests.cs +++ b/test/Orleans.Runtime.Internal.Tests/ActivationsLifeCycleTests/ActivationDataMigrationTests.cs @@ -1,15 +1,21 @@ #nullable enable using Microsoft.Extensions.DependencyInjection; +using Orleans.CodeGeneration; +using Orleans.Metadata; using Orleans.Runtime; +using Orleans.Runtime.Diagnostics; +using Orleans.Serialization.Invocation; using Orleans.TestingHost; using TestExtensions; using UnitTests.GrainInterfaces; +using UnitTests.Grains; using Xunit; namespace UnitTests.ActivationsLifeCycleTests; [TestSuite("BVT")] [TestProvider("None")] +[TestArea("Runtime")] [TestCategory("BVT"), TestCategory("Migration")] public class ActivationDataMigrationTests(ActivationDataMigrationTests.Fixture fixture) : IClassFixture { @@ -53,6 +59,491 @@ private async Task GetActivation(CancellationToken cancellationT return Assert.IsType(directory.FindTarget(grainId)); } + [Fact] + public async Task TryStartMigrationReturnsFalseDuringSynchronousActivationStartup() + { + var startupFixture = new SynchronousMigrationFixture(); + await startupFixture.InitializeAsync(); + var (grainId, scenario) = startupFixture.CreateScenario(); + using var lifecycleSubscription = startupFixture.ObserveLifecycle(); + Task? startTask = null; + ActivationData? context = null; + + try + { + startTask = Task.Run(() => startupFixture.StartActivation(grainId)); + var observation = await startupFixture.Probe.Observation.WaitAsync( + TimeSpan.FromSeconds(30), + TestContext.Current.CancellationToken); + context = observation.Context; + + Assert.False(observation.Result); + Assert.Equal(ActivationState.Creating, context.State); + Assert.Same(context, startupFixture.ActivationDirectory.FindTarget(grainId)); + Assert.Equal(1, scenario.CreateCount); + Assert.Equal(1, scenario.ConstructorCount); + Assert.Equal(0, scenario.OnActivateCount); + Assert.Equal(0, scenario.GetEventCount("Deactivating")); + Assert.Equal(0, scenario.GetEventCount("Deactivated")); + Assert.Equal(0, scenario.DisposeStartedCount); + + startupFixture.Probe.Release(); + Assert.Same( + context, + await startTask.WaitAsync( + TimeSpan.FromSeconds(30), + TestContext.Current.CancellationToken)); + await scenario.WaitForEvent("Activated").WaitAsync( + TimeSpan.FromSeconds(30), + TestContext.Current.CancellationToken); + var (message, request) = startupFixture.CreateRequest( + context, + scenario, + payload: "synchronous-migration", + requestContextValue: "synchronous-migration-context"); + context.ReceiveMessage(message); + await scenario.WaitForEvent("Response").WaitAsync( + TimeSpan.FromSeconds(30), + TestContext.Current.CancellationToken); + + Assert.Equal( + "synchronous-migration:synchronous-migration-context", + request.Result); + Assert.Equal(ActivationState.Valid, context.State); + Assert.Equal(1, scenario.GetEventCount("Created")); + Assert.Equal(1, scenario.GetEventCount("Activated")); + Assert.Equal(1, scenario.OnActivateCount); + Assert.Equal(1, scenario.RequestInvocationCount); + Assert.Same(context, startupFixture.ActivationDirectory.FindTarget(grainId)); + Assert.Single(scenario.Events.Select(entry => entry.ActivationId).Distinct()); + Assert.All(scenario.Events, entry => Assert.Equal(grainId, entry.GrainId)); + AssertEventOrder( + scenario, + "SynchronousMigrationAttempted", + "Created", + "LifecycleStartLow", + "LifecycleStartHigh", + "OnActivateEntered", + "OnActivateCompleted", + "Activated", + "RequestInvoked", + "Response"); + } + finally + { + startupFixture.Probe.Release(); + scenario.ReleaseActivation(); + scenario.ReleaseDisposal(); + await AwaitIgnoringFailure(startTask); + await CleanupAsync(startupFixture.ActivationDirectory, context, scenario); + startupFixture.Hooks.RemoveScenario(grainId); + await startupFixture.DisposeAsync(); + } + } + + [Fact] + public async Task TryStartMigrationReturnsFalseDuringAsyncActivationStartupAndSucceedsAfterActivation() + { + var startupFixture = new ActivationStartupTestFixture(); + await startupFixture.InitializeAsync(); + var (grainId, scenario) = startupFixture.CreateScenario( + ActivationStartupCompletion.AsynchronousSuccess, + ActivationStartupDisposal.Synchronous); + using var lifecycleProbe = new MigrationLifecycleProbe(scenario); + ActivationData? context = null; + + try + { + context = startupFixture.StartActivation(grainId); + await scenario.WaitForEvent("OnActivateEntered").WaitAsync( + TimeSpan.FromSeconds(30), + TestContext.Current.CancellationToken); + + Assert.Equal(ActivationState.Activating, context.State); + Assert.False( + context.TryStartMigration( + requestContext: null, + TestContext.Current.CancellationToken)); + Assert.Equal(ActivationState.Activating, context.State); + Assert.Same(context, startupFixture.ActivationDirectory.FindTarget(grainId)); + Assert.DoesNotContain( + lifecycleProbe.Events, + static evt => evt is GrainLifecycleEvents.Deactivating); + Assert.DoesNotContain( + lifecycleProbe.Events, + static evt => evt is GrainLifecycleEvents.Deactivated); + Assert.Equal(0, scenario.DisposeStartedCount); + Assert.Equal(0, scenario.ScopeDisposeCount); + + scenario.ReleaseActivation(); + await scenario.WaitForEvent("Activated").WaitAsync( + TimeSpan.FromSeconds(30), + TestContext.Current.CancellationToken); + Assert.Equal(ActivationState.Valid, context.State); + + var deactivating = scenario.WaitForEvent("Deactivating"); + var deactivated = scenario.WaitForEvent("Deactivated"); + Assert.True( + context.TryStartMigration( + requestContext: null, + TestContext.Current.CancellationToken)); + await deactivating.WaitAsync( + TimeSpan.FromSeconds(30), + TestContext.Current.CancellationToken); + + Assert.Equal(ActivationState.Deactivating, context.State); + var migration = Assert.Single( + lifecycleProbe.Events.OfType()); + Assert.Same(context, migration.GrainContext); + Assert.Equal(DeactivationReasonCode.Migrating, migration.Reason.ReasonCode); + + await deactivated.WaitAsync( + TimeSpan.FromSeconds(30), + TestContext.Current.CancellationToken); + Assert.Equal(ActivationState.Invalid, context.State); + Assert.Null(startupFixture.ActivationDirectory.FindTarget(grainId)); + Assert.Equal(1, scenario.CreateCount); + Assert.Equal(1, scenario.ConstructorCount); + Assert.Equal(1, scenario.OnActivateCount); + Assert.Equal(1, scenario.OnDeactivateCount); + Assert.Equal(1, scenario.GetEventCount("Created")); + Assert.Equal(1, scenario.GetEventCount("Activated")); + Assert.Equal(1, scenario.GetEventCount("Deactivating")); + Assert.Equal(1, scenario.GetEventCount("Deactivated")); + Assert.Equal(1, scenario.DisposeStartedCount); + Assert.Equal(1, scenario.DisposeCompletedCount); + Assert.Equal(1, scenario.ScopeDisposeCount); + Assert.Single(scenario.Events.Select(entry => entry.ActivationId).Distinct()); + Assert.All(scenario.Events, entry => Assert.Equal(grainId, entry.GrainId)); + AssertEventOrder( + scenario, + "Created", + "LifecycleStartLow", + "LifecycleStartHigh", + "OnActivateEntered", + "OnActivateCompleted", + "Activated", + "Deactivating", + "OnDeactivateEntered", + "OnDeactivateCompleted", + "LifecycleStopHigh", + "LifecycleStopLow", + "ActivatorDisposeStarted", + "ActivatorDisposeCompleted", + "ScopeDisposed", + "Deactivated"); + } + finally + { + scenario.ReleaseActivation(); + scenario.ReleaseDisposal(); + await CleanupAsync(startupFixture.ActivationDirectory, context, scenario); + startupFixture.RemoveScenario(grainId); + await startupFixture.DisposeAsync(); + } + } + + private static void AssertEventOrder( + ActivationStartupScenario scenario, + params string[] expected) + { + var expectedNames = expected.ToHashSet(); + Assert.Equal( + expected, + scenario.Events + .Where(entry => expectedNames.Contains(entry.Name)) + .Select(entry => entry.Name)); + Assert.All(expected, name => Assert.Equal(1, scenario.GetEventCount(name))); + } + + private static async Task CleanupAsync( + ActivationDirectory directory, + ActivationData? context, + ActivationStartupScenario scenario) + { + if (context is null || context.State is ActivationState.Invalid) + { + return; + } + + var deactivated = scenario.WaitForEvent("Deactivated"); + context.Deactivate(new( + DeactivationReasonCode.ApplicationRequested, + "Migration startup test cleanup.")); + await deactivated.WaitAsync( + TimeSpan.FromSeconds(30), + TestContext.Current.CancellationToken); + Assert.Null(directory.FindTarget(context.GrainId)); + } + + private static async Task AwaitIgnoringFailure(Task? task) + { + if (task is null) + { + return; + } + + try + { + await task.WaitAsync( + TimeSpan.FromSeconds(30), + TestContext.Current.CancellationToken); + } + catch + { + } + } + + private sealed class MigrationLifecycleProbe : + IObserver, + IDisposable + { + private readonly ActivationStartupScenario _scenario; + private readonly IDisposable _subscription; + private readonly List _events = []; + + public MigrationLifecycleProbe(ActivationStartupScenario scenario) + { + _scenario = scenario; + _subscription = GrainLifecycleEvents.AllEvents.Subscribe(this); + } + + public IReadOnlyList Events + { + get + { + lock (_events) + { + return _events.ToArray(); + } + } + } + + public void OnCompleted() + { + } + + public void OnError(Exception error) + { + } + + public void OnNext(GrainLifecycleEvents.LifecycleEvent value) + { + if (!value.GrainContext.GrainId.Equals(_scenario.GrainId)) + { + return; + } + + lock (_events) + { + _events.Add(value); + } + + var name = value switch + { + GrainLifecycleEvents.Created => "Created", + GrainLifecycleEvents.Activated => "Activated", + GrainLifecycleEvents.Deactivating => "Deactivating", + GrainLifecycleEvents.Deactivated => "Deactivated", + _ => null, + }; + if (name is not null) + { + _scenario.Record(name, value.GrainContext); + } + } + + public void Dispose() => _subscription.Dispose(); + } + + private sealed class SynchronousMigrationFixture : BaseTestClusterFixture + { + protected override void ConfigureTestCluster(TestClusterBuilder builder) + { + builder.Options.InitialSilosCount = 1; + builder.AddSiloBuilderConfigurator(); + } + + private InProcessSiloHandle PrimarySilo => (InProcessSiloHandle)HostedCluster.Primary!; + + private IServiceProvider Services => PrimarySilo.SiloHost.Services; + + public ActivationStartupTestHooks Hooks => + Services.GetRequiredService(); + + public SynchronousMigrationProbe Probe => + Services.GetRequiredService(); + + public ActivationDirectory ActivationDirectory => + Services.GetRequiredService(); + + public (GrainId GrainId, ActivationStartupScenario Scenario) CreateScenario() + { + var grainType = Services.GetRequiredService() + .GetGrainType(typeof(ActivationStartupTestGrain)); + var grainId = GrainId.Create(grainType, Guid.NewGuid().ToString("N")); + return ( + grainId, + Hooks.CreateScenario( + grainId, + ActivationStartupCompletion.ImmediateSuccess, + ActivationStartupDisposal.Synchronous)); + } + + public ActivationData StartActivation(GrainId grainId) => + Assert.IsType( + Services.GetRequiredService() + .GetOrCreateActivation( + grainId, + requestContextData: null, + rehydrationContext: null)); + + public (Message Message, ActivationStartupRequest Request) CreateRequest( + ActivationData context, + ActivationStartupScenario scenario, + string payload, + string requestContextValue) + { + var request = new ActivationStartupRequest(scenario, payload, recordResponse: true); + var message = Services.GetRequiredService() + .CreateMessage(request, InvokeMethodOptions.OneWay); + message.SetInfiniteTimeToLive(); + message.RequestContextData = new() + { + [ActivationStartupTestHooks.RequestContextKey] = requestContextValue, + }; + message.SendingGrain = GrainId.Create( + "migration-startup-sender", + Guid.NewGuid().ToString("N")); + message.SendingSilo = PrimarySilo.SiloAddress; + message.TargetGrain = context.GrainId; + message.TargetSilo = PrimarySilo.SiloAddress; + return (message, request); + } + + public IDisposable ObserveLifecycle() => + GrainLifecycleEvents.AllEvents.Subscribe(new LifecycleObserver(Hooks)); + + private sealed class SiloConfigurator : ISiloConfigurator + { + public void Configure(ISiloBuilder hostBuilder) + { + hostBuilder.Services.AddSingleton(); + hostBuilder.Services.AddSingleton(); + hostBuilder.Services.AddScoped(); + hostBuilder.Services.AddSingleton(); + hostBuilder.Services.AddSingleton< + IConfigureGrainTypeComponents, + ActivationStartupConfigurator>(); + } + } + + private sealed class ActivationStartupConfigurator( + GrainClassMap grainClassMap, + SynchronousMigrationActivator activator) + : IConfigureGrainTypeComponents + { + public void Configure( + GrainType grainType, + GrainProperties properties, + GrainTypeSharedContext shared) + { + if (grainClassMap.TryGetGrainClass(grainType, out var grainClass) + && grainClass == typeof(ActivationStartupTestGrain)) + { + shared.SetComponent(activator); + } + } + } + + private sealed class LifecycleObserver(ActivationStartupTestHooks hooks) + : IObserver + { + public void OnCompleted() + { + } + + public void OnError(Exception error) + { + } + + public void OnNext(GrainLifecycleEvents.LifecycleEvent value) + { + if (!hooks.TryGetScenario(value.GrainContext.GrainId, out var scenario)) + { + return; + } + + var name = value switch + { + GrainLifecycleEvents.Created => "Created", + GrainLifecycleEvents.Activated => "Activated", + GrainLifecycleEvents.Deactivating => "Deactivating", + GrainLifecycleEvents.Deactivated => "Deactivated", + _ => null, + }; + if (name is not null) + { + scenario!.Record(name, value.GrainContext); + } + } + } + } + + private sealed class SynchronousMigrationProbe + { + private readonly TaskCompletionSource<(ActivationData Context, bool Result)> _observation = + new(TaskCreationOptions.RunContinuationsAsynchronously); + private readonly TaskCompletionSource _release = + new(TaskCreationOptions.RunContinuationsAsynchronously); + + public Task<(ActivationData Context, bool Result)> Observation => _observation.Task; + + public void ObserveAndWait(ActivationData context, bool result) + { + _observation.TrySetResult((context, result)); + try + { + _release.Task.WaitAsync(TimeSpan.FromSeconds(30)).GetAwaiter().GetResult(); + } + catch (TimeoutException exception) + { + throw new TimeoutException( + $"Timed out releasing synchronous migration probe for grain '{context.GrainId}', activation '{context.ActivationId}', result '{result}'.", + exception); + } + } + + public void Release() => _release.TrySetResult(); + } + + private sealed class SynchronousMigrationActivator( + ActivationStartupTestHooks hooks, + SynchronousMigrationProbe probe) : IGrainActivator + { + public object CreateInstance(IGrainContext context) + { + var scenario = hooks.GetRequiredScenario(context.GrainId); + scenario.ObserveCreate(context); + context.ActivationServices.GetRequiredService() + .Attach(context); + var instance = new ActivationStartupTestGrain(context, hooks); + var activation = Assert.IsType(context); + var result = activation.TryStartMigration(requestContext: null); + scenario.Record("SynchronousMigrationAttempted", context); + probe.ObserveAndWait(activation, result); + return instance; + } + + public ValueTask DisposeInstance(IGrainContext context, object instance) + { + var scenario = hooks.GetRequiredScenario(context.GrainId); + scenario.ObserveDisposeStarted(context); + scenario.ObserveDisposeCompleted(context); + return ValueTask.CompletedTask; + } + } + public class Fixture : BaseTestClusterFixture { protected override void ConfigureTestCluster(TestClusterBuilder builder) diff --git a/test/Orleans.Runtime.Internal.Tests/ActivationsLifeCycleTests/ActivationStartupTestFixture.cs b/test/Orleans.Runtime.Internal.Tests/ActivationsLifeCycleTests/ActivationStartupTestFixture.cs new file mode 100644 index 00000000000..9cda1796338 --- /dev/null +++ b/test/Orleans.Runtime.Internal.Tests/ActivationsLifeCycleTests/ActivationStartupTestFixture.cs @@ -0,0 +1,207 @@ +using System.Reflection; +using Microsoft.Extensions.DependencyInjection; +using Orleans.CodeGeneration; +using Orleans.Metadata; +using Orleans.Runtime; +using Orleans.Runtime.Diagnostics; +using Orleans.Serialization.Invocation; +using Orleans.TestingHost; +using TestExtensions; +using UnitTests.Grains; +using Xunit; + +namespace UnitTests.ActivationsLifeCycleTests; + +public sealed class ActivationStartupTestFixture : BaseTestClusterFixture +{ + protected override void ConfigureTestCluster(TestClusterBuilder builder) + { + builder.Options.InitialSilosCount = 1; + builder.AddSiloBuilderConfigurator(); + } + + public InProcessSiloHandle PrimarySilo => (InProcessSiloHandle)HostedCluster.Primary!; + + public IServiceProvider Services => PrimarySilo.SiloHost.Services; + + public ActivationStartupTestHooks Hooks => Services.GetRequiredService(); + + internal ActivationDirectory ActivationDirectory => Services.GetRequiredService(); + + public (GrainId GrainId, ActivationStartupScenario Scenario) CreateScenario( + ActivationStartupCompletion completion, + ActivationStartupDisposal disposal) + { + var grainType = Services.GetRequiredService().GetGrainType(typeof(ActivationStartupTestGrain)); + var grainId = GrainId.Create(grainType, Guid.NewGuid().ToString("N")); + return (grainId, Hooks.CreateScenario(grainId, completion, disposal)); + } + + internal ActivationData StartActivation(GrainId grainId, string? requestContextValue = null) + { + Dictionary? requestContext = requestContextValue is null + ? null + : new() { [ActivationStartupTestHooks.RequestContextKey] = requestContextValue }; + return Assert.IsType( + Services.GetRequiredService().GetOrCreateActivation(grainId, requestContext, rehydrationContext: null)); + } + + internal (Message Message, ActivationStartupRequest Request) CreateRequest( + ActivationData context, + ActivationStartupScenario scenario, + string payload, + string requestContextValue, + bool recordResponse, + InvokeMethodOptions invokeMethodOptions = InvokeMethodOptions.OneWay) + { + var request = new ActivationStartupRequest(scenario, payload, recordResponse); + var message = Services.GetRequiredService().CreateMessage(request, invokeMethodOptions); + message.SetInfiniteTimeToLive(); + message.RequestContextData = new() + { + [ActivationStartupTestHooks.RequestContextKey] = requestContextValue, + }; + message.SendingGrain = GrainId.Create("activation-startup-sender", Guid.NewGuid().ToString("N")); + message.SendingSilo = PrimarySilo.SiloAddress; + message.TargetGrain = context.GrainId; + message.TargetSilo = PrimarySilo.SiloAddress; + return (message, request); + } + + public IDisposable ObserveLifecycle() => + GrainLifecycleEvents.AllEvents.Subscribe(new LifecycleObserver(Hooks)); + + public void RemoveScenario(GrainId grainId) => Hooks.RemoveScenario(grainId); + + private sealed class SiloConfigurator : ISiloConfigurator + { + public void Configure(ISiloBuilder hostBuilder) + { + hostBuilder.Services.AddSingleton(); + hostBuilder.Services.AddScoped(); + hostBuilder.Services.AddSingleton(); + } + } + + private sealed class ActivationStartupTestActivator( + GrainClassMap grainClassMap, + ActivationStartupTestHooks hooks) : IGrainActivator, IConfigureGrainTypeComponents + { + public void Configure(GrainType grainType, GrainProperties properties, GrainTypeSharedContext shared) + { + if (grainClassMap.TryGetGrainClass(grainType, out var grainClass) + && grainClass == typeof(ActivationStartupTestGrain)) + { + shared.SetComponent(this); + } + } + + public object CreateInstance(IGrainContext context) + { + var scenario = hooks.GetRequiredScenario(context.GrainId); + scenario.ObserveCreate(context); + context.ActivationServices.GetRequiredService().Attach(context); + return new ActivationStartupTestGrain(context, hooks); + } + + public async ValueTask DisposeInstance(IGrainContext context, object instance) + { + var scenario = hooks.GetRequiredScenario(context.GrainId); + scenario.ObserveDisposeStarted(context); + if (scenario.Disposal is ActivationStartupDisposal.Asynchronous) + { + await scenario.DisposalRelease; + } + + scenario.ObserveDisposeCompleted(context); + } + } + + private sealed class LifecycleObserver(ActivationStartupTestHooks hooks) + : IObserver + { + public void OnCompleted() + { + } + + public void OnError(Exception error) + { + } + + public void OnNext(GrainLifecycleEvents.LifecycleEvent value) + { + if (!hooks.TryGetScenario(value.GrainContext.GrainId, out var scenario)) + { + return; + } + + var name = value switch + { + GrainLifecycleEvents.Created => "Created", + GrainLifecycleEvents.Activated => "Activated", + GrainLifecycleEvents.Deactivating => "Deactivating", + GrainLifecycleEvents.Deactivated => "Deactivated", + _ => null, + }; + if (name is not null) + { + scenario!.Record(name, value.GrainContext); + } + } + } +} + +internal sealed class ActivationStartupRequest( + ActivationStartupScenario scenario, + string payload, + bool recordResponse) : IInvokable +{ + private static readonly MethodInfo Method = + typeof(IActivationStartupTestGrain).GetMethod(nameof(IActivationStartupTestGrain.Invoke))!; + + private IActivationStartupTestGrain? _target; + private string? _result; + + public string? Result => Volatile.Read(ref _result); + + public object? GetTarget() => _target; + + public void SetTarget(ITargetHolder holder) + { + _target = (IActivationStartupTestGrain)holder.GetTarget()!; + } + + public async ValueTask Invoke() + { + var result = await _target!.Invoke(payload); + Volatile.Write(ref _result, result); + if (recordResponse) + { + scenario.Record("Response", ((IGrainBase)_target).GrainContext); + } + + return Response.FromResult(result); + } + + public int GetArgumentCount() => 1; + + public object? GetArgument(int index) => + index == 0 ? payload : throw new ArgumentOutOfRangeException(nameof(index)); + + public void SetArgument(int index, object value) => + throw new NotSupportedException("The activation startup request is immutable."); + + public string GetMethodName() => nameof(IActivationStartupTestGrain.Invoke); + + public string GetInterfaceName() => typeof(IActivationStartupTestGrain).FullName!; + + public string GetActivityName() => $"{GetInterfaceName()}/{GetMethodName()}"; + + public MethodInfo GetMethod() => Method; + + public Type GetInterfaceType() => typeof(IActivationStartupTestGrain); + + public void Dispose() + { + } +} diff --git a/test/Orleans.Runtime.Internal.Tests/ActivationsLifeCycleTests/ActivationStartupTests.cs b/test/Orleans.Runtime.Internal.Tests/ActivationsLifeCycleTests/ActivationStartupTests.cs new file mode 100644 index 00000000000..fb928163331 --- /dev/null +++ b/test/Orleans.Runtime.Internal.Tests/ActivationsLifeCycleTests/ActivationStartupTests.cs @@ -0,0 +1,682 @@ +using System.Diagnostics; +using Microsoft.Extensions.DependencyInjection; +using Orleans.CodeGeneration; +using Orleans.Diagnostics; +using Orleans.Runtime; +using Orleans.Runtime.Diagnostics; +using Orleans.TestingHost.Diagnostics; +using TestExtensions; +using UnitTests.Grains; +using Xunit; + +namespace UnitTests.ActivationsLifeCycleTests; + +[TestSuite("BVT")] +[TestProvider("None")] +[TestArea("Runtime")] +public sealed class ActivationStartupTests(ActivationStartupTestFixture fixture) + : IClassFixture +{ + private static readonly TimeSpan Timeout = TimeSpan.FromSeconds(30); + + [Fact] + public async Task FailureBeforeStartAbortsRejectsRequestAndUnregistersPreparedContext() + { + var (grainId, scenario) = fixture.CreateScenario( + ActivationStartupCompletion.ImmediateSuccess, + ActivationStartupDisposal.Synchronous); + var failActivityCreation = new AsyncLocal(); + var expected = new InvalidOperationException("activity-start-fault"); + using var collector = new DiagnosticEventCollector(DispatcherEvents.ListenerName); + var activationCollector = fixture.Services.GetRequiredService(); + var baselineActivationCount = activationCollector._activationCount; + ActivationData? context = null; + ActivationStartupRequest? request = null; + Task? rejected = null; + using var listener = new ActivityListener + { + ShouldListenTo = static source => + source.Name == ActivitySources.LifecycleActivitySourceName, + Sample = (ref ActivityCreationOptions _) => OnSample(), + SampleUsingParentId = (ref ActivityCreationOptions _) => OnSample(), + }; + ActivitySource.AddActivityListener(listener); + + try + { + failActivityCreation.Value = true; + var actual = Assert.Throws(() => fixture.StartActivation(grainId)); + + Assert.Same(expected, actual); + Assert.NotNull(context); + await context.Deactivated.WaitAsync(Timeout, TestContext.Current.CancellationToken); + Assert.NotNull(rejected); + var rejectionEvent = await rejected; + var rejection = Assert.IsType(rejectionEvent.Payload); + Assert.Same(request, rejection.Message.BodyObject); + Assert.Equal(Message.RejectionTypes.Transient, rejection.RejectionType); + Assert.Equal( + "Activation startup was aborted.", + Assert.IsType(rejection.Exception).Message); + Assert.Null(fixture.ActivationDirectory.FindTarget(grainId)); + Assert.Equal(0, scenario.CreateCount); + Assert.Equal(0, scenario.ConstructorCount); + Assert.Equal(0, scenario.OnActivateCount); + Assert.Equal(ActivationState.Invalid, context.State); + Assert.Equal(baselineActivationCount, activationCollector._activationCount); + + var lateRequestData = fixture.CreateRequest( + context, + scenario, + payload: "late-must-not-run", + requestContextValue: "late-pre-start", + recordResponse: true); + var lateRejected = collector.WaitForEventAsync( + nameof(DispatcherEvents.Rejected), + diagnosticEvent => diagnosticEvent.Payload is DispatcherEvents.Rejected rejection + && ReferenceEquals(rejection.Message, lateRequestData.Message), + Timeout, + TestContext.Current.CancellationToken); + context.ReceiveMessage(lateRequestData.Message); + var lateRejection = Assert.IsType((await lateRejected).Payload); + Assert.Equal(Message.RejectionTypes.Transient, lateRejection.RejectionType); + Assert.Equal( + "Activation startup was aborted.", + Assert.IsType(lateRejection.Exception).Message); + Assert.Equal(0, scenario.RequestInvocationCount); + } + finally + { + failActivityCreation.Value = false; + fixture.RemoveScenario(grainId); + } + + ActivitySamplingResult OnSample() + { + if (!failActivityCreation.Value) + { + return ActivitySamplingResult.None; + } + + context = Assert.IsType(fixture.ActivationDirectory.FindTarget(grainId)); + var requestData = fixture.CreateRequest( + context, + scenario, + payload: "must-not-run", + requestContextValue: "pre-start", + recordResponse: true); + request = requestData.Request; + rejected = collector.WaitForEventAsync( + nameof(DispatcherEvents.Rejected), + diagnosticEvent => diagnosticEvent.Payload is DispatcherEvents.Rejected rejection + && ReferenceEquals(rejection.Message, requestData.Message), + Timeout, + TestContext.Current.CancellationToken); + context.ReceiveMessage(requestData.Message); + throw expected; + } + } + + [Fact] + public async Task AsyncActivation_OrdersLifecycleCallbacksAndCleanup() + { + var (grainId, scenario) = fixture.CreateScenario( + ActivationStartupCompletion.AsynchronousSuccess, + ActivationStartupDisposal.Asynchronous); + using var lifecycleSubscription = fixture.ObserveLifecycle(); + var expected = new[] + { + "Created", + "LifecycleStartLow", + "LifecycleStartHigh", + "OnActivateEntered", + "OnActivateCompleted", + "Activated", + "RequestInvoked", + "Deactivating", + "OnDeactivateEntered", + "OnDeactivateCompleted", + "LifecycleStopHigh", + "LifecycleStopLow", + "ActivatorDisposeStarted", + "ActivatorDisposeCompleted", + "ScopeDisposed", + "Deactivated", + }; + var eventTasks = expected.Select(scenario.WaitForEvent).ToArray(); + ActivationData? context = null; + + try + { + context = fixture.StartActivation(grainId); + await scenario.WaitForEvent("OnActivateEntered").WaitAsync(Timeout, TestContext.Current.CancellationToken); + + scenario.ReleaseActivation(); + await scenario.WaitForEvent("Activated").WaitAsync(Timeout, TestContext.Current.CancellationToken); + + var (message, _) = fixture.CreateRequest( + context, + scenario, + payload: "lifecycle", + requestContextValue: "lifecycle-request", + recordResponse: false); + context.ReceiveMessage(message); + await scenario.WaitForEvent("RequestInvoked").WaitAsync(Timeout, TestContext.Current.CancellationToken); + + context.Deactivate( + new(DeactivationReasonCode.ApplicationRequested, "Lifecycle test complete."), + TestContext.Current.CancellationToken); + await scenario.WaitForEvent("ActivatorDisposeStarted").WaitAsync(Timeout, TestContext.Current.CancellationToken); + Assert.Equal(0, scenario.DisposeCompletedCount); + Assert.Equal(0, scenario.ScopeDisposeCount); + + scenario.ReleaseDisposal(); + await scenario.WaitForEvent("Deactivated").WaitAsync(Timeout, TestContext.Current.CancellationToken); + await Task.WhenAll(eventTasks).WaitAsync(Timeout, TestContext.Current.CancellationToken); + + Assert.Equal(expected, scenario.Events.Select(static entry => entry.Name)); + Assert.All(expected, name => Assert.Equal(1, scenario.GetEventCount(name))); + Assert.Equal(1, scenario.ConstructorCount); + Assert.Equal(1, scenario.OnActivateCount); + Assert.Equal(1, scenario.OnDeactivateCount); + Assert.Equal(1, scenario.RequestInvocationCount); + AssertIdentity(scenario, grainId); + Assert.Null(fixture.ActivationDirectory.FindTarget(grainId)); + } + finally + { + await CleanupAsync(context, scenario); + fixture.RemoveScenario(grainId); + } + } + + [Fact] + public async Task RequestAdmittedDuringAsyncActivation_CompletesAfterStartup() + { + var (grainId, scenario) = fixture.CreateScenario( + ActivationStartupCompletion.AsynchronousSuccess, + ActivationStartupDisposal.Synchronous); + using var lifecycleSubscription = fixture.ObserveLifecycle(); + var entered = scenario.WaitForEvent("OnActivateEntered"); + var response = scenario.WaitForEvent("Response"); + ActivationData? context = null; + + try + { + context = fixture.StartActivation(grainId); + await entered.WaitAsync(Timeout, TestContext.Current.CancellationToken); + var (message, request) = fixture.CreateRequest( + context, + scenario, + payload: "success", + requestContextValue: "request-success", + recordResponse: true); + + context.ReceiveMessage(message); + scenario.Record("RequestAdmitted", context); + + Assert.Equal(0, scenario.RequestInvocationCount); + Assert.Null(request.Result); + + scenario.ReleaseActivation(); + await response.WaitAsync(Timeout, TestContext.Current.CancellationToken); + + Assert.Equal("success:request-success", request.Result); + Assert.Equal(1, scenario.RequestInvocationCount); + Assert.Equal("request-success", scenario.RequestContextValue); + AssertEventSubset( + scenario, + "RequestAdmitted", + "OnActivateCompleted", + "Activated", + "RequestInvoked", + "Response"); + AssertIdentity(scenario, grainId); + } + finally + { + await CleanupAsync(context, scenario); + fixture.RemoveScenario(grainId); + } + } + + [Fact] + public async Task RequestAdmittedDuringAsyncActivationFailure_IsRejectedWithoutInvocation() + { + var (grainId, scenario) = fixture.CreateScenario( + ActivationStartupCompletion.AsynchronousFailure, + ActivationStartupDisposal.Synchronous); + using var lifecycleSubscription = fixture.ObserveLifecycle(); + using var collector = new DiagnosticEventCollector(DispatcherEvents.ListenerName); + var entered = scenario.WaitForEvent("OnActivateEntered"); + var failed = scenario.WaitForEvent("OnActivateFailed"); + var deactivated = scenario.WaitForEvent("Deactivated"); + ActivationData? context = null; + + try + { + context = fixture.StartActivation(grainId); + await entered.WaitAsync(Timeout, TestContext.Current.CancellationToken); + var (message, request) = fixture.CreateRequest( + context, + scenario, + payload: "must-not-run", + requestContextValue: "request-failure", + recordResponse: true); + var rejected = collector.WaitForEventAsync( + nameof(DispatcherEvents.Rejected), + diagnosticEvent => diagnosticEvent.Payload is DispatcherEvents.Rejected rejection + && ReferenceEquals(rejection.Message, message), + Timeout, + TestContext.Current.CancellationToken); + + context.ReceiveMessage(message); + scenario.Record("RequestAdmitted", context); + Assert.Equal(0, scenario.RequestInvocationCount); + + scenario.ReleaseActivation(); + await failed.WaitAsync(Timeout, TestContext.Current.CancellationToken); + var rejectionEvent = await rejected; + await deactivated.WaitAsync(Timeout, TestContext.Current.CancellationToken); + + var rejection = Assert.IsType(rejectionEvent.Payload); + var exception = Assert.IsType(rejection.Exception); + Assert.Equal("activate-fault", exception.Message); + Assert.Equal(Message.RejectionTypes.Transient, rejection.RejectionType); + Assert.Contains("Failed to activate grain", rejection.Reason); + Assert.Same(message, rejection.Message); + Assert.Null(request.Result); + Assert.Equal(0, scenario.RequestInvocationCount); + Assert.Equal(0, scenario.GetEventCount("Activated")); + Assert.Equal(1, scenario.GetEventCount("Deactivating")); + Assert.Equal(1, scenario.GetEventCount("Deactivated")); + Assert.Equal(1, scenario.DisposeStartedCount); + Assert.Equal(1, scenario.DisposeCompletedCount); + Assert.Equal(1, scenario.ScopeDisposeCount); + Assert.Null(fixture.ActivationDirectory.FindTarget(grainId)); + AssertIdentity(scenario, grainId); + } + finally + { + await CleanupAsync(context, scenario); + fixture.RemoveScenario(grainId); + } + } + + [Fact] + public async Task CancellationDuringAsyncActivation_AbortsAndCleansUpExactlyOnce() + { + var (grainId, scenario) = fixture.CreateScenario( + ActivationStartupCompletion.Cancellation, + ActivationStartupDisposal.Asynchronous); + using var lifecycleSubscription = fixture.ObserveLifecycle(); + var entered = scenario.WaitForEvent("OnActivateEntered"); + var cancellationObserved = scenario.WaitForEvent("CancellationObserved"); + var disposeStarted = scenario.WaitForEvent("ActivatorDisposeStarted"); + var deactivated = scenario.WaitForEvent("Deactivated"); + ActivationData? context = null; + + try + { + context = fixture.StartActivation(grainId); + await entered.WaitAsync(Timeout, TestContext.Current.CancellationToken); + + context.Deactivate( + new(DeactivationReasonCode.RuntimeRequested, "Cancel startup."), + TestContext.Current.CancellationToken); + await cancellationObserved.WaitAsync(Timeout, TestContext.Current.CancellationToken); + await disposeStarted.WaitAsync(Timeout, TestContext.Current.CancellationToken); + Assert.Equal(0, scenario.DisposeCompletedCount); + Assert.Equal(0, scenario.ScopeDisposeCount); + + scenario.ReleaseDisposal(); + await deactivated.WaitAsync(Timeout, TestContext.Current.CancellationToken); + + Assert.Equal(0, scenario.GetEventCount("Activated")); + Assert.Equal(0, scenario.RequestInvocationCount); + Assert.Equal(0, scenario.OnDeactivateCount); + Assert.Equal(1, scenario.GetEventCount("Deactivating")); + Assert.Equal(1, scenario.GetEventCount("Deactivated")); + Assert.Equal(1, scenario.DisposeStartedCount); + Assert.Equal(1, scenario.DisposeCompletedCount); + Assert.Equal(1, scenario.ScopeDisposeCount); + Assert.Null(fixture.ActivationDirectory.FindTarget(grainId)); + AssertIdentity(scenario, grainId); + } + finally + { + await CleanupAsync(context, scenario); + fixture.RemoveScenario(grainId); + } + } + + [Fact] + public async Task SiloShutdownDuringAsyncActivation_CancelsAndCleansUpExactlyOnce() + { + var shutdownFixture = new ActivationStartupTestFixture(); + await shutdownFixture.InitializeAsync(); + var (grainId, scenario) = shutdownFixture.CreateScenario( + ActivationStartupCompletion.Cancellation, + ActivationStartupDisposal.Asynchronous); + var shutdownHooks = shutdownFixture.Hooks; + var activationDirectory = shutdownFixture.ActivationDirectory; + var primarySilo = shutdownFixture.PrimarySilo; + using var lifecycleSubscription = shutdownFixture.ObserveLifecycle(); + var entered = scenario.WaitForEvent("OnActivateEntered"); + var cancellationObserved = scenario.WaitForEvent("CancellationObserved"); + var disposeStarted = scenario.WaitForEvent("ActivatorDisposeStarted"); + var deactivated = scenario.WaitForEvent("Deactivated"); + using var collector = new DiagnosticEventCollector(DispatcherEvents.ListenerName); + Task? stopTask = null; + ActivationData? context = null; + + try + { + context = shutdownFixture.StartActivation(grainId); + await entered.WaitAsync(Timeout, TestContext.Current.CancellationToken); + var (message, request) = shutdownFixture.CreateRequest( + context, + scenario, + payload: "shutdown-pending", + requestContextValue: "shutdown-request", + recordResponse: true, + invokeMethodOptions: InvokeMethodOptions.None); + var rejected = collector.WaitForEventAsync( + nameof(DispatcherEvents.Rejected), + diagnosticEvent => diagnosticEvent.Payload is DispatcherEvents.Rejected rejection + && ReferenceEquals(rejection.Message, message), + Timeout, + TestContext.Current.CancellationToken); + + context.ReceiveMessage(message); + Assert.Equal(0, scenario.RequestInvocationCount); + Assert.Null(request.Result); + + stopTask = shutdownFixture.HostedCluster.StopSiloAsync( + primarySilo, + cancellationToken: TestContext.Current.CancellationToken); + await cancellationObserved.WaitAsync(Timeout, TestContext.Current.CancellationToken); + await disposeStarted.WaitAsync(Timeout, TestContext.Current.CancellationToken); + + scenario.ReleaseDisposal(); + await deactivated.WaitAsync(Timeout, TestContext.Current.CancellationToken); + await stopTask.WaitAsync(Timeout, TestContext.Current.CancellationToken); + var rejectionEvent = await rejected; + var rejection = Assert.IsType(rejectionEvent.Payload); + + Assert.Same(message, rejection.Message); + Assert.Equal(Message.RejectionTypes.Unrecoverable, rejection.RejectionType); + var unavailable = Assert.IsType(rejection.Exception); + Assert.Equal( + $"Silo '{primarySilo.SiloAddress}' is shutting down.", + unavailable.Message); + Assert.Null(rejection.Reason); + Assert.Null(request.Result); + Assert.Equal(0, scenario.GetEventCount("Activated")); + Assert.Equal(0, scenario.RequestInvocationCount); + Assert.Equal(1, scenario.GetEventCount("Deactivating")); + Assert.Equal(1, scenario.GetEventCount("Deactivated")); + Assert.Equal(1, scenario.DisposeStartedCount); + Assert.Equal(1, scenario.DisposeCompletedCount); + Assert.Equal(1, scenario.ScopeDisposeCount); + Assert.Equal(ActivationState.Invalid, context.State); + Assert.Null(activationDirectory.FindTarget(grainId)); + AssertIdentity(scenario, grainId); + } + finally + { + scenario.ReleaseActivation(); + scenario.ReleaseDisposal(); + if (stopTask is not null) + { + await stopTask.WaitAsync(Timeout, TestContext.Current.CancellationToken); + } + + shutdownHooks.RemoveScenario(grainId); + await shutdownFixture.DisposeAsync(); + } + } + + [Fact] + public async Task ActivationStartup_DoesNotFlowExecutionContextOrRequestContext() + { + var (grainIdA, scenarioA) = fixture.CreateScenario( + ActivationStartupCompletion.AsynchronousSuccess, + ActivationStartupDisposal.Synchronous); + var (grainIdB, scenarioB) = fixture.CreateScenario( + ActivationStartupCompletion.AsynchronousSuccess, + ActivationStartupDisposal.Synchronous); + using var lifecycleSubscription = fixture.ObserveLifecycle(); + var enteredA = scenarioA.WaitForEvent("OnActivateEntered"); + var enteredB = scenarioB.WaitForEvent("OnActivateEntered"); + var responseA = scenarioA.WaitForEvent("Response"); + var responseB = scenarioB.WaitForEvent("Response"); + var originalAmbient = fixture.Hooks.AmbientValue; + var originalRequestContext = RequestContext.Get(ActivationStartupTestHooks.RequestContextKey); + ActivationData? contextA = null; + ActivationData? contextB = null; + + try + { + fixture.Hooks.AmbientValue = "caller-ambient"; + RequestContext.Set(ActivationStartupTestHooks.RequestContextKey, "caller-request"); + + contextA = fixture.StartActivation(grainIdA, "activate-a"); + contextB = fixture.StartActivation(grainIdB, "activate-b"); + await Task.WhenAll(enteredA, enteredB).WaitAsync(Timeout, TestContext.Current.CancellationToken); + + var (messageA, requestA) = fixture.CreateRequest( + contextA, + scenarioA, + payload: "a", + requestContextValue: "request-a", + recordResponse: true); + var (messageB, requestB) = fixture.CreateRequest( + contextB, + scenarioB, + payload: "b", + requestContextValue: "request-b", + recordResponse: true); + contextA.ReceiveMessage(messageA); + scenarioA.Record("RequestAdmitted", contextA); + contextB.ReceiveMessage(messageB); + scenarioB.Record("RequestAdmitted", contextB); + Assert.Equal(0, scenarioA.RequestInvocationCount); + Assert.Equal(0, scenarioB.RequestInvocationCount); + + scenarioA.ReleaseActivation(); + scenarioB.ReleaseActivation(); + await Task.WhenAll(responseA, responseB).WaitAsync(Timeout, TestContext.Current.CancellationToken); + + Assert.Null(scenarioA.ConstructorAmbientValue); + Assert.Null(scenarioB.ConstructorAmbientValue); + Assert.Null(scenarioA.ConstructorRequestContextValue); + Assert.Null(scenarioB.ConstructorRequestContextValue); + Assert.Null(scenarioA.OnActivateAmbientValue); + Assert.Null(scenarioB.OnActivateAmbientValue); + Assert.Equal("activate-a", scenarioA.OnActivateRequestContextValue); + Assert.Equal("activate-b", scenarioB.OnActivateRequestContextValue); + Assert.Null(scenarioA.RequestAmbientValue); + Assert.Null(scenarioB.RequestAmbientValue); + Assert.Equal("request-a", scenarioA.RequestContextValue); + Assert.Equal("request-b", scenarioB.RequestContextValue); + Assert.Equal("a:request-a", requestA.Result); + Assert.Equal("b:request-b", requestB.Result); + Assert.Equal(1, scenarioA.ConstructorCount); + Assert.Equal(1, scenarioB.ConstructorCount); + Assert.Equal(1, scenarioA.OnActivateCount); + Assert.Equal(1, scenarioB.OnActivateCount); + Assert.Equal(1, scenarioA.RequestInvocationCount); + Assert.Equal(1, scenarioB.RequestInvocationCount); + Assert.Equal("caller-ambient", fixture.Hooks.AmbientValue); + Assert.Equal("caller-request", RequestContext.Get(ActivationStartupTestHooks.RequestContextKey)); + } + finally + { + fixture.Hooks.AmbientValue = originalAmbient; + if (originalRequestContext is null) + { + RequestContext.Remove(ActivationStartupTestHooks.RequestContextKey); + } + else + { + RequestContext.Set(ActivationStartupTestHooks.RequestContextKey, originalRequestContext); + } + await CleanupAsync(contextA, scenarioA); + await CleanupAsync(contextB, scenarioB); + fixture.RemoveScenario(grainIdA); + fixture.RemoveScenario(grainIdB); + } + } + + [Theory] + [InlineData(ActivationStartupCompletion.ImmediateSuccess, ActivationStartupDisposal.Synchronous)] + [InlineData(ActivationStartupCompletion.AsynchronousSuccess, ActivationStartupDisposal.Asynchronous)] + [InlineData(ActivationStartupCompletion.ImmediateFailure, ActivationStartupDisposal.Synchronous)] + [InlineData(ActivationStartupCompletion.AsynchronousFailure, ActivationStartupDisposal.Asynchronous)] + [InlineData(ActivationStartupCompletion.Cancellation, ActivationStartupDisposal.Asynchronous)] + public async Task ActivationStartup_CleanupOccursExactlyOnce( + ActivationStartupCompletion completion, + ActivationStartupDisposal disposal) + { + var (grainId, scenario) = fixture.CreateScenario(completion, disposal); + using var lifecycleSubscription = fixture.ObserveLifecycle(); + var activated = scenario.WaitForEvent("Activated"); + var entered = scenario.WaitForEvent("OnActivateEntered"); + var failed = scenario.WaitForEvent("OnActivateFailed"); + var cancellationObserved = scenario.WaitForEvent("CancellationObserved"); + var disposeStarted = scenario.WaitForEvent("ActivatorDisposeStarted"); + var disposeCompleted = scenario.WaitForEvent("ActivatorDisposeCompleted"); + var scopeDisposed = scenario.WaitForEvent("ScopeDisposed"); + var deactivated = scenario.WaitForEvent("Deactivated"); + ActivationData? context = null; + + try + { + context = fixture.StartActivation(grainId); + switch (completion) + { + case ActivationStartupCompletion.ImmediateSuccess: + await activated.WaitAsync(Timeout, TestContext.Current.CancellationToken); + context.Deactivate( + new(DeactivationReasonCode.ApplicationRequested, "Matrix success complete."), + TestContext.Current.CancellationToken); + break; + case ActivationStartupCompletion.AsynchronousSuccess: + await entered.WaitAsync(Timeout, TestContext.Current.CancellationToken); + scenario.ReleaseActivation(); + await activated.WaitAsync(Timeout, TestContext.Current.CancellationToken); + context.Deactivate( + new(DeactivationReasonCode.ApplicationRequested, "Matrix success complete."), + TestContext.Current.CancellationToken); + break; + case ActivationStartupCompletion.ImmediateFailure: + await failed.WaitAsync(Timeout, TestContext.Current.CancellationToken); + break; + case ActivationStartupCompletion.AsynchronousFailure: + await entered.WaitAsync(Timeout, TestContext.Current.CancellationToken); + scenario.ReleaseActivation(); + await failed.WaitAsync(Timeout, TestContext.Current.CancellationToken); + break; + case ActivationStartupCompletion.Cancellation: + await entered.WaitAsync(Timeout, TestContext.Current.CancellationToken); + context.Deactivate( + new(DeactivationReasonCode.RuntimeRequested, "Matrix cancellation."), + TestContext.Current.CancellationToken); + await cancellationObserved.WaitAsync(Timeout, TestContext.Current.CancellationToken); + break; + default: + throw new ArgumentOutOfRangeException(nameof(completion)); + } + + await disposeStarted.WaitAsync(Timeout, TestContext.Current.CancellationToken); + if (disposal is ActivationStartupDisposal.Asynchronous) + { + Assert.Equal(0, scenario.DisposeCompletedCount); + Assert.Equal(0, scenario.ScopeDisposeCount); + scenario.ReleaseDisposal(); + } + + await deactivated.WaitAsync(Timeout, TestContext.Current.CancellationToken); + await Task.WhenAll(disposeCompleted, scopeDisposed).WaitAsync(Timeout, TestContext.Current.CancellationToken); + + Assert.Equal(1, scenario.CreateCount); + Assert.Equal(1, scenario.ConstructorCount); + Assert.Equal(1, scenario.OnActivateCount); + Assert.Equal(1, scenario.DisposeStartedCount); + Assert.Equal(1, scenario.DisposeCompletedCount); + Assert.Equal(1, scenario.ScopeDisposeCount); + Assert.Equal( + completion is ActivationStartupCompletion.ImmediateSuccess + or ActivationStartupCompletion.AsynchronousSuccess ? 1 : 0, + scenario.OnDeactivateCount); + Assert.Equal( + completion is ActivationStartupCompletion.ImmediateSuccess + or ActivationStartupCompletion.AsynchronousSuccess ? 1 : 0, + scenario.GetEventCount("Activated")); + Assert.Equal(1, scenario.GetEventCount("Deactivated")); + Assert.Null(fixture.ActivationDirectory.FindTarget(grainId)); + + var events = scenario.Events; + Assert.True(IndexOf(events, "ActivatorDisposeStarted") < IndexOf(events, "ActivatorDisposeCompleted")); + Assert.True(IndexOf(events, "ActivatorDisposeCompleted") < IndexOf(events, "ScopeDisposed")); + Assert.True(IndexOf(events, "ScopeDisposed") < IndexOf(events, "Deactivated")); + Assert.Equal("Deactivated", events[^1].Name); + AssertIdentity(scenario, grainId); + + var eventSnapshot = scenario.Events; + await context.DisposeAsync(); + await context.DisposeAsync(); + Assert.Equal(eventSnapshot, scenario.Events); + Assert.Equal(1, scenario.DisposeCompletedCount); + Assert.Equal(1, scenario.ScopeDisposeCount); + } + finally + { + await CleanupAsync(context, scenario); + fixture.RemoveScenario(grainId); + } + } + + private static async Task CleanupAsync(ActivationData? context, ActivationStartupScenario scenario) + { + scenario.ReleaseActivation(); + scenario.ReleaseDisposal(); + if (context is null) + { + return; + } + + var deactivated = scenario.WaitForEvent("Deactivated"); + context.Deactivate(new(DeactivationReasonCode.ApplicationRequested, "Test cleanup.")); + await deactivated.WaitAsync(Timeout, TestContext.Current.CancellationToken); + } + + private static void AssertEventSubset(ActivationStartupScenario scenario, params string[] expected) + { + var expectedSet = expected.ToHashSet(); + Assert.Equal(expected, scenario.Events.Where(entry => expectedSet.Contains(entry.Name)).Select(static entry => entry.Name)); + Assert.All(expected, name => Assert.Equal(1, scenario.GetEventCount(name))); + } + + private static void AssertIdentity(ActivationStartupScenario scenario, GrainId grainId) + { + var events = scenario.Events; + Assert.NotEmpty(events); + Assert.All(events, entry => Assert.Equal(grainId, entry.GrainId)); + Assert.Single(events.Select(static entry => entry.ActivationId).Distinct()); + } + + private static int IndexOf(IReadOnlyList events, string name) + { + for (var i = 0; i < events.Count; i++) + { + if (events[i].Name == name) + { + return i; + } + } + + return -1; + } +} diff --git a/test/Orleans.Runtime.Internal.Tests/ActivationsLifeCycleTests/StatelessWorkerActivationStartupTests.cs b/test/Orleans.Runtime.Internal.Tests/ActivationsLifeCycleTests/StatelessWorkerActivationStartupTests.cs new file mode 100644 index 00000000000..9c2b5d93bb4 --- /dev/null +++ b/test/Orleans.Runtime.Internal.Tests/ActivationsLifeCycleTests/StatelessWorkerActivationStartupTests.cs @@ -0,0 +1,725 @@ +using System.Reflection; +using Microsoft.Extensions.DependencyInjection; +using Orleans.CodeGeneration; +using Orleans.Metadata; +using Orleans.Runtime; +using Orleans.Runtime.Diagnostics; +using Orleans.Serialization.Invocation; +using Orleans.TestingHost; +using Orleans.TestingHost.Diagnostics; +using TestExtensions; +using UnitTests.Grains; +using Xunit; + +namespace UnitTests.ActivationsLifeCycleTests; + +[TestSuite("BVT")] +[TestProvider("None")] +[TestArea("Runtime")] +public sealed class StatelessWorkerActivationStartupTests( + StatelessWorkerActivationStartupTests.Fixture fixture) + : IClassFixture +{ + private static readonly TimeSpan Timeout = TimeSpan.FromSeconds(30); + private static readonly FieldInfo WorkersField = + typeof(StatelessWorkerGrainContext).GetField( + "_workers", + BindingFlags.Instance | BindingFlags.NonPublic) + ?? throw new InvalidOperationException("The stateless-worker child list was not found."); + + [Fact] + public async Task StatelessWorkerChildIsPublishedBeforeSynchronousStartAndDoesNotInvokeEarly() + { + var (grainId, scenario) = fixture.CreateScenario( + ActivationStartupCompletion.ImmediateSuccess, + ActivationStartupDisposal.Synchronous); + var gate = fixture.CreateConstructorGate(grainId, CatalogActivationStartupOutcome.Success); + using var eventSubscription = fixture.ObserveEvents(); + using var collector = new DiagnosticEventCollector(StatelessWorkerEvents.ListenerName); + var wrapper = fixture.GetOrCreateContext(grainId); + var workerCreatedTask = WaitForWorkerCreatedAsync(collector, wrapper); + var responseTask = scenario.WaitForEvent("Response"); + ActivationData? worker = null; + Task? admissionTask = null; + + try + { + var (message, request) = fixture.CreateRequest( + wrapper, + scenario, + payload: "synchronous", + requestContextValue: "sync-request"); + + admissionTask = Task.Run( + () => wrapper.ReceiveMessage(message), + TestContext.Current.CancellationToken); + worker = Assert.IsType( + await gate.Entered.WaitAsync(Timeout, TestContext.Current.CancellationToken)); + + Assert.Same(worker, Assert.Single(GetWorkers(wrapper))); + Assert.Equal(1, gate.EntryCount); + Assert.Equal(1, scenario.CreateCount); + Assert.Equal(1, scenario.ConstructorCount); + Assert.Equal(0, scenario.OnActivateCount); + Assert.Equal(0, scenario.RequestInvocationCount); + Assert.Null(request.Result); + Assert.Empty(GetWorkerCreatedEvents(collector, wrapper)); + + gate.Release(); + await admissionTask.WaitAsync(Timeout, TestContext.Current.CancellationToken); + var workerCreated = await workerCreatedTask; + await responseTask.WaitAsync(Timeout, TestContext.Current.CancellationToken); + + Assert.Same(wrapper, workerCreated.Context); + Assert.Same(worker, workerCreated.WorkerContext); + Assert.Equal(1, workerCreated.WorkerCount); + Assert.Equal( + $"synchronous:{worker.ActivationId}:sync-request", + request.Result); + Assert.Equal(1, scenario.CreateCount); + Assert.Equal(1, scenario.ConstructorCount); + Assert.Equal(1, scenario.OnActivateCount); + Assert.Equal(1, scenario.RequestInvocationCount); + Assert.Same(worker, Assert.Single(GetWorkers(wrapper))); + AssertScenarioEventOrder( + scenario, + "ConstructorEntered", + "ConstructorBlocked", + "WorkerCreated", + "LifecycleStartLow", + "LifecycleStartHigh", + "OnActivateEntered", + "OnActivateCompleted", + "RequestInvoked", + "Response"); + AssertIdentity(scenario, grainId, worker.ActivationId); + Assert.Empty(GetContextTerminatedEvents(collector, wrapper)); + Assert.Empty(GetMessageForwardedEvents(collector, grainId)); + } + finally + { + gate.Release(); + if (admissionTask is not null) + { + await admissionTask.WaitAsync(Timeout, TestContext.Current.CancellationToken); + } + + await CleanupAsync(wrapper, worker, scenario, collector); + fixture.RemoveScenario(grainId); + } + } + + [Fact] + public async Task StatelessWorkerRequestWaitsForAsyncStartupBeforeInvocation() + { + var (grainId, scenario) = fixture.CreateScenario( + ActivationStartupCompletion.AsynchronousSuccess, + ActivationStartupDisposal.Synchronous); + using var eventSubscription = fixture.ObserveEvents(); + using var collector = new DiagnosticEventCollector(StatelessWorkerEvents.ListenerName); + var wrapper = fixture.GetOrCreateContext(grainId); + var workerCreatedTask = WaitForWorkerCreatedAsync(collector, wrapper); + var activationEnteredTask = scenario.WaitForEvent("OnActivateEntered"); + var responseTask = scenario.WaitForEvent("Response"); + ActivationData? worker = null; + + try + { + var (message, request) = fixture.CreateRequest( + wrapper, + scenario, + payload: "asynchronous", + requestContextValue: "async-request"); + + wrapper.ReceiveMessage(message); + var workerCreated = await workerCreatedTask; + worker = Assert.IsType(workerCreated.WorkerContext); + await activationEnteredTask.WaitAsync(Timeout, TestContext.Current.CancellationToken); + + Assert.Same(worker, Assert.Single(GetWorkers(wrapper))); + Assert.Equal(1, workerCreated.WorkerCount); + Assert.Equal(1, scenario.ConstructorCount); + Assert.Equal(1, scenario.OnActivateCount); + Assert.Equal(0, scenario.RequestInvocationCount); + Assert.Null(request.Result); + + scenario.ReleaseActivation(); + await responseTask.WaitAsync(Timeout, TestContext.Current.CancellationToken); + + Assert.Equal( + $"asynchronous:{worker.ActivationId}:async-request", + request.Result); + Assert.Equal(1, scenario.RequestInvocationCount); + Assert.Equal("async-request", scenario.RequestContextValue); + AssertScenarioEventOrder( + scenario, + "WorkerCreated", + "OnActivateEntered", + "OnActivateCompleted", + "RequestInvoked", + "Response"); + AssertIdentity(scenario, grainId, worker.ActivationId); + Assert.Same(worker, Assert.Single(GetWorkers(wrapper))); + Assert.Single(GetWorkerCreatedEvents(collector, wrapper)); + Assert.Empty(GetContextTerminatedEvents(collector, wrapper)); + Assert.Empty(GetMessageForwardedEvents(collector, grainId)); + } + finally + { + await CleanupAsync(wrapper, worker, scenario, collector); + fixture.RemoveScenario(grainId); + } + } + + [Theory] + [InlineData(StatelessWorkerStartupFailure.Constructor)] + [InlineData(StatelessWorkerStartupFailure.AsynchronousActivation)] + public async Task StatelessWorkerStartupFailureRemovesChildAndDisposesResourcesExactlyOnce( + StatelessWorkerStartupFailure failure) + { + var completion = failure is StatelessWorkerStartupFailure.Constructor + ? ActivationStartupCompletion.ImmediateSuccess + : ActivationStartupCompletion.AsynchronousFailure; + var disposal = failure is StatelessWorkerStartupFailure.Constructor + ? ActivationStartupDisposal.Synchronous + : ActivationStartupDisposal.Asynchronous; + var (grainId, scenario) = fixture.CreateScenario(completion, disposal); + var gate = failure is StatelessWorkerStartupFailure.Constructor + ? fixture.CreateConstructorGate(grainId, CatalogActivationStartupOutcome.ConstructorFailure) + : null; + using var eventSubscription = fixture.ObserveEvents(); + using var statelessEvents = new DiagnosticEventCollector(StatelessWorkerEvents.ListenerName); + using var dispatcherEvents = new DiagnosticEventCollector(DispatcherEvents.ListenerName); + var wrapper = fixture.GetOrCreateContext(grainId); + var workerCreatedTask = WaitForWorkerCreatedAsync(statelessEvents, wrapper); + var contextTerminatedTask = WaitForContextTerminatedAsync(statelessEvents, wrapper); + var deactivatedTask = scenario.WaitForEvent("Deactivated"); + var scopeDisposedTask = scenario.WaitForEvent("ScopeDisposed"); + ActivationData? worker = null; + Task? admissionTask = null; + + try + { + var (message, request) = fixture.CreateRequest( + wrapper, + scenario, + payload: "must-not-run", + requestContextValue: $"failure-{failure}"); + var rejectedTask = WaitForRejectionAsync(dispatcherEvents, message); + + if (gate is not null) + { + admissionTask = Task.Run( + () => wrapper.ReceiveMessage(message), + TestContext.Current.CancellationToken); + worker = Assert.IsType( + await gate.Entered.WaitAsync(Timeout, TestContext.Current.CancellationToken)); + Assert.Same(worker, Assert.Single(GetWorkers(wrapper))); + Assert.Equal(0, scenario.RequestInvocationCount); + gate.Release(); + await admissionTask.WaitAsync(Timeout, TestContext.Current.CancellationToken); + } + else + { + wrapper.ReceiveMessage(message); + await scenario.WaitForEvent("OnActivateEntered") + .WaitAsync(Timeout, TestContext.Current.CancellationToken); + worker = Assert.IsType((await workerCreatedTask).WorkerContext); + Assert.Same(worker, Assert.Single(GetWorkers(wrapper))); + Assert.Equal(0, scenario.RequestInvocationCount); + scenario.ReleaseActivation(); + await scenario.WaitForEvent("OnActivateFailed") + .WaitAsync(Timeout, TestContext.Current.CancellationToken); + await scenario.WaitForEvent("ActivatorDisposeStarted") + .WaitAsync(Timeout, TestContext.Current.CancellationToken); + Assert.Equal(0, scenario.DisposeCompletedCount); + Assert.Equal(0, scenario.ScopeDisposeCount); + scenario.ReleaseDisposal(); + } + + var workerCreated = await workerCreatedTask; + worker ??= Assert.IsType(workerCreated.WorkerContext); + var rejectionEvent = await rejectedTask; + var terminated = await contextTerminatedTask; + await Task.WhenAll(deactivatedTask, scopeDisposedTask) + .WaitAsync(Timeout, TestContext.Current.CancellationToken); + + var rejection = Assert.IsType(rejectionEvent.Payload); + var expectedException = gate?.ConstructorException ?? scenario.ActivationException; + var expectedExceptionMessage = failure is StatelessWorkerStartupFailure.Constructor + ? "constructor-fault" + : "activate-fault"; + Assert.Same(expectedException, rejection.Exception); + Assert.Equal(expectedExceptionMessage, rejection.Exception!.Message); + Assert.Equal(Message.RejectionTypes.Transient, rejection.RejectionType); + Assert.Equal( + failure is StatelessWorkerStartupFailure.Constructor + ? "Error constructing grain instance." + : "Failed to activate grain.", + rejection.Reason); + Assert.Same(message, rejection.Message); + Assert.Null(request.Result); + Assert.Equal(0, scenario.RequestInvocationCount); + Assert.Equal(0, scenario.GetEventCount("Activated")); + Assert.Equal(1, scenario.GetEventCount("Deactivated")); + Assert.Equal( + failure is StatelessWorkerStartupFailure.Constructor ? 0 : 1, + scenario.DisposeStartedCount); + Assert.Equal( + failure is StatelessWorkerStartupFailure.Constructor ? 0 : 1, + scenario.DisposeCompletedCount); + Assert.Equal(1, scenario.ScopeDisposeCount); + Assert.Equal(0, terminated.WorkerCount); + Assert.Empty(GetWorkers(wrapper)); + Assert.Single(GetWorkerCreatedEvents(statelessEvents, wrapper)); + Assert.Single(GetContextTerminatedEvents(statelessEvents, wrapper)); + Assert.Empty(GetMessageForwardedEvents(statelessEvents, grainId)); + if (failure is StatelessWorkerStartupFailure.Constructor) + { + AssertScenarioEventOrder( + scenario, + "ConstructorEntered", + "ConstructorBlocked", + "ConstructorFailed", + "Deactivating", + "WorkerCreated", + "ScopeDisposed", + "Deactivated"); + } + else + { + AssertScenarioEventOrder( + scenario, + "ConstructorEntered", + "WorkerCreated", + "OnActivateEntered", + "OnActivateFailed", + "Deactivating", + "ActivatorDisposeStarted", + "ActivatorDisposeCompleted", + "ScopeDisposed", + "Deactivated"); + } + + AssertIdentity(scenario, grainId, worker.ActivationId); + + var eventSnapshot = scenario.Events.ToArray(); + await worker.DisposeAsync(); + await worker.DisposeAsync(); + Assert.Equal(eventSnapshot, scenario.Events); + Assert.Equal(1, scenario.ScopeDisposeCount); + Assert.Equal( + failure is StatelessWorkerStartupFailure.Constructor ? 0 : 1, + scenario.DisposeCompletedCount); + Assert.Single(GetContextTerminatedEvents(statelessEvents, wrapper)); + } + finally + { + gate?.Release(); + if (admissionTask is not null) + { + await admissionTask.WaitAsync(Timeout, TestContext.Current.CancellationToken); + } + + await CleanupAsync(wrapper, worker, scenario, statelessEvents); + fixture.RemoveScenario(grainId); + } + } + + private static IReadOnlyList GetWorkers(StatelessWorkerGrainContext context) => + [.. Assert.IsType>(WorkersField.GetValue(context))]; + + private static IReadOnlyList GetWorkerCreatedEvents( + DiagnosticEventCollector collector, + IGrainContext context) => + [.. collector.Events + .Select(static evt => evt.Payload) + .OfType() + .Where(evt => ReferenceEquals(evt.Context, context))]; + + private static IReadOnlyList GetContextTerminatedEvents( + DiagnosticEventCollector collector, + IGrainContext context) => + [.. collector.Events + .Select(static evt => evt.Payload) + .OfType() + .Where(evt => ReferenceEquals(evt.Context, context))]; + + private static IReadOnlyList GetMessageForwardedEvents( + DiagnosticEventCollector collector, + GrainId grainId) => + [.. collector.Events + .Select(static evt => evt.Payload) + .OfType() + .Where(evt => evt.GrainId.Equals(grainId))]; + + private static async Task WaitForWorkerCreatedAsync( + DiagnosticEventCollector collector, + IGrainContext context) + { + var evt = await collector.WaitForEventAsync( + nameof(StatelessWorkerEvents.WorkerCreated), + diagnosticEvent => diagnosticEvent.Payload is StatelessWorkerEvents.WorkerCreated created + && ReferenceEquals(created.Context, context), + Timeout, + TestContext.Current.CancellationToken); + return Assert.IsType(evt.Payload); + } + + private static async Task WaitForContextTerminatedAsync( + DiagnosticEventCollector collector, + IGrainContext context) + { + var evt = await collector.WaitForEventAsync( + nameof(StatelessWorkerEvents.ContextTerminated), + diagnosticEvent => diagnosticEvent.Payload is StatelessWorkerEvents.ContextTerminated terminated + && ReferenceEquals(terminated.Context, context), + Timeout, + TestContext.Current.CancellationToken); + return Assert.IsType(evt.Payload); + } + + private static Task WaitForRejectionAsync( + DiagnosticEventCollector collector, + Message message) => + collector.WaitForEventAsync( + nameof(DispatcherEvents.Rejected), + diagnosticEvent => diagnosticEvent.Payload is DispatcherEvents.Rejected rejected + && ReferenceEquals(rejected.Message, message), + Timeout, + TestContext.Current.CancellationToken); + + private static async Task CleanupAsync( + StatelessWorkerGrainContext wrapper, + ActivationData? worker, + ActivationStartupScenario scenario, + DiagnosticEventCollector collector) + { + scenario.ReleaseActivation(); + scenario.ReleaseDisposal(); + if (worker is not null && scenario.GetEventCount("Deactivated") == 0) + { + var deactivated = scenario.WaitForEvent("Deactivated"); + worker.Deactivate(new(DeactivationReasonCode.ApplicationRequested, "Test cleanup.")); + await deactivated.WaitAsync(Timeout, TestContext.Current.CancellationToken); + } + + if (worker is not null) + { + await WaitForContextTerminatedAsync(collector, wrapper); + } + + await wrapper.DisposeAsync(); + } + + private static void AssertScenarioEventOrder( + ActivationStartupScenario scenario, + params string[] expected) + { + var expectedSet = expected.ToHashSet(); + Assert.Equal( + expected, + scenario.Events + .Where(entry => expectedSet.Contains(entry.Name)) + .Select(static entry => entry.Name)); + Assert.All(expected, name => Assert.Equal(1, scenario.GetEventCount(name))); + } + + private static void AssertIdentity( + ActivationStartupScenario scenario, + GrainId grainId, + ActivationId activationId) + { + Assert.NotEmpty(scenario.Events); + Assert.All( + scenario.Events, + entry => + { + Assert.Equal(grainId, entry.GrainId); + Assert.Equal(activationId, entry.ActivationId); + }); + } + + public enum StatelessWorkerStartupFailure + { + Constructor, + AsynchronousActivation, + } + + public sealed class Fixture : BaseTestClusterFixture + { + protected override void ConfigureTestCluster(TestClusterBuilder builder) + { + builder.Options.InitialSilosCount = 1; + builder.AddSiloBuilderConfigurator(); + } + + internal IServiceProvider Services => + ((InProcessSiloHandle)HostedCluster.Primary!).SiloHost.Services; + + internal ActivationDirectory ActivationDirectory => + Services.GetRequiredService(); + + private ActivationStartupTestHooks Hooks => + Services.GetRequiredService(); + + private CatalogActivationStartupTestHooks CatalogHooks => + Services.GetRequiredService(); + + internal (GrainId GrainId, ActivationStartupScenario Scenario) CreateScenario( + ActivationStartupCompletion completion, + ActivationStartupDisposal disposal) + { + var grainType = Services.GetRequiredService() + .GetGrainType(typeof(StatelessWorkerActivationStartupTestGrain)); + var grainId = GrainId.Create(grainType, Guid.NewGuid().ToString("N")); + return (grainId, Hooks.CreateScenario(grainId, completion, disposal)); + } + + internal ActivationStartupScenario CreateScenario( + GrainId grainId, + ActivationStartupCompletion completion, + ActivationStartupDisposal disposal) => + Hooks.CreateScenario(grainId, completion, disposal); + + internal CatalogActivationStartupGate CreateConstructorGate( + GrainId grainId, + CatalogActivationStartupOutcome outcome) => + CatalogHooks.CreateGate(grainId, outcome); + + internal StatelessWorkerGrainContext GetOrCreateContext(GrainId grainId) => + Assert.IsType( + Services.GetRequiredService() + .GetOrCreateActivation( + grainId, + requestContextData: null, + rehydrationContext: null)); + + internal (Message Message, StatelessWorkerActivationStartupRequest Request) CreateRequest( + StatelessWorkerGrainContext context, + ActivationStartupScenario scenario, + string payload, + string requestContextValue) + { + var request = new StatelessWorkerActivationStartupRequest(scenario, payload); + var message = Services.GetRequiredService() + .CreateMessage(request, InvokeMethodOptions.OneWay); + message.SetInfiniteTimeToLive(); + message.RequestContextData = new() + { + [ActivationStartupTestHooks.RequestContextKey] = requestContextValue, + }; + message.SendingGrain = GrainId.Create( + "stateless-worker-activation-startup-sender", + Guid.NewGuid().ToString("N")); + message.SendingSilo = ((InProcessSiloHandle)HostedCluster.Primary!).SiloAddress; + message.TargetGrain = context.GrainId; + message.TargetSilo = ((InProcessSiloHandle)HostedCluster.Primary!).SiloAddress; + return (message, request); + } + + internal IDisposable ObserveEvents() => + new CompositeDisposable( + GrainLifecycleEvents.AllEvents.Subscribe(new LifecycleObserver(Hooks)), + StatelessWorkerEvents.AllEvents.Subscribe(new StatelessWorkerObserver(Hooks))); + + internal void RemoveScenario(GrainId grainId) + { + CatalogHooks.RemoveGate(grainId); + Hooks.RemoveScenario(grainId); + } + + private sealed class SiloConfigurator : ISiloConfigurator + { + public void Configure(ISiloBuilder hostBuilder) + { + hostBuilder.Services.AddSingleton(); + hostBuilder.Services.AddSingleton(); + hostBuilder.Services.AddScoped(); + hostBuilder.Services.AddSingleton(); + hostBuilder.Services.AddSingleton< + IConfigureGrainTypeComponents, + StatelessWorkerActivationStartupConfigurator>(); + } + } + + private sealed class StatelessWorkerActivationStartupConfigurator( + GrainClassMap grainClassMap, + StatelessWorkerActivationStartupTestActivator activator) + : IConfigureGrainTypeComponents + { + public void Configure( + GrainType grainType, + GrainProperties properties, + GrainTypeSharedContext shared) + { + if (grainClassMap.TryGetGrainClass(grainType, out var grainClass) + && grainClass == typeof(StatelessWorkerActivationStartupTestGrain)) + { + shared.SetComponent(activator); + } + } + } + + private sealed class StatelessWorkerActivationStartupTestActivator( + ActivationStartupTestHooks hooks, + CatalogActivationStartupTestHooks catalogHooks) : IGrainActivator + { + public object CreateInstance(IGrainContext context) + { + var scenario = hooks.GetRequiredScenario(context.GrainId); + scenario.ObserveCreate(context); + context.ActivationServices + .GetRequiredService() + .Attach(context); + var instance = new StatelessWorkerActivationStartupTestGrain(context, hooks); + + if (catalogHooks.TryGetGate(context.GrainId, out var gate)) + { + gate!.EnterAndWait(context, scenario); + if (gate.Outcome is CatalogActivationStartupOutcome.ConstructorFailure) + { + scenario.Record("ConstructorFailed", context); + throw gate.ConstructorException; + } + } + + return instance; + } + + public async ValueTask DisposeInstance(IGrainContext context, object instance) + { + var scenario = hooks.GetRequiredScenario(context.GrainId); + scenario.ObserveDisposeStarted(context); + if (scenario.Disposal is ActivationStartupDisposal.Asynchronous) + { + await scenario.DisposalRelease; + } + + scenario.ObserveDisposeCompleted(context); + } + } + + private sealed class LifecycleObserver(ActivationStartupTestHooks hooks) + : IObserver + { + public void OnCompleted() + { + } + + public void OnError(Exception error) + { + } + + public void OnNext(GrainLifecycleEvents.LifecycleEvent value) + { + if (!hooks.TryGetScenario(value.GrainContext.GrainId, out var scenario)) + { + return; + } + + var name = value switch + { + GrainLifecycleEvents.Created => "Created", + GrainLifecycleEvents.Activated => "Activated", + GrainLifecycleEvents.Deactivating => "Deactivating", + GrainLifecycleEvents.Deactivated => "Deactivated", + _ => null, + }; + if (name is not null) + { + scenario!.Record(name, value.GrainContext); + } + } + } + + private sealed class StatelessWorkerObserver(ActivationStartupTestHooks hooks) + : IObserver + { + public void OnCompleted() + { + } + + public void OnError(Exception error) + { + } + + public void OnNext(StatelessWorkerEvents.StatelessWorkerEvent value) + { + if (value is StatelessWorkerEvents.WorkerCreated created + && hooks.TryGetScenario(created.GrainId, out var scenario)) + { + scenario!.Record("WorkerCreated", created.WorkerContext); + } + } + } + + private sealed class CompositeDisposable(params IDisposable[] disposables) : IDisposable + { + public void Dispose() + { + foreach (var disposable in disposables) + { + disposable.Dispose(); + } + } + } + } + + internal sealed class StatelessWorkerActivationStartupRequest( + ActivationStartupScenario scenario, + string payload) : IInvokable + { + private static readonly MethodInfo Method = + typeof(IStatelessWorkerActivationStartupTestGrain) + .GetMethod(nameof(IStatelessWorkerActivationStartupTestGrain.Invoke))!; + + private IStatelessWorkerActivationStartupTestGrain? _target; + private readonly TaskCompletionSource _completion = + new(TaskCreationOptions.RunContinuationsAsynchronously); + private string? _result; + + public Task Completion => _completion.Task; + + public string? Result => Volatile.Read(ref _result); + + public object? GetTarget() => _target; + + public void SetTarget(ITargetHolder holder) + { + _target = (IStatelessWorkerActivationStartupTestGrain)holder.GetTarget()!; + } + + public async ValueTask Invoke() + { + var result = await _target!.Invoke(payload); + Volatile.Write(ref _result, result); + scenario.Record("Response", ((IGrainBase)_target).GrainContext); + _completion.TrySetResult(); + return Response.FromResult(result); + } + + public int GetArgumentCount() => 1; + + public object? GetArgument(int index) => + index == 0 ? payload : throw new ArgumentOutOfRangeException(nameof(index)); + + public void SetArgument(int index, object value) => + throw new NotSupportedException("The stateless-worker startup request is immutable."); + + public string GetMethodName() => nameof(IStatelessWorkerActivationStartupTestGrain.Invoke); + + public string GetInterfaceName() => + typeof(IStatelessWorkerActivationStartupTestGrain).FullName!; + + public string GetActivityName() => $"{GetInterfaceName()}/{GetMethodName()}"; + + public MethodInfo GetMethod() => Method; + + public Type GetInterfaceType() => typeof(IStatelessWorkerActivationStartupTestGrain); + + public void Dispose() + { + } + } +} diff --git a/test/Orleans.Runtime.Tests/GrainActivatorTests.cs b/test/Orleans.Runtime.Tests/GrainActivatorTests.cs index 89edd22ec37..3f11f9dda36 100644 --- a/test/Orleans.Runtime.Tests/GrainActivatorTests.cs +++ b/test/Orleans.Runtime.Tests/GrainActivatorTests.cs @@ -1,8 +1,11 @@ using System.Diagnostics.CodeAnalysis; +using System.Diagnostics.Metrics; using Microsoft.Extensions.DependencyInjection; +using Microsoft.Extensions.Diagnostics.Metrics.Testing; using Microsoft.Extensions.Logging.Abstractions; using Orleans.Metadata; using Orleans.Runtime; +using Orleans.Serialization.Session; using Orleans.TestingHost; using TestExtensions; using UnitTests.GrainInterfaces; @@ -50,6 +53,7 @@ public void Configure(ISiloBuilder hostBuilder) // This allows it to selectively apply to specific grain types services.AddSingleton(); services.AddSingleton(); + services.AddSingleton(); }); } } @@ -112,6 +116,147 @@ public async Task GrainContextIsConfiguredBeforeGrainConstruction() Assert.True(state.WasConfiguredAtConstruction); } + [Fact] + public async Task ContextCreationStartsActivationSynchronouslyOnActivationSchedulerWithCleanExecutionContext() + { + var state = ActivationOrderingState.Instance; + var ambientState = new object(); + state.Arm(); + state.SetAmbientState(ambientState); + + var primary = Assert.IsType(fixture.HostedCluster.Primary); + var services = primary.ServiceProvider; + var grainType = services.GetRequiredService().GetGrainType(typeof(ExplicitlyRegisteredSimpleDIGrain)); + var grainId = GrainId.Create(grainType, Guid.NewGuid().ToString()); + var address = GrainAddress.NewActivationAddress(primary.SiloAddress, grainId); + ActivationData? context = null; + + try + { + context = Assert.IsType(services.GetRequiredService().CreateInstance(address)); + + Assert.NotNull(context.GrainInstance); + Assert.True(state.WasConfiguredAtConstruction); + Assert.True(state.HadSchedulerAffinityAtConstruction); + Assert.True(state.HadRuntimeContextAtConstruction); + Assert.Null(state.RequestContextAtConstruction); + Assert.Null(state.TransactionStateAtConstruction); + Assert.Same(ambientState, RequestContext.Get(ActivationOrderingState.RequestContextKey)); + Assert.Same(ambientState, state.TransactionState); + } + finally + { + state.ClearAmbientState(); + if (context is not null) + { + context.Deactivate( + new(DeactivationReasonCode.ApplicationRequested, "Test completed."), + TestContext.Current.CancellationToken); + await context.Deactivated.WaitAsync( + TimeSpan.FromSeconds(10), + TestContext.Current.CancellationToken); + } + } + } + + [Fact] + public async Task BuiltInActivatorCreateContextStartsContext() + { + var primary = Assert.IsType(fixture.HostedCluster.Primary); + var services = primary.ServiceProvider; + var grainType = services.GetRequiredService().GetGrainType(typeof(ExplicitlyRegisteredSimpleDIGrain)); + var address = GrainAddress.NewActivationAddress( + primary.SiloAddress, + GrainId.Create(grainType, Guid.NewGuid().ToString())); + IGrainContextActivator? activator = null; + foreach (var provider in services.GetServices()) + { + if (provider.TryGet(grainType, out activator)) + { + break; + } + } + + Assert.NotNull(activator); + var context = Assert.IsType(activator.CreateContext(address, [])); + try + { + Assert.NotNull(context.GrainInstance); + } + finally + { + context.Deactivate( + new(DeactivationReasonCode.ApplicationRequested, "Test completed."), + TestContext.Current.CancellationToken); + await context.Deactivated.WaitAsync( + TimeSpan.FromSeconds(10), + TestContext.Current.CancellationToken); + } + } + + [Fact] + public void ConfigurationFailureBeforeConstructionDoesNotDecrementActiveGrainCount() + { + var primary = Assert.IsType(fixture.HostedCluster.Primary); + var services = primary.ServiceProvider; + var grainType = services.GetRequiredService() + .GetGrainType(typeof(ExplicitlyRegisteredSimpleDIGrain)); + var grainId = GrainId.Create(grainType, Guid.NewGuid().ToString()); + var address = GrainAddress.NewActivationAddress(primary.SiloAddress, grainId); + var grainTypeName = services.GetRequiredService() + .GetComponents(grainType) + .GrainTypeName; + using var collector = new MetricCollector( + services.GetRequiredService(), + "Microsoft.Orleans", + InstrumentNames.GRAIN_COUNTS); + + ConfigurationFailureState.Arm(grainId); + try + { + var exception = Assert.Throws( + () => services.GetRequiredService().CreateInstance(address)); + Assert.Equal("configuration-fault", exception.Message); + } + finally + { + ConfigurationFailureState.Clear(); + } + + Assert.DoesNotContain( + collector.GetMeasurementSnapshot(), + measurement => Equals(measurement.Tags["type"], grainTypeName)); + } + + [Fact] + public void ConfigurationFailureDisposesRehydrationContext() + { + var primary = Assert.IsType(fixture.HostedCluster.Primary); + var services = primary.ServiceProvider; + var grainType = services.GetRequiredService() + .GetGrainType(typeof(ExplicitlyRegisteredSimpleDIGrain)); + var grainId = GrainId.Create(grainType, Guid.NewGuid().ToString()); + using var migrationContext = new MigrationContext( + services.GetRequiredService()); + migrationContext.AddBytes("test", [1, 2, 3]); + + ConfigurationFailureState.Arm(grainId); + try + { + var exception = Assert.Throws( + () => services.GetRequiredService() + .GetOrCreateActivation(grainId, requestContextData: null, migrationContext)); + Assert.Equal("configuration-fault", exception.Message); + } + finally + { + ConfigurationFailureState.Clear(); + } + + Assert.False(migrationContext.TryGetBytes("test", out _)); + Assert.Null(services.GetRequiredService().FindTarget(grainId)); + } + /// /// Custom grain activator that bypasses dependency injection entirely. /// Implements both IGrainActivator (for creation/disposal) and IConfigureGrainTypeComponents @@ -179,21 +324,114 @@ public bool TryGetConfigurator( } } + private sealed class ConfigurationFailureConfiguratorProvider(GrainClassMap grainClassMap) : IConfigureGrainContextProvider + { + public bool TryGetConfigurator( + GrainType grainType, + GrainProperties properties, + [NotNullWhen(true)] out IConfigureGrainContext? configurator) + { + if (grainClassMap.TryGetGrainClass(grainType, out var grainClass) + && grainClass == typeof(ExplicitlyRegisteredSimpleDIGrain)) + { + configurator = ConfigurationFailureConfigurator.Instance; + return true; + } + + configurator = null; + return false; + } + } + + private sealed class ConfigurationFailureConfigurator : IConfigureGrainContext + { + public static ConfigurationFailureConfigurator Instance { get; } = new(); + + public void Configure(IGrainContext context) + { + if (ConfigurationFailureState.ShouldFail(context.GrainId)) + { + throw new InvalidOperationException("configuration-fault"); + } + } + } + + private static class ConfigurationFailureState + { + private static readonly object Lock = new(); + private static GrainId _grainId; + private static bool _armed; + + public static void Arm(GrainId grainId) + { + lock (Lock) + { + _grainId = grainId; + _armed = true; + } + } + + public static bool ShouldFail(GrainId grainId) + { + lock (Lock) + { + return _armed && _grainId.Equals(grainId); + } + } + + public static void Clear() + { + lock (Lock) + { + _armed = false; + _grainId = default; + } + } + } + private sealed class ActivationOrderingState : IConfigureGrainContext { + private readonly AsyncLocal _transactionState = new(); private int _armed; + private int _hadRuntimeContextAtConstruction; + private int _hadSchedulerAffinityAtConstruction; private int _wasConfiguredAtConstruction; + private object? _requestContextAtConstruction; + private object? _transactionStateAtConstruction; public static ActivationOrderingState Instance { get; } = new(); + public const string RequestContextKey = "activation-ordering"; + + public bool HadRuntimeContextAtConstruction => Volatile.Read(ref _hadRuntimeContextAtConstruction) != 0; + public bool HadSchedulerAffinityAtConstruction => Volatile.Read(ref _hadSchedulerAffinityAtConstruction) != 0; + public object? RequestContextAtConstruction => Volatile.Read(ref _requestContextAtConstruction); + public object? TransactionState => _transactionState.Value; + public object? TransactionStateAtConstruction => Volatile.Read(ref _transactionStateAtConstruction); public bool WasConfiguredAtConstruction => Volatile.Read(ref _wasConfiguredAtConstruction) != 0; public void Arm() { + Volatile.Write(ref _hadRuntimeContextAtConstruction, 0); + Volatile.Write(ref _hadSchedulerAffinityAtConstruction, 0); + Volatile.Write(ref _requestContextAtConstruction, null); + Volatile.Write(ref _transactionStateAtConstruction, null); Volatile.Write(ref _wasConfiguredAtConstruction, 0); Volatile.Write(ref _armed, 1); } + public void SetAmbientState(object value) + { + RequestContext.Set(RequestContextKey, value); + _transactionState.Value = value; + } + + public void ClearAmbientState() + { + RequestContext.Remove(RequestContextKey); + _transactionState.Value = null; + } + public void Configure(IGrainContext context) { if (Volatile.Read(ref _armed) == 0) @@ -214,6 +452,14 @@ public void ObserveConstruction(IGrainContext context) Volatile.Write( ref _wasConfiguredAtConstruction, context.GetComponent() is not null ? 1 : 0); + Volatile.Write( + ref _hadSchedulerAffinityAtConstruction, + ReferenceEquals(((ActivationData)context).ActivationTaskScheduler, TaskScheduler.Current) ? 1 : 0); + Volatile.Write( + ref _hadRuntimeContextAtConstruction, + ReferenceEquals(context, RuntimeContext.Current) ? 1 : 0); + Volatile.Write(ref _requestContextAtConstruction, RequestContext.Get(RequestContextKey)); + Volatile.Write(ref _transactionStateAtConstruction, _transactionState.Value); } }