diff --git a/src/coreclr/System.Private.CoreLib/src/System/Runtime/CompilerServices/AsyncHelpers.CoreCLR.cs b/src/coreclr/System.Private.CoreLib/src/System/Runtime/CompilerServices/AsyncHelpers.CoreCLR.cs index 90f74e4b867457..7bf6956a918132 100644 --- a/src/coreclr/System.Private.CoreLib/src/System/Runtime/CompilerServices/AsyncHelpers.CoreCLR.cs +++ b/src/coreclr/System.Private.CoreLib/src/System/Runtime/CompilerServices/AsyncHelpers.CoreCLR.cs @@ -3,6 +3,7 @@ using System.Diagnostics; using System.Diagnostics.CodeAnalysis; +using System.Diagnostics.Tracing; using System.Reflection; using System.Runtime.InteropServices; using System.Runtime.Versioning; @@ -907,8 +908,7 @@ internal void InstrumentedHandleSuspended(AsyncInstrumentation.Flags flags, ref if (AsyncInstrumentation.IsEnabled.AsyncDebugger(flags)) { Continuation? nextContinuation = state.SentinelContinuation!.Next; - - AsyncDebugger.HandleSuspended(nextContinuation); + AsyncDebugger.HandleSuspended(this, nextContinuation); if (!HandleSuspended(ref state)) { @@ -1078,7 +1078,7 @@ private unsafe void InstrumentedDispatchContinuations(AsyncInstrumentation.Flags asyncDispatcherInfo.NextContinuation = MoveContinuationState(); refDispatcherInfo = &asyncDispatcherInfo; - RuntimeAsyncInstrumentationHelpers.ResumeRuntimeAsyncContext(this, ref asyncDispatcherInfo, flags); + RuntimeAsyncInstrumentationHelpers.ResumeRuntimeAsyncContext(this, ref asyncDispatcherInfo, flags, asyncDispatcherInfo.NextContinuation); while (true) { @@ -1778,7 +1778,7 @@ public static void SyncPointCheck(ref AsyncDispatcherInfo info, AsyncInstrumenta } [MethodImpl(MethodImplOptions.AggressiveInlining)] - public static void ResumeRuntimeAsyncContext(Task task, ref AsyncDispatcherInfo info, AsyncInstrumentation.Flags flags) + public static void ResumeRuntimeAsyncContext(Task task, ref AsyncDispatcherInfo info, AsyncInstrumentation.Flags flags, Continuation? continuation) { info.CurrentTask = task; AsyncProfiler.InitInfo(ref info.AsyncProfilerInfo); @@ -1794,7 +1794,7 @@ public static void ResumeRuntimeAsyncContext(Task task, ref AsyncDispatcherInfo if (AsyncInstrumentation.IsEnabled.AsyncDebugger(flags)) { - AsyncDebugger.ResumeAsyncContext(task); + AsyncDebugger.ResumeAsyncContext(task, continuation); } } } @@ -1951,8 +1951,9 @@ public static void CreateAsyncContext(Task task) TplEventSource.Log.TraceOperationBegin(task.Id, "System.Runtime.CompilerServices.AsyncHelpers+RuntimeAsyncTask", 0); } - public static void ResumeAsyncContext(Task task) + public static void ResumeAsyncContext(Task task, Continuation? continuation) { + OutputTaskWaitEnd(task, continuation); TplEventSource.Log.TraceSynchronousWorkBegin(task.Id, CausalitySynchronousWork.Execution); } @@ -2005,16 +2006,37 @@ public static void CompleteAsyncMethod(Continuation curContinuation) Task.RemoveRuntimeAsyncContinuationTimestamp(curContinuation); } - public static void HandleSuspended(Continuation? nextContinuation) + public static void HandleSuspended(Task task, Continuation? nextContinuation) { if (nextContinuation != null) { Task.TryAddRuntimeAsyncContinuationChainTimestamps(nextContinuation); } + + OutputTaskWaitBegin(task, nextContinuation); + } + + private static void OutputTaskWaitBegin(Task task, Continuation? continuation) + { + if (continuation is RuntimeAsyncTaskContinuation { Task: Task awaitedTask }) + { + TplEventSource log = TplEventSource.Log; + if (log.IsEnabled(EventLevel.Informational, TplEventSource.Keywords.TaskTransfer | TplEventSource.Keywords.Tasks)) + { + log.TaskWaitBegin( + task.m_taskScheduler?.Id ?? TaskScheduler.Default.Id, + task.Id, + awaitedTask.Id, + TplEventSource.TaskWaitBehavior.Asynchronous, + task.Id); + } + } } public static void HandleSuspendedFailed(Task task, Continuation? nextContinuation) { + OutputTaskWaitEnd(task, nextContinuation); + if (nextContinuation != null) { Task.RemoveRuntimeAsyncTask(task, nextContinuation); @@ -2024,6 +2046,21 @@ public static void HandleSuspendedFailed(Task task, Continuation? nextContinuati Task.RemoveRuntimeAsyncTask(task); } } + + private static void OutputTaskWaitEnd(Task task, Continuation? continuation) + { + if (continuation is RuntimeAsyncTaskContinuation { Task: Task awaitedTask }) + { + TplEventSource log = TplEventSource.Log; + if (log.IsEnabled(EventLevel.Verbose, TplEventSource.Keywords.Tasks)) + { + log.TaskWaitEnd( + task.m_taskScheduler?.Id ?? TaskScheduler.Default.Id, + task.Id, + awaitedTask.Id); + } + } + } } } } diff --git a/src/libraries/System.Runtime/tests/System.Threading.Tasks.Tests/System.Runtime.CompilerServices/RuntimeAsyncTests.cs b/src/libraries/System.Runtime/tests/System.Threading.Tasks.Tests/System.Runtime.CompilerServices/RuntimeAsyncTests.cs index 33418d77ced9b0..f562c0529e669a 100644 --- a/src/libraries/System.Runtime/tests/System.Threading.Tasks.Tests/System.Runtime.CompilerServices/RuntimeAsyncTests.cs +++ b/src/libraries/System.Runtime/tests/System.Threading.Tasks.Tests/System.Runtime.CompilerServices/RuntimeAsyncTests.cs @@ -6,6 +6,7 @@ using System.Diagnostics.Tracing; using System.Linq; using System.Reflection; +using System.Threading; using Microsoft.DotNet.RemoteExecutor; using Microsoft.DotNet.XUnitExtensions; using Xunit; @@ -21,6 +22,7 @@ public class RuntimeAsyncTests private static readonly FieldInfo s_continuationTimestampsField = GetCorLibClassStaticField("System.Threading.Tasks.Task", "s_runtimeAsyncContinuationTimestamps"); private static readonly FieldInfo s_activeTasksField = GetCorLibClassStaticField("System.Threading.Tasks.Task", "s_currentActiveTasks"); private static readonly FieldInfo s_activeFlagsField = GetCorLibClassStaticField("System.Runtime.CompilerServices.AsyncInstrumentation", "s_activeFlags"); + private static readonly FieldInfo s_taskIdField = GetCorLibClassInstanceField("System.Threading.Tasks.Task", "m_taskId"); private static readonly object s_debuggerLock = new object(); @@ -93,6 +95,23 @@ private static FieldInfo GetCorLibClassStaticField(string className, string fiel return field; } + private static FieldInfo GetCorLibClassInstanceField(string className, string fieldName) + { + Type? classType = typeof(object).Assembly.GetType(className); + if (classType == null) + { + throw new InvalidOperationException($"Type '{className}' doesn't exist in System.Private.CoreLib."); + } + + FieldInfo? field = classType.GetField(fieldName, BindingFlags.Public | BindingFlags.NonPublic | BindingFlags.Instance); + if (field == null) + { + throw new InvalidOperationException($"Expected instance field '{fieldName}' to exist on type '{className}'."); + } + + return field; + } + private static int GetTaskTimestampCount() => s_taskTimestampsField.GetValue(null) is Dictionary dict ? dict.Count : 0; @@ -171,6 +190,12 @@ static async Task FuncThatWaitsTwice(TaskCompletionSource tcs1, TaskCompletionSo await tcs2.Task; } + [System.Runtime.CompilerServices.RuntimeAsyncMethodGeneration(true)] + static async Task FuncThatWaitsOnce(Task task) + { + await task; + } + [System.Runtime.CompilerServices.RuntimeAsyncMethodGeneration(true)] static async Task FuncThatInspectsContinuationTimestamps(TaskCompletionSource tcs, Action callback) { @@ -405,8 +430,8 @@ public void RuntimeAsync_TimestampsTrackedWhileInFlight() { AttachDebugger(); - var tcs1 = new TaskCompletionSource(); - var tcs2 = new TaskCompletionSource(); + var tcs1 = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var tcs2 = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); Task inflight = FuncThatWaitsTwice(tcs1, tcs2); // Task is suspended on tcs1 — should be in active tasks @@ -575,6 +600,169 @@ public void RuntimeAsync_TplEvents() }).Dispose(); } + [ConditionalFact(typeof(RuntimeAsyncTests), nameof(IsRemoteExecutorAndRuntimeAsyncSupported))] + public void RuntimeAsync_TaskWaitEvents() + { + RemoteExecutor.Invoke(() => + { + const int TaskWaitBeginId = 10; + const int TaskWaitEndId = 11; + + AttachDebugger(); + + var events = new ConcurrentQueue(); + using (var listener = new TestEventListener("System.Threading.Tasks.TplEventSource", EventLevel.Verbose)) + { + listener.RunWithCallback(events.Enqueue, () => + { + var tcs1 = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var tcs2 = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + Task runtimeAsyncTask = FuncThatWaitsTwice(tcs1, tcs2); + int runtimeAsyncTaskId = runtimeAsyncTask.Id; + int firstAwaitedTaskId = tcs1.Task.Id; + int secondAwaitedTaskId = tcs2.Task.Id; + + Assert.True( + SpinWait.SpinUntil( + () => events.Any(e => + e.EventId == TaskWaitBeginId && + (int)e.Payload![1]! == runtimeAsyncTaskId && + (int)e.Payload[2]! == firstAwaitedTaskId && + (int)e.Payload[4]! == runtimeAsyncTaskId), + TimeSpan.FromSeconds(30)), + "Expected the RuntimeAsync task to suspend on the first task."); + Assert.DoesNotContain(events, e => + e.EventId == TaskWaitEndId && + (int)e.Payload![1]! == runtimeAsyncTaskId && + (int)e.Payload![2]! == firstAwaitedTaskId); + + tcs1.SetResult(); + + Assert.True( + SpinWait.SpinUntil( + () => events.Any(e => + e.EventId == TaskWaitBeginId && + (int)e.Payload![1]! == runtimeAsyncTaskId && + (int)e.Payload[2]! == secondAwaitedTaskId), + TimeSpan.FromSeconds(30)), + "Expected the RuntimeAsync task to suspend on the second task."); + Assert.Contains(events, e => + e.EventId == TaskWaitEndId && + (int)e.Payload![1]! == runtimeAsyncTaskId && + (int)e.Payload[2]! == firstAwaitedTaskId); + Assert.Contains(events, e => + e.EventId == TaskWaitBeginId && + (int)e.Payload![1]! == runtimeAsyncTaskId && + (int)e.Payload[2]! == secondAwaitedTaskId && + (int)e.Payload[4]! == runtimeAsyncTaskId); + Assert.DoesNotContain(events, e => + e.EventId == TaskWaitEndId && + (int)e.Payload![1]! == runtimeAsyncTaskId && + (int)e.Payload![2]! == secondAwaitedTaskId); + + tcs2.SetResult(); + runtimeAsyncTask.GetAwaiter().GetResult(); + + Assert.True( + SpinWait.SpinUntil( + () => events.Any(e => + e.EventId == TaskWaitEndId && + (int)e.Payload![1]! == runtimeAsyncTaskId && + (int)e.Payload[2]! == secondAwaitedTaskId), + TimeSpan.FromSeconds(30)), + "Expected the RuntimeAsync task to finish waiting on the second task."); + Assert.Contains(events, e => + e.EventId == TaskWaitEndId && + (int)e.Payload![1]! == runtimeAsyncTaskId && + (int)e.Payload[2]! == secondAwaitedTaskId); + }); + } + + DetachDebugger(); + }).Dispose(); + } + + [ConditionalFact(typeof(RuntimeAsyncTests), nameof(IsRemoteExecutorAndRuntimeAsyncSupported))] + public void RuntimeAsync_TaskWaitEndEventsForFaultedAndCanceledTasks() + { + RemoteExecutor.Invoke(() => + { + const int TaskWaitBeginId = 10; + const int TaskWaitEndId = 11; + + AttachDebugger(); + + var events = new ConcurrentQueue(); + using (var listener = new TestEventListener("System.Threading.Tasks.TplEventSource", EventLevel.Verbose)) + { + listener.RunWithCallback(events.Enqueue, () => + { + foreach (bool cancel in new[] { false, true }) + { + var tcs = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + Task runtimeAsyncTask = FuncThatWaitsOnce(tcs.Task); + int runtimeAsyncTaskId = runtimeAsyncTask.Id; + int awaitedTaskId = tcs.Task.Id; + + Assert.True( + SpinWait.SpinUntil( + () => events.Any(e => + e.EventId == TaskWaitBeginId && + (int)e.Payload![1]! == runtimeAsyncTaskId && + (int)e.Payload[2]! == awaitedTaskId), + TimeSpan.FromSeconds(30)), + $"Expected a wait-begin event for the {(cancel ? "canceled" : "faulted")} task."); + + if (cancel) + { + tcs.SetCanceled(); + Assert.ThrowsAny(() => runtimeAsyncTask.GetAwaiter().GetResult()); + } + else + { + tcs.SetException(new InvalidOperationException()); + Assert.Throws(() => runtimeAsyncTask.GetAwaiter().GetResult()); + } + + Assert.True( + SpinWait.SpinUntil( + () => events.Any(e => + e.EventId == TaskWaitEndId && + (int)e.Payload![1]! == runtimeAsyncTaskId && + (int)e.Payload[2]! == awaitedTaskId), + TimeSpan.FromSeconds(30)), + $"Expected a wait-end event for the {(cancel ? "canceled" : "faulted")} task."); + } + }); + } + + DetachDebugger(); + }).Dispose(); + } + + [ConditionalFact(typeof(RuntimeAsyncTests), nameof(IsRemoteExecutorAndRuntimeAsyncSupported))] + public void RuntimeAsync_TaskWaitEventsArePayForPlay() + { + RemoteExecutor.Invoke(() => + { + AttachDebugger(); + + var tcs1 = new TaskCompletionSource(); + var tcs2 = new TaskCompletionSource(); + Task runtimeAsyncTask = FuncThatWaitsTwice(tcs1, tcs2); + + Assert.Equal(0, (int)s_taskIdField.GetValue(tcs1.Task)!); + + tcs1.SetResult(); + Assert.Equal(0, (int)s_taskIdField.GetValue(tcs2.Task)!); + + tcs2.SetResult(); + runtimeAsyncTask.GetAwaiter().GetResult(); + + DetachDebugger(); + }).Dispose(); + } + [ConditionalFact(typeof(RuntimeAsyncTests), nameof(IsRemoteExecutorAndRuntimeAsyncSupported))] public void RuntimeAsync_SuspensionTimestampsAreDistinct() { @@ -601,6 +789,8 @@ public void RuntimeAsync_NoTplEventsWithoutDebugger() { RemoteExecutor.Invoke(() => { + const int TaskWaitBeginId = 10; + const int TaskWaitEndId = 11; const int TraceOperationBeginId = 14; const string RuntimeAsyncTaskOperationName = "System.Runtime.CompilerServices.AsyncHelpers+RuntimeAsyncTask"; @@ -608,6 +798,9 @@ public void RuntimeAsync_NoTplEventsWithoutDebugger() // The AsyncDebugger guard should prevent the V2 async instrumentation // from emitting any TPL causality events. var events = new ConcurrentQueue(); + int runtimeAsyncTaskId = 0; + int firstAwaitedTaskId = 0; + int secondAwaitedTaskId = 0; using (var listener = new TestEventListener("System.Threading.Tasks.TplEventSource", EventLevel.Verbose)) { listener.RunWithCallback(events.Enqueue, () => @@ -616,9 +809,28 @@ public void RuntimeAsync_NoTplEventsWithoutDebugger() { Func().GetAwaiter().GetResult(); } + + var tcs1 = new TaskCompletionSource(); + var tcs2 = new TaskCompletionSource(); + Task runtimeAsyncTask = FuncThatWaitsTwice(tcs1, tcs2); + runtimeAsyncTaskId = runtimeAsyncTask.Id; + firstAwaitedTaskId = tcs1.Task.Id; + secondAwaitedTaskId = tcs2.Task.Id; + tcs1.SetResult(); + tcs2.SetResult(); + runtimeAsyncTask.GetAwaiter().GetResult(); }); } + Assert.DoesNotContain(events, e => + e.EventId == TaskWaitBeginId && + (int)e.Payload![1]! == runtimeAsyncTaskId && + ((int)e.Payload[2]! == firstAwaitedTaskId || (int)e.Payload[2]! == secondAwaitedTaskId)); + Assert.DoesNotContain(events, e => + e.EventId == TaskWaitEndId && + (int)e.Payload![1]! == runtimeAsyncTaskId && + ((int)e.Payload[2]! == firstAwaitedTaskId || (int)e.Payload[2]! == secondAwaitedTaskId)); + // TraceOperationBegin with the RuntimeAsyncTask operation name is uniquely // emitted by V2 async instrumentation. It must not appear without a debugger. Assert.DoesNotContain(events, e =>